diff --git a/gemm_scatter_reduce/tle_inter_node_reduce_scatter.py b/gemm_scatter_reduce/tle_inter_node_reduce_scatter.py new file mode 100644 index 0000000000..afd75f58aa --- /dev/null +++ b/gemm_scatter_reduce/tle_inter_node_reduce_scatter.py @@ -0,0 +1,546 @@ +from __future__ import annotations + +import dataclasses +from typing import Optional + +import torch +import triton +import triton.language as tl +import triton.experimental.tle.language as tle + + +# Reduce-Scatter context: stores GPU/node organization, buffers, communication +# pointers, and CUDA streams. +@dataclasses.dataclass +class TleReduceScatter2DContext: + + # Auto-generated constructor. + max_M: int + N: int + rank: int + world_size: int + local_world_size: int + dtype: torch.dtype + with_gemm_output: bool + + gemm_out_dev_comm_ptr: Optional[int] + gemm_out_dev_mem_ptr: Optional[int] + gemm_out_buf: Optional[torch.Tensor] + scatter_buf: torch.Tensor + rs_per_node_buf: torch.Tensor + p2p_buf: torch.Tensor + signal_buf: torch.Tensor + scatter_dev_comm_ptr: int + scatter_dev_mem_ptr: int + rs_per_node_dev_comm_ptr: int + rs_per_node_dev_mem_ptr: int + p2p_dev_comm_ptr: int + p2p_dev_mem_ptr: int + signal_dev_comm_ptr: int + signal_dev_mem_ptr: int + + reduction_stream: torch.cuda.Stream + p2p_stream: torch.cuda.Stream + num_sync_sms: int + num_p2p_sms: int + num_reduction_sms: int + num_scatter_sms: int + + scatter_signal_buf: torch.Tensor = dataclasses.field(init=False) + rs_per_node_signal_buf: torch.Tensor = dataclasses.field(init=False) + local_rank: int = dataclasses.field(init=False) + node_id: int = dataclasses.field(init=False) + nnodes: int = dataclasses.field(init=False) + _finalized: bool = dataclasses.field(init=False, default=False) + + # Compute local_rank, node_id, nnodes, scatter_signal_buf, and validate + # arguments. + def __post_init__(self): + if self.world_size < 2: + raise ValueError("TLE reduce-scatter requires at least two GPUs") + if self.world_size % self.local_world_size: + raise ValueError("world_size must be divisible by local_world_size") + if self.max_M % self.world_size: + raise ValueError("max_M must be divisible by world_size") + self.local_rank = self.rank % self.local_world_size + self.node_id = self.rank // self.local_world_size + self.nnodes = self.world_size // self.local_world_size + if self.nnodes != 1: + raise NotImplementedError("TLE has no node-level remote primitive yet; this example only " + "implements the nnodes == 1 specialization") + if self.local_rank != self.rank: + raise ValueError("single-node rank must equal local_rank") + if self.signal_buf.numel() < 2 * self.world_size: + raise ValueError("signal_buf must contain scatter and per-node signals") + if self.num_scatter_sms < 1: + raise ValueError("num_scatter_sms must be positive") + + # Take the first world_size elements of signal_buf as scatter_signal_buf, + # and the next world_size elements as rs_per_node_signal_buf. + self.scatter_signal_buf = self.signal_buf[:self.world_size] + # The last world_size elements of signal_buf are used as + # rs_per_node_signal_buf for multi-node. + self.rs_per_node_signal_buf = self.signal_buf[self.world_size:2 * self.world_size] + + @property + def num_rs_sms(self) -> int: + if self.nnodes == 1: + return self.num_scatter_sms + return (self.num_scatter_sms + self.num_sync_sms + self.num_p2p_sms + self.num_reduction_sms) + + def finalize(self): + """Collectively release TLE/FlagCX after the context is no longer used.""" + if self._finalized: + return + torch.cuda.synchronize() + tle.cleanup_communicator() + self._finalized = True + + def reset_barriers(self): + self.signal_buf.zero_() + + +# Intra-node scatter: write each target shard from this rank to the target rank. +@triton.jit +def _scatter_kernel( + input_ptr, + local_scatter_ptr, + dev_mem_ptr, + ready_ptr, + M_per_rank, + N, + LOCAL_RANK: tl.constexpr, + WORLD_SIZE: tl.constexpr, + SCATTER_NODE_SLICE_OFFSET_ELEMS: tl.constexpr, + WAIT_FOR_READY: tl.constexpr, + BLOCK_M: tl.constexpr, + BLOCK_N: tl.constexpr, +): + + pid = tl.program_id(0) + num_pid = tl.num_programs(0) + num_tiles_m = tl.cdiv(M_per_rank, BLOCK_M) + num_tiles_n = tl.cdiv(N, BLOCK_N) + tiles_per_peer = num_tiles_m * num_tiles_n + row_offs = tl.arange(0, BLOCK_M) + col_offs = tl.arange(0, BLOCK_N) + + # On destination rank t, source rank LOCAL_RANK owns slot LOCAL_RANK. + source_slot_offset_elems = LOCAL_RANK * M_per_rank * N + for step in range(WORLD_SIZE): + target_rank = (LOCAL_RANK + step + 1) % WORLD_SIZE + + # ready_ptr + target_rank points to the ready flag for the target rank + # on the current GPU. Poll until the GEMM data produced by this GPU for + # target_rank is complete. + if WAIT_FOR_READY: + while tl.atomic_add(ready_ptr + target_rank, 0, sem="acquire", scope="gpu") == 0: + pass + + if target_rank == LOCAL_RANK: + remote_base = local_scatter_ptr + source_slot_offset_elems + else: + remote_base = tle.remote( + dev_mem_ptr, + space="device", + dtype=input_ptr.dtype.element_ty, + shard_id=target_rank, + offset=SCATTER_NODE_SLICE_OFFSET_ELEMS + source_slot_offset_elems, + ) + + for local_tile in range(pid, tiles_per_peer, num_pid): + tile_m = local_tile // num_tiles_n + tile_n = local_tile % num_tiles_n + input_row = target_rank * M_per_rank + tile_m * BLOCK_M + input_col = tile_n * BLOCK_N + input_ptrs = (input_ptr + (input_row + row_offs[:, None]) * N + input_col + col_offs[None, :]) + row_mask = (input_row + row_offs[:, None]) < (target_rank + 1) * M_per_rank + col_mask = (input_col + col_offs[None, :]) < N + values = tl.load(input_ptrs, mask=row_mask & col_mask, other=0.0) + + scatter_row = tile_m * BLOCK_M + scatter_ptrs = (remote_base + (scatter_row + row_offs[:, None]) * N + input_col + col_offs[None, :]) + scatter_row_mask = (scatter_row + row_offs[:, None]) < M_per_rank + tl.store(scatter_ptrs, values, mask=scatter_row_mask & col_mask) + + +# Device-side intra-node barrier: wait for all GPUs' scatter writes to become +# visible. +@triton.jit +def _device_barrier_kernel(dev_comm_ptr, WORLD_SIZE: tl.constexpr): + # This is a FlagCX/TLE intra-node device barrier, not a CUDA CTA barrier. + tle.distributed_barrier( + comm_ptr=dev_comm_ptr, + space="device", + group_kind="block", + barrier_kind="sync", + order="acqrel", + index=0, + ) + + +# TMA reduction: read each source on this GPU and write the summed result. +@triton.jit +def _ring_reduce_tma_kernel( + local_scatter_ptr, + output_ptr, + M_per_rank, + N, + LOCAL_RANK: tl.constexpr, + WORLD_SIZE: tl.constexpr, + BLOCK_M: tl.constexpr, + BLOCK_N: tl.constexpr, +): + + pid = tl.program_id(0) + num_pid = tl.num_programs(0) + num_tiles_m = tl.cdiv(M_per_rank, BLOCK_M) + num_tiles_n = tl.cdiv(N, BLOCK_N) + total_tiles = num_tiles_m * num_tiles_n + scatter_desc = tl.make_tensor_descriptor( + local_scatter_ptr, + shape=[M_per_rank * WORLD_SIZE, N], + strides=[N, 1], + block_shape=[BLOCK_M, BLOCK_N], + ) + output_desc = tl.make_tensor_descriptor( + output_ptr, + shape=[M_per_rank, N], + strides=[N, 1], + block_shape=[BLOCK_M, BLOCK_N], + ) + + for tile_id in range(pid, total_tiles, num_pid): + tile_m = tile_id // num_tiles_n + tile_n = tile_id % num_tiles_n + row = tile_m * BLOCK_M + col = tile_n * BLOCK_N + source_rank = (LOCAL_RANK + 1) % WORLD_SIZE + accum = scatter_desc.load([row + source_rank * M_per_rank, col]) + + for i in range(1, WORLD_SIZE): + source_rank = (LOCAL_RANK + i + 1) % WORLD_SIZE + accum += scatter_desc.load([row + source_rank * M_per_rank, col]) + output_desc.store([row, col], accum) + + +# Context factory: allocate and register all communication buffers. +def create_tle_reduce_scatter_2d_ctx(max_M: int, N: int, rank: int, world_size: int, local_world_size: int, + dtype: torch.dtype, with_gemm_output: bool = False, + reduction_stream: Optional[torch.cuda.Stream] = None, num_reduction_sms: int = 15, + num_scatter_sms: int = 16) -> TleReduceScatter2DContext: + + if world_size < 2: + raise ValueError("TLE reduce-scatter requires at least two GPUs") + if world_size % local_world_size: + raise ValueError("world_size must be divisible by local_world_size") + if max_M % world_size: + raise ValueError("max_M must be divisible by world_size") + if world_size // local_world_size != 1: + raise NotImplementedError("TLE has no node-level remote primitive yet; this example only " + "implements the nnodes == 1 specialization") + + # Buffer allocation. + per_node_rows = max_M // local_world_size + with torch.cuda.use_mem_pool(tle.get_mem_pool()): + gemm_out_buf = (torch.empty((max_M, N), dtype=dtype, device="cuda") if with_gemm_output else None) + scatter_buf = torch.empty((max_M, N), dtype=dtype, device="cuda") + rs_per_node_buf = torch.empty((per_node_rows, N), dtype=dtype, device="cuda") + p2p_buf = torch.empty((per_node_rows, N), dtype=dtype, device="cuda") + signal_buf = torch.empty((2 * world_size, ), dtype=torch.int32, device="cuda") + signal_buf.zero_() + + # Create multiple communication windows. + scatter_dev_comm_ptr, scatter_dev_mem_ptr = tle.create_comm_tensor(scatter_buf) + + rs_per_node_dev_comm_ptr, rs_per_node_dev_mem_ptr = tle.create_comm_tensor(rs_per_node_buf) + + p2p_dev_comm_ptr, p2p_dev_mem_ptr = tle.create_comm_tensor(p2p_buf) + + signal_dev_comm_ptr, signal_dev_mem_ptr = tle.create_comm_tensor(signal_buf) + + gemm_out_dev_comm_ptr, gemm_out_dev_mem_ptr = ((None, None) + if gemm_out_buf is None else tle.create_comm_tensor(gemm_out_buf)) + return TleReduceScatter2DContext( + max_M=max_M, + N=N, + rank=rank, + world_size=world_size, + local_world_size=local_world_size, + dtype=dtype, + with_gemm_output=with_gemm_output, + gemm_out_buf=gemm_out_buf, + scatter_buf=scatter_buf, + gemm_out_dev_comm_ptr=gemm_out_dev_comm_ptr, + gemm_out_dev_mem_ptr=gemm_out_dev_mem_ptr, + rs_per_node_buf=rs_per_node_buf, + p2p_buf=p2p_buf, + signal_buf=signal_buf, + scatter_dev_comm_ptr=scatter_dev_comm_ptr, + scatter_dev_mem_ptr=scatter_dev_mem_ptr, + rs_per_node_dev_comm_ptr=rs_per_node_dev_comm_ptr, + rs_per_node_dev_mem_ptr=rs_per_node_dev_mem_ptr, + p2p_dev_comm_ptr=p2p_dev_comm_ptr, + p2p_dev_mem_ptr=p2p_dev_mem_ptr, + signal_dev_comm_ptr=signal_dev_comm_ptr, + signal_dev_mem_ptr=signal_dev_mem_ptr, + reduction_stream=(reduction_stream if reduction_stream is not None else torch.cuda.Stream(priority=-1)), + p2p_stream=torch.cuda.Stream(priority=-1), + num_sync_sms=0, + num_p2p_sms=1, + num_reduction_sms=num_reduction_sms, + num_scatter_sms=num_scatter_sms, + ) + + +def _set_tma_allocator(): + + def alloc_fn(size: int, alignment: int, stream: Optional[int]): + return torch.empty(size, device="cuda", dtype=torch.int8) + + triton.set_allocator(alloc_fn) + + +# Perform intra-node scatter and partial reduction per target node. +# Per-target-node local Reduce-Scatter and P2P. +def reduce_scatter_for_each_node(input_tensor: torch.Tensor, stream: torch.cuda.Stream, ctx: TleReduceScatter2DContext, + ready_flags: Optional[torch.Tensor] = None) -> torch.Tensor: + + # Intra-node Reduce-Scatter and the subsequent inter-node P2P. + world_size = ctx.world_size + local_world_size = ctx.local_world_size + local_rank = ctx.local_rank + reduction_stream = ctx.reduction_stream + num_reduction_sms = ctx.num_reduction_sms + nnodes = ctx.nnodes + node_id = ctx.node_id + rs_per_node_buf = ctx.rs_per_node_buf + p2p_buf = ctx.p2p_buf + M, N = input_tensor.shape + M_per_rank = M // world_size + M_per_node = M_per_rank * local_world_size + + # Set the number of scatter CTAs. + scatter_grid = lambda META: (min( + triton.cdiv(M_per_rank, META["BLOCK_M"]) * triton.cdiv(N, META["BLOCK_N"]), + ctx.num_scatter_sms, + ), ) + + def reduce_launch_config(num_sms: int): + if num_sms == -1: + return (lambda META: (triton.cdiv(M_per_rank, META["BLOCK_M"]) * triton.cdiv(N, META["BLOCK_N"]), )), 64, 4 + return (lambda META: + (min(triton.cdiv(M_per_rank, META["BLOCK_M"]) * triton.cdiv(N, META["BLOCK_N"]), num_sms), )), 128, 8 + + # Plain RS has no upstream GEMM producer, so it does not wait on a signal; + # when fused with GEMM, ready_flags[target_rank] indicates the corresponding + # GEMM tile is complete. + ready_ptr = ready_flags if ready_flags is not None else input_tensor + + with torch.cuda.stream(stream): + for n in range(nnodes): + # Same node-level swizzle as tutorial 06: the final round targets the + # current node. + cur_node_id = (node_id + n + 1) % nnodes + + # Take all rows belonging to the target node; for single-node this is + # the full input. + input_intra_node = input_tensor[cur_node_id * M_per_node:(cur_node_id + 1) * M_per_node] + + # Buffer where the target node's data will be stored. + scatter_for_node = ctx.scatter_buf[cur_node_id * M_per_node:(cur_node_id + 1) * M_per_node] + + # Buffer holding the partial reduction result computed by this node for + # the target node. + rs_per_node_output = rs_per_node_buf[cur_node_id * M_per_rank:(cur_node_id + 1) * M_per_rank] + + # Signals are ordered by global rank; this round uses only the + # local-rank segment for the target node. + + # Starting signal index for the target node. + signal_start = cur_node_id * local_world_size + + # Flag used when coordinating with GEMM. + if ready_flags is not None: + ready_for_node = ready_flags[signal_start:signal_start + local_world_size] + else: + ready_for_node = input_tensor + + # dev_mem_ptr points to the full scatter_buf window; remote access + # first jumps to the staging slice of the current target node. Offset is + # in elements, not bytes. + scatter_node_slice_offset_elems = cur_node_id * M_per_node * N + + # Scatter data for the target node within the current node. + _scatter_kernel[scatter_grid]( + input_intra_node, + scatter_for_node, + ctx.scatter_dev_mem_ptr, + ready_for_node, + M_per_rank, + N, + LOCAL_RANK=local_rank, + WORLD_SIZE=local_world_size, + SCATTER_NODE_SLICE_OFFSET_ELEMS=scatter_node_slice_offset_elems, + WAIT_FOR_READY=ready_flags is not None, + BLOCK_M=256, + BLOCK_N=128, + num_warps=4, + ) + + # Wait for all GPUs' intra-node scatter writes to become visible. + _device_barrier_kernel[(local_world_size, )](ctx.scatter_dev_comm_ptr, WORLD_SIZE=local_world_size) + + # Limit reduction SMs for non-final nodes; the final node uses the full + # tile grid. + node_reduce_sms = (-1 if n == nnodes - 1 else num_reduction_sms) + + reduce_grid, reduce_block_n, reduce_warps = reduce_launch_config(node_reduce_sms) + + if stream is not reduction_stream: + reduction_stream.wait_stream(stream) + + with torch.cuda.stream(reduction_stream): + # Reduce only among the local_world_size GPUs of the current node. + # Result is written to rs_per_node_output. + _ring_reduce_tma_kernel[reduce_grid]( + scatter_for_node, + rs_per_node_output, + M_per_rank, + N, + LOCAL_RANK=local_rank, + WORLD_SIZE=local_world_size, + BLOCK_M=256, + BLOCK_N=reduce_block_n, + num_warps=reduce_warps, + ) + + # TLE node-space remote/P2P primitive is not yet wired up. + if nnodes > 1: + pass + + if stream is not reduction_stream: + stream.wait_stream(reduction_stream) + if nnodes == 1: + return rs_per_node_buf[:M_per_rank * nnodes] + return p2p_buf[:M_per_rank * nnodes] + + +def reduce_scatter_multi_node(input_tensor: torch.Tensor, stream: torch.cuda.Stream, ctx: TleReduceScatter2DContext, + output: torch.Tensor, ready_flags: Optional[torch.Tensor] = None) -> torch.Tensor: + + M, N = input_tensor.shape + M_per_rank = M // ctx.world_size + ctx.p2p_stream.wait_stream(stream) + + # Intra-node reduce-scatter; returns each node's reduction result. In the + # single-node case, this step completes the operation. + rs_result_per_node = reduce_scatter_for_each_node(input_tensor, stream, ctx, ready_flags) + + final_grid = lambda META: (triton.cdiv(M_per_rank, META["BLOCK_M"]) * triton.cdiv(N, META["BLOCK_N"]), ) + + # In single-node case, pass local_rank=0, world_size=1 to copy + # rs_result_per_node directly to output. + with torch.cuda.stream(stream): + _ring_reduce_tma_kernel[final_grid]( + rs_result_per_node, + output, + M_per_rank, + N, + LOCAL_RANK=ctx.node_id, + WORLD_SIZE=ctx.nnodes, + BLOCK_M=256, + BLOCK_N=64, + num_warps=4, + ) + return output + + +# Dispatch function that calls reduce_scatter_multi_node. +def reduce_scatter_2d_op(input_tensor: torch.Tensor, ctx: TleReduceScatter2DContext, + output: Optional[torch.Tensor] = None, + ready_flags: Optional[torch.Tensor] = None) -> torch.Tensor: + + M, N = input_tensor.shape + if input_tensor.dtype != ctx.dtype or N != ctx.N: + raise ValueError("input shape/dtype does not match reduce-scatter context") + if M > ctx.max_M or M % ctx.world_size: + raise ValueError("M must divide world_size and fit in the context") + M_per_rank = M // ctx.world_size + if M_per_rank < 256: + raise ValueError("M_per_rank must be >= 256 for the TMA reduce kernel") + if output is None: + output = torch.empty((M_per_rank, N), dtype=input_tensor.dtype, device=input_tensor.device) + if tuple(output.shape) != (M_per_rank, N): + raise ValueError("output has an invalid reduce-scatter shape") + if ready_flags is not None and ready_flags.numel() != ctx.world_size: + raise ValueError("ready_flags must contain one entry per target rank") + + _set_tma_allocator() + reduction_stream = ctx.reduction_stream + scatter_stream = torch.cuda.current_stream() + if scatter_stream is reduction_stream: + raise ValueError("scatter_stream and reduction_stream must be distinct") + + reduction_stream.wait_stream(scatter_stream) + + with torch.cuda.stream(scatter_stream): + _device_barrier_kernel[(ctx.world_size, )](ctx.scatter_dev_comm_ptr, WORLD_SIZE=ctx.world_size) + + output = reduce_scatter_multi_node(input_tensor, scatter_stream, ctx, output, ready_flags) + + with torch.cuda.stream(scatter_stream): + ctx.reset_barriers() + return output + + +# PyTorch baseline implementation. +def torch_rs(input_tensor: torch.Tensor, TP_GROUP) -> torch.Tensor: + output = torch.empty((input_tensor.shape[0] // TP_GROUP.size(), input_tensor.shape[1]), dtype=input_tensor.dtype, + device=input_tensor.device) + dist.reduce_scatter_tensor(output, input_tensor, group=TP_GROUP) + return output + + +def main(): + # get_mem_pool initializes both NCCL's process group and the FlagCX runtime. + tle.get_mem_pool() + rank = dist.get_rank() + TP_GROUP = dist.group.WORLD + world_size = dist.get_world_size() + local_rank = int(os.environ.get("LOCAL_RANK", rank)) + local_world_size = int(os.environ.get("LOCAL_WORLD_SIZE", world_size)) + torch.cuda.set_device(local_rank) + + if world_size < 2: + print("This example needs at least two GPUs", file=sys.stderr) + return + if local_world_size != world_size: + raise NotImplementedError("TLE node-level reduce-scatter is not implemented") + if torch.cuda.get_device_capability()[0] < 9: + print("Skip: the TMA reduce kernel requires sm90 or newer") + tle.cleanup_communicator() + return + + dtype = torch.bfloat16 + + M, N = 8192, 16384 + # + ctx = create_tle_reduce_scatter_2d_ctx(M, N, rank, world_size, local_world_size, dtype) + input_tensor = torch.rand((M, N), dtype=dtype, device="cuda") + + # PyTorch baseline implementation. + torch_output = torch_rs(input_tensor, TP_GROUP) + torch.cuda.synchronize() + + output = reduce_scatter_2d_op(input_tensor, ctx) + torch.cuda.current_stream().wait_stream(ctx.reduction_stream) + torch.cuda.synchronize() + + torch.testing.assert_close(torch_output, output, atol=6e-2, rtol=6e-2) + print(f"[Rank {rank}] TLE single-node reduce-scatter passed: {tuple(output.shape)}") + ctx.finalize() + + +if __name__ == "__main__": + main() diff --git a/gemm_scatter_reduce/tle_overlapping_gemm_reduce_scatter.py b/gemm_scatter_reduce/tle_overlapping_gemm_reduce_scatter.py new file mode 100644 index 0000000000..c34f80e79c --- /dev/null +++ b/gemm_scatter_reduce/tle_overlapping_gemm_reduce_scatter.py @@ -0,0 +1,423 @@ +from __future__ import annotations + +import dataclasses +import os +import statistics +import sys +from typing import Optional + +import torch +import torch.distributed as dist +import triton +import triton.runtime +import triton.language as tl +import triton.experimental.tle.language as tle + +from tle_inter_node_reduce_scatter import ( + TleReduceScatter2DContext, + create_tle_reduce_scatter_2d_ctx, + reduce_scatter_2d_op, +) + + +@dataclasses.dataclass +class TleGemmReduceScatterContext: + """Tutorial-08-style hierarchical GEMM+RS context.""" + + rs_ctx: TleReduceScatter2DContext + output_dtype: torch.dtype + rs_stream: torch.cuda.Stream + num_gemm_sms: int + BLOCK_M: int = 128 + BLOCK_N: int = 256 + BLOCK_K: int = 64 + GROUP_M: int = 8 + STAGES: int = 3 + + def __post_init__(self): + # rs_stream drives scatter; reduction_stream consumes each completed + # scatter stage. They must be distinct to create the pipeline. + if self.rs_stream is self.rs_ctx.reduction_stream: + raise ValueError("rs_stream and reduction_stream must be distinct") + + def finalize(self): + """Match tutorial 08 by releasing the owned RS context.""" + self.rs_ctx.finalize() + + def get_gemm_out_buf(self, input_tensor: torch.Tensor) -> torch.Tensor: + if self.rs_ctx.gemm_out_buf is None: + raise RuntimeError("GEMM context must reserve a symmetric GEMM output buffer") + return self.rs_ctx.gemm_out_buf[:input_tensor.shape[0]] + + +# Create the context: allocate the GEMM output buffer, scatter signal, and +# Reduce-Scatter buffers. +def create_gemm_rs_context(max_M: int, N: int, rank: int, world_size: int, local_world_size: int, + output_dtype: torch.dtype, rs_stream: torch.cuda.Stream, BLOCK_M: int = 128, + BLOCK_N: int = 256, BLOCK_K: int = 64, GROUP_M: int = 8, STAGES: int = 3, + num_scatter_sms: int = 16) -> TleGemmReduceScatterContext: + """Build GEMM+RS state on the caller-owned RS stream.""" + if max_M % world_size: + raise ValueError("max_M must be divisible by world_size") + + # rs_stream is the scatter stream. The factory creates a separate high + # priority reduction stream, which consumes each completed scatter stage. + rs_ctx = create_tle_reduce_scatter_2d_ctx(max_M, N, rank, world_size, local_world_size, output_dtype, + with_gemm_output=True, num_scatter_sms=num_scatter_sms) + + total_sms = torch.cuda.get_device_properties("cuda").multi_processor_count + num_gemm_sms = total_sms - rs_ctx.num_rs_sms + if num_gemm_sms < 1: + raise ValueError("reduce-scatter SM reservation leaves no SM for GEMM") + + return TleGemmReduceScatterContext(rs_ctx=rs_ctx, output_dtype=output_dtype, rs_stream=rs_stream, + num_gemm_sms=num_gemm_sms, BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K, + GROUP_M=GROUP_M, STAGES=STAGES) + + +# GEMM computation and completion notification. +# counter_ptr = workspace +# ready_ptr = scatter_signal +# scatter_signal has size world_size; each workspace[j] counts how many output +# tiles for rank j have completed. +# +@triton.jit +def kernel_gemm_rs_producer_persistent( + a_ptr, + b_ptr, + c_ptr, + M, + N, + K, + ready_ptr, + counter_ptr, + RANK: tl.constexpr, + LOCAL_WORLD_SIZE: tl.constexpr, + WORLD_SIZE: tl.constexpr, + BLOCK_M: tl.constexpr, + BLOCK_N: tl.constexpr, + BLOCK_K: tl.constexpr, + GROUP_M: tl.constexpr, + NUM_SMS: tl.constexpr, +): + + start_pid = tl.program_id(0) + num_pid_m = tl.cdiv(M, BLOCK_M) + num_pid_n = tl.cdiv(N, BLOCK_N) + k_tiles = tl.cdiv(K, BLOCK_K) + num_tiles = num_pid_m * num_pid_n + node_id = RANK // LOCAL_WORLD_SIZE + nnodes = WORLD_SIZE // LOCAL_WORLD_SIZE + tiles_per_sm = num_tiles // NUM_SMS + if start_pid < num_tiles % NUM_SMS: + tiles_per_sm += 1 + + a_desc = tl.make_tensor_descriptor(a_ptr, shape=[M, K], strides=[K, 1], block_shape=[BLOCK_M, BLOCK_K]) + b_desc = tl.make_tensor_descriptor(b_ptr, shape=[N, K], strides=[K, 1], block_shape=[BLOCK_N, BLOCK_K]) + c_desc = tl.make_tensor_descriptor(c_ptr, shape=[M, N], strides=[N, 1], block_shape=[BLOCK_M, BLOCK_N]) + + M_per_rank = M // WORLD_SIZE + tiles_m_per_rank = M_per_rank // BLOCK_M + tiles_per_group = GROUP_M * num_pid_n + tile_id = start_pid - NUM_SMS + k_tile = -1 + pid_m = 0 + pid_n = 0 + offs_am = 0 + offs_bn = 0 + accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32) + + for _ in range(0, k_tiles * tiles_per_sm): + + k_tile = tl.where(k_tile == k_tiles - 1, 0, k_tile + 1) + + if k_tile == 0: + tile_id += NUM_SMS + group_id = tile_id // tiles_per_group + first_pid_m = group_id * GROUP_M + group_size_m = min(num_pid_m - first_pid_m, GROUP_M) + logical_pid_m = first_pid_m + tile_id % group_size_m + pid_n = (tile_id % tiles_per_group) // group_size_m + + m_rank = logical_pid_m // tiles_m_per_rank + pid_m_intra_rank = logical_pid_m - m_rank * tiles_m_per_rank + m_node_id = m_rank // LOCAL_WORLD_SIZE + m_local_rank = m_rank % LOCAL_WORLD_SIZE + swizzle_m_node_id = (m_node_id + node_id + 1) % nnodes + swizzle_m_local_rank = (m_local_rank + RANK + 1) % LOCAL_WORLD_SIZE + swizzle_m_rank = (swizzle_m_node_id * LOCAL_WORLD_SIZE + swizzle_m_local_rank) + pid_m = swizzle_m_rank * tiles_m_per_rank + pid_m_intra_rank + offs_am = pid_m * BLOCK_M + offs_bn = pid_n * BLOCK_N + + a = a_desc.load([offs_am, k_tile * BLOCK_K]) + b = b_desc.load([offs_bn, k_tile * BLOCK_K]) + accumulator = tl.dot(a, b.T, accumulator) + + # If this is the last K-dimension tile. + if k_tile == k_tiles - 1: + + # Write GEMM output. + c_desc.store([offs_am, offs_bn], accumulator.to(c_ptr.dtype.element_ty)) + + # Determine which target rank's M-row slice this output tile belongs to. + counter_start = offs_am // M_per_rank + counter_end = (offs_am + BLOCK_M - 1) // M_per_rank + counter_end = min(counter_end, WORLD_SIZE - 1) + + for counter_id in range(counter_start, counter_end + 1): + m_start = M_per_rank * counter_id + m_end = M_per_rank * (counter_id + 1) - 1 + tiled_m_start = m_start // BLOCK_M + tiled_m_end = m_end // BLOCK_M + tiled_m_size = tiled_m_end - tiled_m_start + 1 + tiled_n = tl.cdiv(N, BLOCK_N) + + # Increment once per completed output tile. + prior = tl.atomic_add(counter_ptr + counter_id, 1, sem="release", scope="gpu") + + if prior == tiled_m_size * tiled_n - 1: + # Tutorial 08's dl.notify(..., signal=1) sets the signal to + # 1 rather than accumulating on the old value. Use exchange + # to ensure the signal remains a boolean across repeated + # calls, preventing 1 -> 2 from causing the consumer to + # wait forever. + # old value = ready_ptr[counter_id] + # ready_ptr[counter_id] = 1 + tl.atomic_xchg(ready_ptr + counter_id, 1, sem="release", scope="gpu") + + # Reset accumulator to start computing the next assigned output. + accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32) + + +def gemm_rs_producer_persistent(a: torch.Tensor, b: torch.Tensor, c: torch.Tensor, barrier: torch.Tensor, + workspace: torch.Tensor, world_size: int, local_world_size: int, rank: int, + num_gemm_sms: int, BLOCK_SIZE_M: int = 128, BLOCK_SIZE_N: int = 256, + BLOCK_SIZE_K: int = 64, GROUP_SIZE_M: int = 8, STAGES: int = 3): + """Tutorial-08-style producer wrapper around the TLE Triton kernel.""" + if a.shape[1] != b.shape[1]: + raise ValueError("incompatible GEMM dimensions") + if a.dtype != b.dtype: + raise ValueError("GEMM operands must have the same dtype") + M, local_K = a.shape + N = b.shape[0] + M_per_rank = M // world_size + + if M_per_rank % BLOCK_SIZE_M: + raise ValueError("M_per_rank must be aligned to BLOCK_SIZE_M") + + # TMA descriptors require a global-memory allocator, as in tutorial 08. + def alloc_fn(size: int, alignment: int, stream: Optional[int]): + return torch.empty(size, device="cuda", dtype=torch.int8) + + triton.set_allocator(alloc_fn) + grid = lambda META: (min( + num_gemm_sms, + triton.cdiv(M, META["BLOCK_M"]) * triton.cdiv(N, META["BLOCK_N"]), + ), ) + + # stages: pipeline depth parameter for the Triton kernel launch. + return kernel_gemm_rs_producer_persistent[grid]( + a, + b, + c, + M, + N, + local_K, + barrier, + workspace, + RANK=rank, + LOCAL_WORLD_SIZE=local_world_size, + WORLD_SIZE=world_size, + BLOCK_M=BLOCK_SIZE_M, + BLOCK_N=BLOCK_SIZE_N, + BLOCK_K=BLOCK_SIZE_K, + GROUP_M=GROUP_SIZE_M, + NUM_SMS=num_gemm_sms, + num_warps=8, + num_stages=STAGES, + ) + + +def _pad_to_block_m(input_tensor: torch.Tensor, world_size: int, block_m: int) -> torch.Tensor: + M, K = input_tensor.shape + M_per_rank = M // world_size + padded_M_per_rank = triton.cdiv(M_per_rank, block_m) * block_m + if padded_M_per_rank == M_per_rank: + return input_tensor + reshaped = input_tensor.reshape(world_size, M_per_rank, K) + padded = torch.empty((world_size, padded_M_per_rank, K), dtype=input_tensor.dtype, device=input_tensor.device) + padded[:, :M_per_rank].copy_(reshaped) + return padded.reshape(-1, K) + + +# GEMM + Reduce-Scatter combined implementation. +def gemm_rs_multi_node_persistent_op(input_tensor: torch.Tensor, weight: torch.Tensor, + ctx: TleGemmReduceScatterContext) -> torch.Tensor: + """Tutorial-08 persistent GEMM + per-rank-ready reduce-scatter path.""" + world_size = ctx.rs_ctx.world_size + local_world_size = ctx.rs_ctx.local_world_size + rs_stream = ctx.rs_stream + original_M = input_tensor.shape[0] + original_M_per_rank = original_M // world_size + + input_tensor = _pad_to_block_m(input_tensor, world_size, ctx.BLOCK_M) + M, K = input_tensor.shape + N = weight.shape[0] + if N != ctx.rs_ctx.N or weight.shape[1] != K: + raise ValueError("invalid GEMM dimensions for the reduce-scatter context") + if M > ctx.rs_ctx.max_M: + raise ValueError("padded M exceeds context capacity") + + current_stream = torch.cuda.current_stream() + # Queue this dependency before GEMM. The RS stream then waits per target + # rank on scatter_signal rather than waiting for the complete GEMM launch. + rs_stream.wait_stream(current_stream) + + output = torch.empty((M // world_size, N), dtype=ctx.output_dtype, device="cuda") + + # Shape: [world_size] + # Each workspace[j] counts how many output tiles for rank j have completed. + # When workspace[j] reaches tiled_m_size * tiled_n - 1, + # all tiles for rank j are done, so trigger the signal notification. + workspace = torch.zeros((world_size, ), dtype=torch.int32, device=input_tensor.device) + scatter_signal = ctx.rs_ctx.scatter_signal_buf + gemm_out = ctx.get_gemm_out_buf(input_tensor) + + # Launch compute kernel. + gemm_rs_producer_persistent( + input_tensor, + weight, + gemm_out, + scatter_signal, + workspace, + world_size, + local_world_size, + ctx.rs_ctx.rank, + ctx.num_gemm_sms, + BLOCK_SIZE_M=ctx.BLOCK_M, + BLOCK_SIZE_N=ctx.BLOCK_N, + BLOCK_SIZE_K=ctx.BLOCK_K, + GROUP_SIZE_M=ctx.GROUP_M, + STAGES=ctx.STAGES, + ) + + # Launch communication kernel. + with torch.cuda.stream(rs_stream): + reduce_scatter_2d_op(gemm_out, ctx.rs_ctx, output=output, ready_flags=scatter_signal) + current_stream.wait_stream(rs_stream) + return output[:original_M_per_rank] + + +def gemm_rs_multi_node(input_tensor: torch.Tensor, weight: torch.Tensor, + ctx: TleGemmReduceScatterContext) -> torch.Tensor: + """Tutorial-08 public GEMM + Reduce-Scatter entry point.""" + + return gemm_rs_multi_node_persistent_op(input_tensor, weight, ctx) + + +# PyTorch baseline implementation. +def torch_gemm_rs(input_tensor: torch.Tensor, weight: torch.Tensor, TP_GROUP) -> torch.Tensor: + """PyTorch/NCCL baseline""" + M, _ = input_tensor.shape + N = weight.shape[0] + gemm_out = torch.matmul(input_tensor, weight.T) + output = torch.empty((M // TP_GROUP.size(), N), dtype=gemm_out.dtype, device=input_tensor.device) + dist.reduce_scatter_tensor(output, gemm_out, group=TP_GROUP) + return output + + +# Timing/benchmark code. +def _time_ms(fn, stream: torch.cuda.Stream, warmup: int = 20, iters: int = 200, clear_l2: bool = True) -> float: + """Measure median latency with rank alignment and optional L2 eviction. + + Cache eviction occurs before the start event, so it is not included in the + measured latency. Keep the same setting for TLE and torch baselines. + """ + driver = triton.runtime.driver.active + cache = driver.get_empty_cache_for_benchmark() if clear_l2 else None + + for _ in range(warmup): + fn() + torch.cuda.synchronize() + dist.barrier() + + samples = [] + for _ in range(iters): + start = torch.cuda.Event(enable_timing=True) + end = torch.cuda.Event(enable_timing=True) + + torch.cuda.synchronize() + dist.barrier() + driver.clear_cache(cache) + + start.record(stream) + fn() + end.record(stream) + torch.cuda.synchronize() + samples.append(start.elapsed_time(end)) + return float(statistics.median(samples)) + + +def main(): + + # Initialization. + tle.get_mem_pool() + rank = dist.get_rank() + TP_GROUP = dist.group.WORLD + world_size = dist.get_world_size() + local_rank = int(os.environ.get("LOCAL_RANK", rank)) + local_world_size = int(os.environ.get("LOCAL_WORLD_SIZE", world_size)) + torch.cuda.set_device(local_rank) + + if world_size < 2: + print("This example needs at least two GPUs", file=sys.stderr) + return + + if local_world_size != world_size: + raise NotImplementedError("TLE node-level GEMM reduce-scatter is not implemented") + + if torch.cuda.get_device_capability()[0] < 9: + print("Skip: persistent TMA GEMM requires sm90 or newer") + tle.cleanup_communicator() + return + + # Generate input. + M, N, K = 16384, 12288, 49152 + local_K = K // world_size + dtype = torch.bfloat16 + scale = rank + 1 + input_tensor = (torch.rand((M, local_K), dtype=dtype, device="cuda") * (0.02 * scale) - 0.01 * scale) + weight = (torch.rand((N, local_K), dtype=dtype, device="cuda") * (0.02 * scale) - 0.01 * scale) + + # Matches the validated TLE path: scatter and intra-node reduction share + # the caller-provided high-priority RS stream. + rs_stream = torch.cuda.Stream(priority=-1) + + num_scatter_sms = int(os.environ.get("TLE_SCATTER_SMS", "16")) + ctx = create_gemm_rs_context(M, N, rank, world_size, local_world_size, dtype, rs_stream, + num_scatter_sms=num_scatter_sms) + if rank == 0: + print(f"[Rank 0] TLE scatter CTA budget: {num_scatter_sms}; " + f"GEMM persistent CTA budget: {ctx.num_gemm_sms}") + + torch_output = torch_gemm_rs(input_tensor, weight, TP_GROUP) + + tle_output = gemm_rs_multi_node(input_tensor, weight, ctx) + + torch.cuda.synchronize() + + torch.testing.assert_close(torch_output, tle_output, atol=6e-2, rtol=6e-2) + print(f"[Rank {rank}] TLE overlapping GEMM+RS correctness passed") + + tle_ms = _time_ms(lambda: gemm_rs_multi_node(input_tensor, weight, ctx), torch.cuda.current_stream()) + + torch_ms = _time_ms(lambda: torch_gemm_rs(input_tensor, weight, TP_GROUP), torch.cuda.current_stream()) + + print(f"[Rank {rank}] TLE GEMM+RS median: {tle_ms:.3f} ms") + print(f"[Rank {rank}] torch GEMM+RS median: {torch_ms:.3f} ms") + ctx.finalize() + + +if __name__ == "__main__": + main() diff --git a/python/tutorials/tle/test_tle_intra_node_allgather.py b/python/tutorials/tle/test_tle_intra_node_allgather.py new file mode 100644 index 0000000000..091b03e676 --- /dev/null +++ b/python/tutorials/tle/test_tle_intra_node_allgather.py @@ -0,0 +1,229 @@ +""" +Intra-node AllGather with FlagTree TLE device remote pointers. +This tutorial implements a single-node all-gather operator using +FlagTree TLE. +Run with a FlagTree environment, for example: + + export FLAGCX_MEM_ENABLE=1 + export FLAGCX_USE_HETERO_COMM=1 + export FLAGCX_VMM_ENABLE=0 + export FLAGCX_P2P_DISABLE=1 + export CUDA_VISIBLE_DEVICES=0,1 + # Optional: set FLAGCX_IB_HCA to the HCA list for your machine. + torchrun \ + --nproc_per_node="${NPROC_PER_NODE}" \ + --nnodes=1 \ + --node_rank=0 \ + --master_addr="${MASTER_ADDR}" \ + --master_port="${MASTER_PORT}" \ + "${SCRIPT_DIR}/test_tle_intra_node_allgather.py" + +If you explicitly disabled distributed support with USE_FLAGCX=0, USE_DIST=0, +or USE_TLE_DIST=0, it might be necessary to reset these settings before running this tutorial. +""" + +import os + +import torch +import torch.distributed as dist +import triton +import triton.language as tl +import triton.experimental.tle.language as tle + + +@triton.jit +def _all_gather_push_2d_kernel( + local_ptr, + ag_ptr, + ag_dev_mem, + dev_comm_dptr, # DevComm handle, used to query the current rank within the kernel + mesh: tl.constexpr, + ELEM_PER_RANK: tl.constexpr, + BLOCK: tl.constexpr, +): + peer = tl.program_id(0) + block_id = tl.program_id(1) + local_rank = tle.shard_id(mesh, "device", comm_ptr=dev_comm_dptr) + dst_base = local_rank * ELEM_PER_RANK + offsets = block_id * BLOCK + tl.arange(0, BLOCK) + mask = offsets < ELEM_PER_RANK + vals = tl.load(local_ptr + offsets, mask=mask, other=0.0) + + if peer != local_rank: + dst_ptr = tle.remote( + ag_dev_mem, + shard_id=peer, + space="device", + dtype=ag_ptr.dtype.element_ty, + offset=dst_base, + ) + tl.store(dst_ptr + offsets, vals, mask=mask) + else: + tl.store(ag_ptr + dst_base + offsets, vals, mask=mask) + + +@triton.jit(do_not_specialize=["signal_target"]) +def _all_gather_signal_kernel( + signal_ptr, + signal_dev_mem, + dev_comm_dptr, + mesh: tl.constexpr, + signal_target, +): + peer = tl.program_id(0) + local_rank = tle.shard_id(mesh, "device", comm_ptr=dev_comm_dptr) + + if peer != local_rank: + remote_signal_ptr = tle.remote( + signal_dev_mem, + shard_id=peer, + space="device", + dtype=tl.int32, + offset=local_rank, + ) + # Publish completion with system-scope release semantics. The receiver + # uses an acquire atomic before consuming the corresponding shard. + tl.atomic_xchg( + remote_signal_ptr, + signal_target, + sem="release", + scope="sys", + ) + else: + tl.atomic_xchg( + signal_ptr + local_rank, + signal_target, + sem="release", + scope="sys", + ) + + +@triton.jit(do_not_specialize=["signal_target"]) +def _all_gather_wait_kernel( + signal_ptr, + local_rank, + signal_target, + WORLD_SIZE: tl.constexpr, +): + """Wait until every remote shard in this rank's output is ready.""" + peer = tl.program_id(0) + if peer < WORLD_SIZE and peer != local_rank: + # atomic_add(0) is an acquire load expressed with Triton's public + # atomic API. GE is required because a faster peer may already have + # published a later epoch. + observed = tl.atomic_add( + signal_ptr + peer, + 0, + sem="acquire", + scope="sys", + ) + while observed < signal_target: + observed = tl.atomic_add( + signal_ptr + peer, + 0, + sem="acquire", + scope="sys", + ) + + +def _rank_print(rank: int, *items): + dist.barrier() + for cur_rank in range(dist.get_world_size()): + if cur_rank == rank: + print(*items, flush=True) + dist.barrier() + + +def main(): + local_rank = int(os.environ.get("LOCAL_RANK", "0")) + torch.cuda.set_device(local_rank) + + mem_pool = tle.get_mem_pool() + if mem_pool is None: + raise RuntimeError("FlagCX memory pool is unavailable; check FlagCX build and environment variables.") + + rank = dist.get_rank() # Obtain rank and world_size + world_size = dist.get_world_size() + local_world_size = int(os.getenv("LOCAL_WORLD_SIZE", str(world_size))) + assert world_size == local_world_size, "This tutorial is designed for a single node" + + M = 8192 + N = 12288 + assert M % world_size == 0 + m_per_rank = M // world_size + dtype = torch.float16 + device = torch.device("cuda", local_rank) + + local_data = torch.randn((m_per_rank, N), dtype=dtype, device=device) + + with torch.cuda.use_mem_pool(mem_pool): + ag_buffer = torch.empty((M, N), dtype=dtype, device=device) + signal = torch.empty((world_size, ), dtype=torch.int32, device=device) + + dev_comm_dptr, ag_dev_mem = tle.create_comm_tensor(ag_buffer) + _, signal_dev_mem = tle.create_comm_tensor(signal) + # ag_dev_mem is the device-side DevMem handle address created by FlagCX/TLE, + # signal_dev_mem is used to remotely write to the peer's signal[local_rank]. + # dev_comm_dptr is used on the device side by tle.shard_id(..., comm_ptr=...) to query the current rank. + + golden = torch.empty((M, N), dtype=dtype, device=device) + dist.all_gather_into_tensor(golden, local_data) + + ag_buffer.fill_(-1) + signal.zero_() + + torch.cuda.synchronize() + dist.barrier() + + elem_per_rank = m_per_rank * N + block = 1024 + num_blocks = triton.cdiv(elem_per_rank, block) + + # 2D copy grid: split each peer transfer into independent chunks. + # The signal is written by a second kernel so it is ordered after all copy chunks in this stream. + copy_grid = (world_size, num_blocks) + signal_grid = (world_size, ) + mesh = tle.device_mesh(tle.MeshConfig(device=world_size)) + signal_target = 1 + + def launch_tle_all_gather(): + _all_gather_push_2d_kernel[copy_grid]( + local_data, + ag_buffer, + ag_dev_mem, + dev_comm_dptr, + mesh, + ELEM_PER_RANK=elem_per_rank, + BLOCK=block, + num_warps=4, + ) + _all_gather_signal_kernel[signal_grid]( + signal, + signal_dev_mem, + dev_comm_dptr, + mesh, + signal_target, + num_warps=4, + ) + _all_gather_wait_kernel[signal_grid]( + signal, + rank, + signal_target, + WORLD_SIZE=world_size, + num_warps=1, + ) + + launch_tle_all_gather() + torch.cuda.synchronize() + dist.barrier() + + _rank_print(rank, f"Rank {rank} FlagTree Result:", ag_buffer) + _rank_print(rank, f"Rank {rank} FlagTree Signal:", signal) + assert torch.allclose(golden, ag_buffer, atol=1e-5, rtol=1e-5) + _rank_print(rank, f"Rank {rank} Pass!") + + tle.cleanup_communicator() + + +if __name__ == "__main__": + main() diff --git a/python/tutorials/tle/test_tle_intra_node_allgather.sh b/python/tutorials/tle/test_tle_intra_node_allgather.sh new file mode 100755 index 0000000000..fd025aaafd --- /dev/null +++ b/python/tutorials/tle/test_tle_intra_node_allgather.sh @@ -0,0 +1,30 @@ +#!/usr/bin/env bash +set -euo pipefail + +SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" + + +export FLAGCX_IB_HCA=mlx5_0,mlx5_1,mlx5_2,mlx5_3,mlx5_6,mlx5_7,mlx5_8,mlx5_9 +export FLAGCX_MEM_ENABLE=1 +export FLAGCX_USE_HETERO_COMM=1 +export FLAGCX_VMM_ENABLE=0 +export FLAGCX_P2P_DISABLE=1 +export CUDA_VISIBLE_DEVICES=0,1,2,3 + + +if [[ "${CLEAR_TRITON_CACHE:-0}" == "1" ]]; then + rm -rf "${TRITON_CACHE_DIR:-${HOME}/.triton/cache}" +fi + +NPROC_PER_NODE=4 +MASTER_ADDR=localhost +MASTER_PORT=8333 + + +torchrun \ + --nproc_per_node="${NPROC_PER_NODE}" \ + --nnodes=1 \ + --node_rank=0 \ + --master_addr="${MASTER_ADDR}" \ + --master_port="${MASTER_PORT}" \ + "${SCRIPT_DIR}/test_tle_intra_node_allgather.py" diff --git a/python/tutorials/tle/test_tle_intra_node_reduce_scatter.py b/python/tutorials/tle/test_tle_intra_node_reduce_scatter.py new file mode 100644 index 0000000000..529d8e4891 --- /dev/null +++ b/python/tutorials/tle/test_tle_intra_node_reduce_scatter.py @@ -0,0 +1,213 @@ +""" +Intra-node Reduce-Scatter using FlagTree TLE (Triton Language Extension) +========================================================================= + +This tutorial implements a single-node reduce-scatter operator using +FlagTree TLE. + +""" + +import os +import sys + +import torch +import torch.distributed as dist + +import triton +import triton.language as tl +import triton.experimental.tle.language as tle + + +@triton.jit +def scatter_kernel( + input_ptr, + dev_mem_ptr, + M_per_rank, + N, + LOCAL_RANK: tl.constexpr, + WORLD_SIZE: tl.constexpr, + BLOCK_M: tl.constexpr, + BLOCK_N: tl.constexpr, +): + pid = tl.program_id(0) + num_pid = tl.num_programs(0) + + num_tiles_m = tl.cdiv(M_per_rank, BLOCK_M) + num_tiles_n = tl.cdiv(N, BLOCK_N) + tiles_per_peer = num_tiles_m * num_tiles_n + total_tiles = tiles_per_peer * WORLD_SIZE + + row_offs = tl.arange(0, BLOCK_M) + col_offs = tl.arange(0, BLOCK_N) + + for tile_id in range(pid, total_tiles, num_pid): + peer = tile_id // tiles_per_peer + local_tile = tile_id % tiles_per_peer + tile_m = local_tile // num_tiles_n + tile_n = local_tile % num_tiles_n + + in_row = peer * M_per_rank + tile_m * BLOCK_M + in_col = tile_n * BLOCK_N + in_ptrs = (input_ptr + (in_row + row_offs[:, None]) * N + (in_col + col_offs[None, :])) + data = tl.load(in_ptrs) + + out_row_in_peer = LOCAL_RANK * M_per_rank + tile_m * BLOCK_M + out_col = tile_n * BLOCK_N + out_offset_elems = out_row_in_peer * N + out_col + + remote_base = tle.remote( + dev_mem_ptr, + space="device", + dtype=input_ptr.dtype.element_ty, + shard_id=peer, + offset=out_offset_elems, + ) + remote_ptrs = (remote_base + row_offs[:, None] * N + col_offs[None, :]) + # Write the local data of the current rank to the scatter buffer of the remote peer. + tl.store(remote_ptrs, data) + + +@triton.jit +def ring_reduce_kernel( + local_scatter_ptr, + output_ptr, + M_per_rank, + N, + LOCAL_RANK: tl.constexpr, + WORLD_SIZE: tl.constexpr, + BLOCK_M: tl.constexpr, + BLOCK_N: tl.constexpr, +): + pid = tl.program_id(0) + num_pid = tl.num_programs(0) + + num_tiles_m = tl.cdiv(M_per_rank, BLOCK_M) + num_tiles_n = tl.cdiv(N, BLOCK_N) + total_tiles = num_tiles_m * num_tiles_n + + row_offs = tl.arange(0, BLOCK_M) + col_offs = tl.arange(0, BLOCK_N) + + # swizzled starting slot, same idea as the NVSHMEM tutorial + begin_idx = LOCAL_RANK + + for tile_id in range(pid, total_tiles, num_pid): + tile_m = tile_id // num_tiles_n + tile_n = tile_id % num_tiles_n + + row_in_shard = tile_m * BLOCK_M + col = tile_n * BLOCK_N + + src_rank = (begin_idx + 1) % WORLD_SIZE + row = src_rank * M_per_rank + row_in_shard + ptrs = (local_scatter_ptr + (row + row_offs[:, None]) * N + (col + col_offs[None, :])) + accum = tl.load(ptrs) + + for i in range(1, WORLD_SIZE): + src_rank = (i + begin_idx + 1) % WORLD_SIZE + row = src_rank * M_per_rank + row_in_shard + ptrs = (local_scatter_ptr + (row + row_offs[:, None]) * N + (col + col_offs[None, :])) + data = tl.load(ptrs) + accum += data + + # store to local output + out_row = row_in_shard + out_col = col + out_ptrs = (output_ptr + (out_row + row_offs[:, None]) * N + (out_col + col_offs[None, :])) + tl.store(out_ptrs, accum) + + +def torch_reduce_scatter(input_tensor, group): + M, N = input_tensor.shape + world_size = dist.get_world_size(group) + output = torch.empty((M // world_size, N), dtype=input_tensor.dtype, device=input_tensor.device) + dist.reduce_scatter_tensor(output, input_tensor, group=group) + return output + + +def main(): + + mem_pool = tle.get_mem_pool() # calls tle.init_communicator() + + rank = dist.get_rank() + world_size = dist.get_world_size() + local_rank = int(os.environ.get("LOCAL_RANK", rank)) + torch.cuda.set_device(local_rank) + + print(f"[Rank {rank}/{world_size}] Starting TLE reduce-scatter") + + if world_size < 2: + print("This example needs at least 2 GPUs", file=sys.stderr) + sys.exit(1) + + dtype = torch.bfloat16 + M, N = 8192, 16384 + M_per_rank = M // world_size + + input_tensor = torch.rand((M, N), dtype=dtype, device="cuda") + + with torch.cuda.use_mem_pool(mem_pool): + scatter_buf = torch.empty((M, N), dtype=dtype, device="cuda").clone() + + _, dev_mem_ptr = tle.create_comm_tensor(scatter_buf) + output = torch.empty((M_per_rank, N), dtype=dtype, device="cuda") + + stream = torch.cuda.current_stream() + + torch_output = torch_reduce_scatter(input_tensor, group=None) + + torch.cuda.synchronize() + + grid_scatter = lambda META: (min( + triton.cdiv(M_per_rank, META["BLOCK_M"]) * triton.cdiv(N, META["BLOCK_N"]) * world_size, + 128, + ), ) + with torch.cuda.stream(stream): + scatter_kernel[grid_scatter]( + input_tensor, + dev_mem_ptr, + M_per_rank, + N, + LOCAL_RANK=local_rank, + WORLD_SIZE=world_size, + BLOCK_M=256, + BLOCK_N=64, + num_warps=4, + ) + + torch.cuda.synchronize() + + dist.barrier() + + grid_reduce = lambda META: (triton.cdiv(M_per_rank, META["BLOCK_M"]) * triton.cdiv(N, META["BLOCK_N"]), ) + with torch.cuda.stream(stream): + ring_reduce_kernel[grid_reduce]( + scatter_buf, + output, + M_per_rank, + N, + LOCAL_RANK=local_rank, + WORLD_SIZE=world_size, + BLOCK_M=256, + BLOCK_N=64, + num_warps=4, + ) + + torch.cuda.synchronize() + + # Validate: compare TLE output against PyTorch/NCCL reference + + atol, rtol = 6e-2, 6e-2 + if torch.allclose(torch_output, output, atol=atol, rtol=rtol): + print(f"[Rank {rank}] PASSED") + else: + print(f"[Rank {rank}] FAILED") + print(f"[Rank {rank}] torch_output[:2,:4] = {torch_output[:2,:4]}") + print(f"[Rank {rank}] tle_output[:2,:4] = {output[:2,:4]}") + sys.exit(1) + + tle.cleanup_communicator() + + +if __name__ == "__main__": + main() diff --git a/python/tutorials/tle/test_tle_intra_node_reduce_scatter.sh b/python/tutorials/tle/test_tle_intra_node_reduce_scatter.sh new file mode 100755 index 0000000000..fe724a1b16 --- /dev/null +++ b/python/tutorials/tle/test_tle_intra_node_reduce_scatter.sh @@ -0,0 +1,31 @@ +#!/bin/bash + +rm -rf ~/.triton/cache + +# FlagCX environment variables (tune for your machine if needed) +export FLAGCX_IB_HCA=mlx5_0,mlx5_1,mlx5_2,mlx5_3,mlx5_6,mlx5_7,mlx5_8,mlx5_9 +export FLAGCX_USE_HETERO_COMM=1 +export FLAGCX_MEM_ENABLE=1 +export FLAGCX_VMM_ENABLE=0 +export FLAGCX_P2P_DISABLE=0 +export CUDA_VISIBLE_DEVICES=0,1,2,3 + +run_test() { + local script_dir + script_dir="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" + + torchrun \ + --nproc_per_node=4 \ + --nnodes=1 \ + --node_rank=0 \ + --master_addr=localhost \ + --master_port=8333 \ + "${script_dir}/test_tle_intra_node_reduce_scatter.py" +} + +run_test + +if [ $? -ne 0 ]; then + echo "ERROR: tle_intra_node_reduce_scatter failed" + exit 1 +fi diff --git a/python/tutorials/tle/test_tle_intra_node_reduce_scatter_tma.py b/python/tutorials/tle/test_tle_intra_node_reduce_scatter_tma.py new file mode 100644 index 0000000000..571ef141b1 --- /dev/null +++ b/python/tutorials/tle/test_tle_intra_node_reduce_scatter_tma.py @@ -0,0 +1,286 @@ +import os +import sys +from typing import Optional + +import torch +import torch.distributed as dist + +import triton +import triton.language as tl +import triton.runtime +import triton.experimental.tle.language as tle + + +# Kernel 1: scatter (optimized) +@triton.jit +def scatter_kernel_opt( + input_ptr, + local_scatter_ptr, + dev_mem_ptr, + M_per_rank, + N, + LOCAL_RANK: tl.constexpr, + WORLD_SIZE: tl.constexpr, + BLOCK_M: tl.constexpr, + BLOCK_N: tl.constexpr, +): + pid = tl.program_id(0) + num_pid = tl.num_programs(0) + + num_tiles_m = tl.cdiv(M_per_rank, BLOCK_M) + num_tiles_n = tl.cdiv(N, BLOCK_N) + tiles_per_peer = num_tiles_m * num_tiles_n + + row_offs = tl.arange(0, BLOCK_M) + col_offs = tl.arange(0, BLOCK_N) + + # Every peer reserves a [M_per_rank, N] slot for this rank. + slot_offset_elems = LOCAL_RANK * M_per_rank * N + + for step in range(WORLD_SIZE): + peer = (LOCAL_RANK + step + 1) % WORLD_SIZE + + # Resolve remote base pointer once per peer. + if peer == LOCAL_RANK: + remote_base = local_scatter_ptr + slot_offset_elems + else: + remote_base = tle.remote( + dev_mem_ptr, + space="device", + dtype=input_ptr.dtype.element_ty, + shard_id=peer, + offset=slot_offset_elems, + ) + + for local_tile in range(pid, tiles_per_peer, num_pid): + tile_m = local_tile // num_tiles_n + tile_n = local_tile % num_tiles_n + + in_row = peer * M_per_rank + tile_m * BLOCK_M + in_col = tile_n * BLOCK_N + in_ptrs = (input_ptr + (in_row + row_offs[:, None]) * N + (in_col + col_offs[None, :])) + # Mask out rows/cols that are beyond the actual tensor boundary. + in_row_mask = (in_row + row_offs[:, None]) < (peer + 1) * M_per_rank + in_col_mask = (in_col + col_offs[None, :]) < N + data = tl.load(in_ptrs, mask=in_row_mask & in_col_mask, other=0.0) + + out_row_in_peer = tile_m * BLOCK_M + out_col = tile_n * BLOCK_N + out_ptrs = (remote_base + (out_row_in_peer + row_offs[:, None]) * N + (out_col + col_offs[None, :])) + out_row_mask = (out_row_in_peer + row_offs[:, None]) < M_per_rank + tl.store(out_ptrs, data, mask=out_row_mask & in_col_mask) + + +# Kernel 2: ring reduce with TMA descriptors +@triton.jit +def ring_reduce_kernel_tma( + local_scatter_ptr, + output_ptr, + M_per_rank, + N, + LOCAL_RANK: tl.constexpr, + WORLD_SIZE: tl.constexpr, + BLOCK_M: tl.constexpr, + BLOCK_N: tl.constexpr, +): + pid = tl.program_id(0) + num_pid = tl.num_programs(0) + + num_tiles_m = tl.cdiv(M_per_rank, BLOCK_M) + num_tiles_n = tl.cdiv(N, BLOCK_N) + total_tiles = num_tiles_m * num_tiles_n + + c_desc = tl.make_tensor_descriptor( + local_scatter_ptr, + shape=[M_per_rank * WORLD_SIZE, N], + strides=[N, 1], + block_shape=[BLOCK_M, BLOCK_N], + ) + + output_desc = tl.make_tensor_descriptor( + output_ptr, + shape=[M_per_rank, N], + strides=[N, 1], + block_shape=[BLOCK_M, BLOCK_N], + ) + + begin_idx = LOCAL_RANK + + for tile_id in range(pid, total_tiles, num_pid): + tile_m = tile_id // num_tiles_n + tile_n = tile_id % num_tiles_n + + row_in_shard = tile_m * BLOCK_M + col = tile_n * BLOCK_N + + src_rank = (begin_idx + 1) % WORLD_SIZE + accum = c_desc.load([ + row_in_shard + src_rank * M_per_rank, + col, + ]) + + for i in range(1, WORLD_SIZE): + src_rank = (i + begin_idx + 1) % WORLD_SIZE + data = c_desc.load([ + row_in_shard + src_rank * M_per_rank, + col, + ]) + accum += data + + output_desc.store([row_in_shard, col], accum) + + +# Host-side reference using PyTorch/NCCL +def torch_reduce_scatter(input_tensor, group): + M, N = input_tensor.shape + world_size = dist.get_world_size(group) + output = torch.empty((M // world_size, N), dtype=input_tensor.dtype, device=input_tensor.device) + dist.reduce_scatter_tensor(output, input_tensor, group=group) + return output + + +# TLE reduce-scatter host wrapper (optimized scatter -> barrier -> TMA reduce) +def tle_reduce_scatter( + input_tensor, + scatter_buf, + dev_mem_ptr, + output, + M_per_rank, + N, + local_rank, + world_size, + stream, + num_sms: int = -1, +): + # Scatter launch config mirrors the reduce two-tier design. + if num_sms == -1: + grid_scatter = lambda META: (triton.cdiv(M_per_rank, META["BLOCK_M"]) * triton.cdiv(N, META["BLOCK_N"]), ) + scatter_num_warps = 4 + else: + grid_scatter = lambda META: (min( + triton.cdiv(M_per_rank, META["BLOCK_M"]) * triton.cdiv(N, META["BLOCK_N"]), + 128, + ), ) + scatter_num_warps = 8 + + with torch.cuda.stream(stream): + scatter_kernel_opt[grid_scatter]( + input_tensor, + scatter_buf, + dev_mem_ptr, + M_per_rank, + N, + LOCAL_RANK=local_rank, + WORLD_SIZE=world_size, + BLOCK_M=256, + BLOCK_N=128, + num_warps=scatter_num_warps, + ) + + torch.cuda.synchronize() + dist.barrier() + + def alloc_fn(size: int, alignment: int, stream: Optional[int]): + return torch.empty(size, device="cuda", dtype=torch.int8) + + triton.set_allocator(alloc_fn) + + # Reduce launch config aligned with 05-intra-node-reduce-scatter.py + if num_sms == -1: + grid_reduce = lambda META: (triton.cdiv(M_per_rank, META["BLOCK_M"]) * triton.cdiv(N, META["BLOCK_N"]), ) + with torch.cuda.stream(stream): + ring_reduce_kernel_tma[grid_reduce]( + scatter_buf, + output, + M_per_rank, + N, + LOCAL_RANK=local_rank, + WORLD_SIZE=world_size, + BLOCK_M=256, + BLOCK_N=64, + num_warps=4, + ) + else: + grid_reduce = lambda META: (min( + triton.cdiv(M_per_rank, META["BLOCK_M"]) * triton.cdiv(N, META["BLOCK_N"]), + num_sms, + ), ) + with torch.cuda.stream(stream): + ring_reduce_kernel_tma[grid_reduce]( + scatter_buf, + output, + M_per_rank, + N, + LOCAL_RANK=local_rank, + WORLD_SIZE=world_size, + BLOCK_M=256, + BLOCK_N=128, + num_warps=8, + ) + + +# Main +def main(): + mem_pool = tle.get_mem_pool() + + rank = dist.get_rank() + world_size = dist.get_world_size() + local_rank = int(os.environ.get("LOCAL_RANK", rank)) + torch.cuda.set_device(local_rank) + + print(f"[Rank {rank}/{world_size}] Starting TLE reduce-scatter (4096, 2048)") + + if world_size < 2: + print("This example needs at least 2 GPUs", file=sys.stderr) + sys.exit(1) + + dtype = torch.bfloat16 + M, N = 4096, 2048 + M_per_rank = M // world_size + + if M_per_rank < 256: + print(f"M // world_size = {M_per_rank} < 256, skipping", file=sys.stderr) + sys.exit(1) + + with torch.cuda.use_mem_pool(mem_pool): + scatter_buf = torch.empty((M * N, ), dtype=dtype, device="cuda") + _, scatter_dev_mem_ptr = tle.create_comm_tensor(scatter_buf) + + input_tensor = torch.rand((M, N), dtype=dtype, device="cuda") + scatter_buf = scatter_buf.view(M, N) + output = torch.empty((M_per_rank, N), dtype=dtype, device="cuda") + stream = torch.cuda.current_stream() + num_sms = torch.cuda.get_device_properties(local_rank).multi_processor_count + + torch_output = torch_reduce_scatter(input_tensor, group=None) + torch.cuda.synchronize() + + # Correctness check. + tle_reduce_scatter( + input_tensor, + scatter_buf, + scatter_dev_mem_ptr, + output, + M_per_rank, + N, + local_rank, + world_size, + stream, + num_sms=num_sms, + ) + torch.cuda.synchronize() + + atol, rtol = 6e-2, 6e-2 + if torch.allclose(torch_output, output, atol=atol, rtol=rtol): + print(f"[Rank {rank}] shape={(M, N)} PASSED") + else: + print(f"[Rank {rank}] shape={(M, N)} FAILED") + print(f"[Rank {rank}] torch_output[:2,:4] = {torch_output[:2,:4]}") + print(f"[Rank {rank}] tle_output[:2,:4] = {output[:2,:4]}") + sys.exit(1) + + tle.cleanup_communicator() + + +if __name__ == "__main__": + main() diff --git a/python/tutorials/tle/test_tle_intra_node_reduce_scatter_tma.sh b/python/tutorials/tle/test_tle_intra_node_reduce_scatter_tma.sh new file mode 100755 index 0000000000..a465095a1e --- /dev/null +++ b/python/tutorials/tle/test_tle_intra_node_reduce_scatter_tma.sh @@ -0,0 +1,31 @@ +#!/bin/bash + +rm -rf ~/.triton/cache + +# FlagCX environment variables (tune for your machine if needed) +export FLAGCX_IB_HCA=mlx5_0,mlx5_1,mlx5_2,mlx5_3,mlx5_6,mlx5_7,mlx5_8,mlx5_9 +export FLAGCX_USE_HETERO_COMM=1 +export FLAGCX_MEM_ENABLE=1 +export FLAGCX_VMM_ENABLE=0 +export FLAGCX_P2P_DISABLE=0 +export CUDA_VISIBLE_DEVICES=0,1,2,3 + +run_test() { + local script_dir + script_dir="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" + + torchrun \ + --nproc_per_node=4 \ + --nnodes=1 \ + --node_rank=0 \ + --master_addr=localhost \ + --master_port=8333 \ + "${script_dir}/test_tle_intra_node_reduce_scatter_tma.py" +} + +run_test + +if [ $? -ne 0 ]; then + echo "ERROR: test_tle_intra_node_reduce_scatter_tma failed" + exit 1 +fi diff --git a/python/tutorials/tle/test_tle_intra_node_reduce_scatter_tma_atomic_barrier.py b/python/tutorials/tle/test_tle_intra_node_reduce_scatter_tma_atomic_barrier.py new file mode 100644 index 0000000000..3e10e4a8f2 --- /dev/null +++ b/python/tutorials/tle/test_tle_intra_node_reduce_scatter_tma_atomic_barrier.py @@ -0,0 +1,313 @@ +import os +import sys +from typing import Optional + +import torch +import torch.distributed as dist + +import triton +import triton.language as tl +import triton.runtime +import triton.experimental.tle.language as tle + + +# Kernel 1: scatter (optimized) +@triton.jit +def scatter_kernel_opt( + input_ptr, + local_scatter_ptr, + dev_mem_ptr, + M_per_rank, + N, + LOCAL_RANK: tl.constexpr, + WORLD_SIZE: tl.constexpr, + BLOCK_M: tl.constexpr, + BLOCK_N: tl.constexpr, +): + pid = tl.program_id(0) + num_pid = tl.num_programs(0) + + num_tiles_m = tl.cdiv(M_per_rank, BLOCK_M) + num_tiles_n = tl.cdiv(N, BLOCK_N) + tiles_per_peer = num_tiles_m * num_tiles_n + + row_offs = tl.arange(0, BLOCK_M) + col_offs = tl.arange(0, BLOCK_N) + + slot_offset_elems = LOCAL_RANK * M_per_rank * N + + for step in range(WORLD_SIZE): + peer = (LOCAL_RANK + step + 1) % WORLD_SIZE + + if peer == LOCAL_RANK: + remote_base = local_scatter_ptr + slot_offset_elems + else: + remote_base = tle.remote( + dev_mem_ptr, + space="device", + dtype=input_ptr.dtype.element_ty, + shard_id=peer, + offset=slot_offset_elems, + ) + + for local_tile in range(pid, tiles_per_peer, num_pid): + tile_m = local_tile // num_tiles_n + tile_n = local_tile % num_tiles_n + + in_row = peer * M_per_rank + tile_m * BLOCK_M + in_col = tile_n * BLOCK_N + in_ptrs = (input_ptr + (in_row + row_offs[:, None]) * N + (in_col + col_offs[None, :])) + + in_row_mask = (in_row + row_offs[:, None]) < (peer + 1) * M_per_rank + in_col_mask = (in_col + col_offs[None, :]) < N + data = tl.load(in_ptrs, mask=in_row_mask & in_col_mask, other=0.0) + + out_row_in_peer = tile_m * BLOCK_M + out_col = tile_n * BLOCK_N + out_ptrs = (remote_base + (out_row_in_peer + row_offs[:, None]) * N + (out_col + col_offs[None, :])) + out_row_mask = (out_row_in_peer + row_offs[:, None]) < M_per_rank + tl.store(out_ptrs, data, mask=out_row_mask & in_col_mask) + + +# Kernel 2: ring reduce with TMA descriptors +@triton.jit +def ring_reduce_kernel_tma( + local_scatter_ptr, + output_ptr, + M_per_rank, + N, + LOCAL_RANK: tl.constexpr, + WORLD_SIZE: tl.constexpr, + BLOCK_M: tl.constexpr, + BLOCK_N: tl.constexpr, +): + pid = tl.program_id(0) + num_pid = tl.num_programs(0) + + num_tiles_m = tl.cdiv(M_per_rank, BLOCK_M) + num_tiles_n = tl.cdiv(N, BLOCK_N) + total_tiles = num_tiles_m * num_tiles_n + + c_desc = tl.make_tensor_descriptor( + local_scatter_ptr, + shape=[M_per_rank * WORLD_SIZE, N], + strides=[N, 1], + block_shape=[BLOCK_M, BLOCK_N], + ) + + output_desc = tl.make_tensor_descriptor( + output_ptr, + shape=[M_per_rank, N], + strides=[N, 1], + block_shape=[BLOCK_M, BLOCK_N], + ) + + begin_idx = LOCAL_RANK + + for tile_id in range(pid, total_tiles, num_pid): + tile_m = tile_id // num_tiles_n + tile_n = tile_id % num_tiles_n + + row_in_shard = tile_m * BLOCK_M + col = tile_n * BLOCK_N + + src_rank = (begin_idx + 1) % WORLD_SIZE + accum = c_desc.load([ + row_in_shard + src_rank * M_per_rank, + col, + ]) + + for i in range(1, WORLD_SIZE): + src_rank = (i + begin_idx + 1) % WORLD_SIZE + data = c_desc.load([ + row_in_shard + src_rank * M_per_rank, + col, + ]) + accum += data + + output_desc.store([row_in_shard, col], accum) + + +# Kernel 3: intra-node atomic barrier using tle.remote + atomic_cas +@triton.jit +def atomic_barrier_kernel( + local_flag_ptr, + dev_mem_ptr, + LOCAL_RANK: tl.constexpr, + WORLD_SIZE: tl.constexpr, +): + peer = tl.program_id(0) + + if peer == LOCAL_RANK: + remote_flag = local_flag_ptr + LOCAL_RANK + else: + remote_flag = tle.remote( + dev_mem_ptr, + space="device", + dtype=tl.int32, + shard_id=peer, + offset=LOCAL_RANK, + ) + + while tl.atomic_cas(remote_flag, 0, 1, sem="release", scope="sys") != 0: + pass + + local_slot = local_flag_ptr + peer + while tl.atomic_cas(local_slot, 1, 0, sem="acquire", scope="sys") != 1: + pass + + +# Host-side reference using PyTorch/NCCL +def torch_reduce_scatter(input_tensor, group): + M, N = input_tensor.shape + world_size = dist.get_world_size(group) + output = torch.empty((M // world_size, N), dtype=input_tensor.dtype, device=input_tensor.device) + dist.reduce_scatter_tensor(output, input_tensor, group=group) + return output + + +# TLE reduce-scatter host wrapper (optimized scatter -> atomic barrier -> TMA reduce) +def tle_reduce_scatter_atomic_barrier( + input_tensor, + scatter_buf, + scatter_dev_mem_ptr, + output, + flag_buf, + flag_dev_mem_ptr, + M_per_rank, + N, + local_rank, + world_size, + stream, + num_sms: int = -1, +): + if num_sms == -1: + grid_scatter = lambda META: (triton.cdiv(M_per_rank, META["BLOCK_M"]) * triton.cdiv(N, META["BLOCK_N"]), ) + scatter_num_warps = 4 + else: + grid_scatter = lambda META: (min( + triton.cdiv(M_per_rank, META["BLOCK_M"]) * triton.cdiv(N, META["BLOCK_N"]), + 128, + ), ) + scatter_num_warps = 8 + + with torch.cuda.stream(stream): + scatter_kernel_opt[grid_scatter]( + input_tensor, + scatter_buf, + scatter_dev_mem_ptr, + M_per_rank, + N, + LOCAL_RANK=local_rank, + WORLD_SIZE=world_size, + BLOCK_M=256, + BLOCK_N=128, + num_warps=scatter_num_warps, + ) + + def alloc_fn(size: int, alignment: int, stream: Optional[int]): + return torch.empty(size, device="cuda", dtype=torch.int8) + + triton.set_allocator(alloc_fn) + + if num_sms == -1: + grid_reduce = lambda META: (triton.cdiv(M_per_rank, META["BLOCK_M"]) * triton.cdiv(N, META["BLOCK_N"]), ) + reduce_block_n = 64 + reduce_num_warps = 4 + else: + grid_reduce = lambda META: (min( + triton.cdiv(M_per_rank, META["BLOCK_M"]) * triton.cdiv(N, META["BLOCK_N"]), + num_sms, + ), ) + reduce_block_n = 128 + reduce_num_warps = 8 + + with torch.cuda.stream(stream): + atomic_barrier_kernel[(world_size, )]( + flag_buf, + flag_dev_mem_ptr, + LOCAL_RANK=local_rank, + WORLD_SIZE=world_size, + ) + ring_reduce_kernel_tma[grid_reduce]( + scatter_buf, + output, + M_per_rank, + N, + LOCAL_RANK=local_rank, + WORLD_SIZE=world_size, + BLOCK_M=256, + BLOCK_N=reduce_block_n, + num_warps=reduce_num_warps, + ) + + +# Main +def main(): + mem_pool = tle.get_mem_pool() + + rank = dist.get_rank() + world_size = dist.get_world_size() + local_rank = int(os.environ.get("LOCAL_RANK", rank)) + torch.cuda.set_device(local_rank) + + print(f"[Rank {rank}/{world_size}] Starting TLE reduce-scatter atomic barrier (4096, 2048)") + + if world_size < 2: + print("This example needs at least 2 GPUs", file=sys.stderr) + sys.exit(1) + + dtype = torch.bfloat16 + M, N = 4096, 2048 + M_per_rank = M // world_size + + if M_per_rank < 256: + print(f"M // world_size = {M_per_rank} < 256, skipping", file=sys.stderr) + sys.exit(1) + + with torch.cuda.use_mem_pool(mem_pool): + flag_buf = torch.full((world_size, ), 0, dtype=torch.int32, device="cuda") + scatter_buf = torch.empty((M * N, ), dtype=dtype, device="cuda") + _, flag_dev_mem_ptr = tle.create_comm_tensor(flag_buf) + _, scatter_dev_mem_ptr = tle.create_comm_tensor(scatter_buf) + + input_tensor = torch.rand((M, N), dtype=dtype, device="cuda") + scatter_buf = scatter_buf.view(M, N) + output = torch.empty((M_per_rank, N), dtype=dtype, device="cuda") + stream = torch.cuda.current_stream() + num_sms = torch.cuda.get_device_properties(local_rank).multi_processor_count + + torch_output = torch_reduce_scatter(input_tensor, group=None) + torch.cuda.synchronize() + + # Correctness check. + tle_reduce_scatter_atomic_barrier( + input_tensor, + scatter_buf, + scatter_dev_mem_ptr, + output, + flag_buf, + flag_dev_mem_ptr, + M_per_rank, + N, + local_rank, + world_size, + stream, + num_sms=num_sms, + ) + torch.cuda.synchronize() + + atol, rtol = 6e-2, 6e-2 + if torch.allclose(torch_output, output, atol=atol, rtol=rtol): + print(f"[Rank {rank}] shape={(M, N)} PASSED") + else: + print(f"[Rank {rank}] shape={(M, N)} FAILED") + print(f"[Rank {rank}] torch_output[:2,:4] = {torch_output[:2,:4]}") + print(f"[Rank {rank}] tle_output[:2,:4] = {output[:2,:4]}") + sys.exit(1) + + tle.cleanup_communicator() + + +if __name__ == "__main__": + main() diff --git a/python/tutorials/tle/test_tle_intra_node_reduce_scatter_tma_atomic_barrier.sh b/python/tutorials/tle/test_tle_intra_node_reduce_scatter_tma_atomic_barrier.sh new file mode 100755 index 0000000000..b46046070b --- /dev/null +++ b/python/tutorials/tle/test_tle_intra_node_reduce_scatter_tma_atomic_barrier.sh @@ -0,0 +1,31 @@ +#!/bin/bash + +rm -rf ~/.triton/cache + +# FlagCX environment variables (tune for your machine if needed) +export FLAGCX_IB_HCA=mlx5_0,mlx5_1,mlx5_2,mlx5_3,mlx5_6,mlx5_7,mlx5_8,mlx5_9 +export FLAGCX_USE_HETERO_COMM=1 +export FLAGCX_MEM_ENABLE=1 +export FLAGCX_VMM_ENABLE=0 +export FLAGCX_P2P_DISABLE=0 +export CUDA_VISIBLE_DEVICES=0,1,2,3 + +run_test() { + local script_dir + script_dir="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" + + torchrun \ + --nproc_per_node=4 \ + --nnodes=1 \ + --node_rank=0 \ + --master_addr=localhost \ + --master_port=8333 \ + "${script_dir}/test_tle_intra_node_reduce_scatter_tma_atomic_barrier.py" +} + +run_test + +if [ $? -ne 0 ]; then + echo "ERROR: test_tle_intra_node_reduce_scatter_tma_atomic_barrier failed" + exit 1 +fi