Skip to content

Commit 436128f

Browse files
Elton Hometa-codesync[bot]
authored andcommitted
Design doc: hook-customizable in-kernel collective framework
Summary: Overall design doc (README) for the comms/dsl framework, landing first in the stack so the design is reviewed before the code. Covers the design principles - the framework owns the generic 95% (schedule, multi-peer addressing, signal/wait), the kernel owner writes the 5% (a per-tile hook + a transport); performance is autotuned, not hand-tuned; spectrum of control - plus two worked examples: 1) a custom collective (a2a non-contig) built on the framework in ~10 lines, and 2) the autotuner workflow that populates the optimal kernel config and re-tunes on shape changes with no kernel edits (target: ~5 min when TBD change their shapes/dimensions). Differential Revision: D108105252
1 parent 0036d6b commit 436128f

1 file changed

Lines changed: 112 additions & 0 deletions

File tree

comms/dsl/README.md

Lines changed: 112 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,112 @@
1+
# comms/dsl: customized collectives without the plumbing - a hook in, an autotuned kernel out
2+
3+
`comms/dsl` framework lets a kernel developer (user) build a customized, autotunable collective in DSLs (Triton or CuTe) by writing only the small part that differentiates their kernel: per-tile hooks and optionally transports, while the framework owns the generic ~95% (schedule, multi-peer addressing, signal/wait protocol) and automates the performance tuning.
4+
5+
## Why
6+
7+
Researchers prototype new ideas as kernels in DSLs like Triton/CuTe instead of CUDA/C++ for the fast iteration loop. Scaling the idea needs communication, and that usually leads to repetitive work: days-to-weeks of race-prone stage/signal/wait/fence logic, re-derived per kernel, with bugs that could pass tests and fail at scale. Notably, the part that differentiates the kernel is typically small: a transpose, a quantize, or an accumulate, while a lot around it is plumbing.
8+
9+
For instance, the `all_to_all_single_non_contig` kernel is ~800 lines, of which roughly 95% is generic machinery, including the signaling protocol, the pipeline, symmetric-memory staging, multi-peer addressing, and tile-size math. The remaining ~5% is layout transform, which is what really makes this kernel unique.
10+
11+
Furthermore, reaching correctness is only half the journey. Each kernel is then tuned against production shapes maintained per hardware tier, because the launch parameters that win on one shape/hardware may lose on the others.
12+
13+
comms/dsl breaks both loops. It provides the customizable boilerplate, and it turns performance into an autotuner run.
14+
15+
## The pieces
16+
17+
| Piece | What it is |
18+
|---|---|
19+
| **Transport** | A fabric binding (NVLink today, IB later): cross-GPU memory + signaling state, created once via a `rendezvous` over PyTorch symmetric memory. The framework ships defaults and the user either sets one up or could provide their own. |
20+
| **Endpoint (`PeerEndpoint`)** | How the transport hands a schedule per-peer addresses - `send_dst` (where to write into a peer), `recv_src` (where to read its data), and the signal slots - so no schedule touches raw pointers. Two forms: a host-resolved `PeerEndpoint` for a single peer (used by `send_tile`/`recv_tile`), and a device-side table of all peers that a fused collective indexes by peer in-kernel (e.g. `all_to_all`). |
21+
| **Ops** | The transport's four device functions that act on those addresses: `put` (remote write), `get` (local read), `signal`/`wait` (the data-ready handshake). The only fabric-specific (NVLink vs IB) device code; hooks and schedules call them and stay transport-agnostic. |
22+
| **Hook + `Ctx`** | The 5% the user writes: `produce(ctx) -> tile` and `consume(ctx, tile)`. `Ctx` is the per-tile view the framework hands the hook - input/output pointer, flat index, mask, and within-chunk position. This is where a transpose / gather / quantize / accumulate goes. |
23+
| **Collective** | The shipped schedule (e.g. `all_to_all`) that runs the hook over the fused `peer x block` transfer, owning all addressing and signal/wait. |
24+
| **`Config` + `Key`** | `Config` = the launch tunables the autotuner sweeps; `Key` = how a tuned config is looked up at runtime. Both are defined by the user (see principles). |
25+
| **Adapter** | A thin class that plugs a collective into the comms-owned tuner engine (`comm_tuning`). |
26+
27+
Flow: a `Transport` gives the kernel per-peer addresses; the `Collective` loops tiles, calling the user's `Hook` (compute) and the `ops` (move + signal); a `Config`, looked up by `Key`, sets the launch parameters.
28+
29+
## Design principles
30+
31+
- Own the 95%, expose the 5%. The framework owns the schedule and the race-prone protocol; the user writes a per-tile `produce`/`consume` hook (transpose, gather, quantize, accumulate) and supplies a transport. No fences, no wait loops, no deadlock reasoning in user code.
32+
- Performance is autotuned, and the tuner is generic. The user declares two things: a `Key` (which input properties should map to one tuned config - size, dtype, world size, layout, whatever matters for that kernel) and a `Config` (the launch knobs to sweep). A shared engine sweeps configs offline and emits a `{Key: Config}` map; at runtime the collective rebuilds the same key and looks it up (safe default if absent). Changing shapes means re-running the tuner, not rewriting the kernel.
33+
- Spectrum of control. Plug-and-play collective (write a hook) -> `send_tile`/`recv_tile` (keep your own schedule) -> raw `put`/`get`/`signal`/`wait` ops (full control). Climb only as high as you need.
34+
- Contracts shared, code per-DSL. Triton and CuTe share the same op names, hook roles, and Config/Key shapes; the device kernels are written per DSL (flat pointers vs partitions), since device code is not portable across DSLs.
35+
36+
## Example: a custom collective (a2a non-contig), end to end
37+
38+
The whole workflow for a non-contiguous all-to-all - write it, then autotune it.
39+
40+
### 1. Write it (the ~5%)
41+
42+
A transport, one hook, one call:
43+
44+
```python
45+
from comms.dsl import nvl_rendezvous
46+
from comms.dsl.triton import all_to_all
47+
from comms.dsl.triton.hooks import transpose_produce # the 5%: the layout transform
48+
49+
t = nvl_rendezvous(group, dev, per_peer_bytes=chunk_bytes) # transport (staging + signaling)
50+
all_to_all(t, out, inp, produce=transpose_produce, rows=R) # config=None -> tuned lookup
51+
```
52+
53+
`nvl_rendezvous` allocates the per-peer staging buffer + signal pad once; `rows` describes the 2D tile layout the hook reads; with no `produce` hook this is a plain all-to-all.
54+
55+
The hook is the 5%. `ctx` is the per-tile view (input pointer, flat index, mask, within-chunk position `pos` with `rows`/`cols`); it returns the tile to send - here, the transposed source position, fused into the transfer leg with no extra HBM pass:
56+
57+
```python
58+
@triton.jit
59+
def transpose_produce(ctx):
60+
r = ctx.pos % ctx.rows # position in the transposed [cols, rows] layout
61+
c = ctx.pos // ctx.rows
62+
src = r * ctx.cols + c # source position in the [rows, cols] layout
63+
base = ctx.idx - ctx.pos # chunk base in the input
64+
return tl.load(ctx.ptr + base + src, mask=ctx.mask)
65+
```
66+
67+
The framework runs the fused `peer x block` schedule (multi-peer addressing, signal/wait) in one launch, validated bit-exact against `dist.all_to_all_single` on 4 ranks.
68+
69+
### 2. Autotune it
70+
71+
Write one thin adapter - the engine is generic, so you only declare your `Key` and the `Config`s to sweep:
72+
73+
```python
74+
class A2ATuningAdapter(CommKernelTuningAdapter):
75+
def enumerate_input_specs(self, world_size): ... # the shapes to tune
76+
def make_key(self, spec, world_size): ... # YOUR key (must match runtime)
77+
def enumerate_candidate_configs(self, spec, key): ... # YOUR configs to sweep
78+
def make_inputs(self, spec, *, rank, world_size, device): ... # tensors + nvl_rendezvous
79+
def run_candidate(self, inputs, config, group): # the collective with the tuned config
80+
all_to_all(inputs.transport, inputs.output, inputs.input, rows=inputs.rows, config=config)
81+
return inputs.output
82+
def run_baseline(self, inputs, baseline, group): ... # e.g. NCCL, the speed baseline
83+
def check_correctness(self, candidate, reference): ...
84+
```
85+
86+
Parent mode sweeps every candidate, benchmarks it against your baseline, and checks correctness; select mode picks the winner per key and writes a generated table:
87+
88+
```python
89+
# comms/dsl/triton/generated/a2a_tuned_configs.py (generated; do not hand-edit)
90+
from comms.dsl.triton.collectives_tuning import A2AConfig, A2AKey
91+
92+
TUNED_A2A_CONFIGS = {
93+
A2AKey(world_size=8, dtype="bfloat16", numel=8192, rows=0, transport_kind="NvlTransport"):
94+
A2AConfig(num_blocks=16, block=2048, num_warps=8),
95+
# ... one row per tuned key ...
96+
}
97+
```
98+
99+
Then nothing else changes: the Step 1 call `all_to_all(..., config=None)` rebuilds the key and looks it up, so it now runs the tuned config automatically. When shapes change, re-run the tuner and regenerate the table - no kernel edits.
100+
101+
## What the same pattern enables next
102+
103+
Concrete doables on the shipped transport + ops + Config/Key + tuner:
104+
105+
- reduce-scatter: an accumulate-on-`consume` hook (sum the received tile into the output) over the same schedule family.
106+
- all-gather: the gather schedule with the default copy hook.
107+
- quantized all-to-all: a `produce` hook that casts/scales to fp8 and a `consume` that dequantizes, fusing the conversion into the transfer leg (no extra HBM pass).
108+
- variable-split a2a (MoE dispatch / combine): the same collective with per-rank split sizes added to the `Key` and layout.
109+
110+
Each is a new hook and/or a sibling schedule on the same substrate - not a new comm stack.
111+
112+
Further out, and a bigger milestone than the items above: comm-compute overlap - hiding a compute kernel (e.g. a GEMM or MoE expert) under the collective via pipelining to cover its latency.

0 commit comments

Comments
 (0)