From 6d30a290d01478e2339df843af81aad2fe3cf3d6 Mon Sep 17 00:00:00 2001 From: Rachmanino <18805904201@163.com> Date: Mon, 3 Aug 2026 23:13:15 +0800 Subject: [PATCH 01/30] [Feat] Detect GIN-capable NCCL and plumb TL_ENABLE_NCCL_GIN GIN (GPU-Initiated Networking) is the network mode of the NCCL Device API, first usable in NCCL 2.28.7. Most environments still ship 2.27.5, which has no `nccl_device/gin.h` and no `ncclDevCommCreate`, so the feature has to be compiled out rather than assumed present. The flag has to be defined twice by two independent mechanisms, because the host library and the JIT are separate compilations. `src/cuda/CMakeLists.txt` gates whether `codegen_cuda.cc` emits the `nccl_gin.h` include; `env.py` autodetects an include dir for the JIT, which `contrib/nvcc.py` and both `libgen.py` back ends turn into `-I` plus `-DTL_ENABLE_NCCL_GIN=1`. Keeping the two halves independent means a host build without GIN headers still works, and `TL_NCCL_PATH` can override the JIT half alone. `FindNCCLGin.cmake` insists on `nccl.h`, `nccl_device/gin.h` and the `ncclDevCommCreate` export together: a tree can advertise a new enough version while missing the device path, and only the export proves the library half is there. It also searches the versioned `libnccl.so.2` soname, since pip wheels ship no unversioned dev symlink. --- cmake/FindNCCLGin.cmake | 147 +++++++++++++++++++++++++++ src/cuda/CMakeLists.txt | 27 +++++ tilelang/contrib/nvcc.py | 15 ++- tilelang/env.py | 27 +++++ tilelang/jit/adapter/libgen.py | 6 ++ tilelang/jit/adapter/nvrtc/libgen.py | 9 ++ 6 files changed, 230 insertions(+), 1 deletion(-) create mode 100644 cmake/FindNCCLGin.cmake diff --git a/cmake/FindNCCLGin.cmake b/cmake/FindNCCLGin.cmake new file mode 100644 index 0000000000..efd9be6097 --- /dev/null +++ b/cmake/FindNCCLGin.cmake @@ -0,0 +1,147 @@ +# Detect an NCCL installation that provides the Device API (GIN). +# +# GIN (GPU-Initiated Networking) is what TileScale's inter-node put/signal path +# compiles against. It is *not* present in every NCCL: the device headers first +# ship in 2.28.7, and a pip `nvidia-nccl-cu12` wheel may be several minor +# versions behind that. Rather than gate on the version macro alone, this module +# requires all three things that the device path actually needs: +# +# 1. nccl.h -- host API, and the version macros +# 2. nccl_device/gin.h -- the ncclGin device class +# 3. ncclDevCommCreate -- host-side device-comm bootstrap, in libnccl +# +# A tree can satisfy (1) and report a new enough version while missing (2)/(3), +# which is why the symbol check is not skipped when the header is found. +# +# Sets: NCCLGin_FOUND, NCCL_INCLUDE_DIR, NCCL_LIBRARY, NCCL_VERSION_STRING +# +# Hint with -DNCCL_ROOT=, or let it fall back to the active Python +# environment's nvidia/nccl wheel and then the CUDA toolkit prefix. + +set(_nccl_gin_min_version "2.28.7") + +# Candidate prefixes, most specific first. +set(_nccl_hints "") +if(NCCL_ROOT) + list(APPEND _nccl_hints "${NCCL_ROOT}") +endif() +if(DEFINED ENV{NCCL_ROOT}) + list(APPEND _nccl_hints "$ENV{NCCL_ROOT}") +endif() + +# pip-installed NCCL lives under /nvidia/nccl. Ask the +# interpreter rather than globbing, so we track the env actually in use. +if(Python3_EXECUTABLE OR Python_EXECUTABLE) + if(Python3_EXECUTABLE) + set(_nccl_py "${Python3_EXECUTABLE}") + else() + set(_nccl_py "${Python_EXECUTABLE}") + endif() + execute_process( + COMMAND "${_nccl_py}" -c + "import os,sysconfig;p=os.path.join(sysconfig.get_paths()['purelib'],'nvidia','nccl');print(p if os.path.isdir(p) else '')" + OUTPUT_VARIABLE _nccl_pip_dir + OUTPUT_STRIP_TRAILING_WHITESPACE + ERROR_QUIET) + if(_nccl_pip_dir) + list(APPEND _nccl_hints "${_nccl_pip_dir}") + endif() +endif() + +if(CUDAToolkit_LIBRARY_ROOT) + list(APPEND _nccl_hints "${CUDAToolkit_LIBRARY_ROOT}") +endif() + +find_path(NCCL_INCLUDE_DIR nccl.h + HINTS ${_nccl_hints} + PATH_SUFFIXES include + DOC "Directory containing nccl.h") + +# pip wheels ship only the versioned soname -- there is no libnccl.so +# development symlink -- so the bare "nccl" name that find_library derives is not +# enough. List the versioned file explicitly, and keep the unversioned name first +# so a real system/dev install still wins. +find_library(NCCL_LIBRARY + NAMES nccl libnccl.so.2 libnccl.so.2.dylib + HINTS ${_nccl_hints} + PATH_SUFFIXES lib lib64 + DOC "NCCL shared library") + +set(NCCL_VERSION_STRING "") +set(_nccl_has_gin_header FALSE) +set(_nccl_has_devcomm FALSE) + +if(NCCL_INCLUDE_DIR) + # Version macros. NCCL_VERSION_CODE is not usable here because it is a macro + # expression, so read the three components directly. + foreach(_part MAJOR MINOR PATCH) + file(STRINGS "${NCCL_INCLUDE_DIR}/nccl.h" _line + REGEX "^#define NCCL_${_part} +[0-9]+") + if(_line) + string(REGEX MATCH "[0-9]+" _nccl_${_part} "${_line}") + else() + set(_nccl_${_part} 0) + endif() + endforeach() + set(NCCL_VERSION_STRING "${_nccl_MAJOR}.${_nccl_MINOR}.${_nccl_PATCH}") + + if(EXISTS "${NCCL_INCLUDE_DIR}/nccl_device/gin.h") + set(_nccl_has_gin_header TRUE) + endif() +endif() + +# ncclDevCommCreate is the host entry point the GIN path needs; a tree can carry +# the header and still be linked against a runtime that does not export it. +if(NCCL_LIBRARY AND _nccl_has_gin_header) + if(UNIX AND NOT APPLE) + find_program(_nccl_nm NAMES nm) + if(_nccl_nm) + execute_process( + COMMAND "${_nccl_nm}" -D --defined-only "${NCCL_LIBRARY}" + OUTPUT_VARIABLE _nccl_syms ERROR_QUIET) + if(_nccl_syms MATCHES "ncclDevCommCreate") + set(_nccl_has_devcomm TRUE) + endif() + else() + # No nm available: trust the header plus version gate instead of + # silently disabling GIN on a stripped-down build image. + set(_nccl_has_devcomm TRUE) + endif() + else() + set(_nccl_has_devcomm TRUE) + endif() +endif() + +set(NCCLGin_FOUND FALSE) +if(NCCL_INCLUDE_DIR AND NCCL_LIBRARY AND _nccl_has_gin_header + AND _nccl_has_devcomm + AND NOT NCCL_VERSION_STRING VERSION_LESS _nccl_gin_min_version) + set(NCCLGin_FOUND TRUE) +endif() + +if(NCCLGin_FOUND) + message(STATUS "NCCL GIN: enabled (NCCL ${NCCL_VERSION_STRING} at ${NCCL_INCLUDE_DIR})") +elseif(NCCL_INCLUDE_DIR) + # Found NCCL but cannot use the device path. Say which check failed -- + # "GIN disabled" with no reason is the hard case to debug on a cluster. + # Order matters: the symbol check is skipped when the library is missing, so + # test for the library before blaming its exports. + if(NOT NCCL_LIBRARY) + set(_why "found headers at ${NCCL_INCLUDE_DIR} but no libnccl alongside them") + elseif(NOT _nccl_has_gin_header) + set(_why "no nccl_device/gin.h (needs NCCL >= ${_nccl_gin_min_version})") + elseif(NOT _nccl_has_devcomm) + set(_why "libnccl does not export ncclDevCommCreate") + elseif(NCCL_VERSION_STRING VERSION_LESS _nccl_gin_min_version) + set(_why "NCCL ${NCCL_VERSION_STRING} < ${_nccl_gin_min_version}") + else() + set(_why "incomplete installation") + endif() + message(STATUS "NCCL GIN: disabled -- ${_why}. " + "Inter-node kernels will fall back to intra-node paths.") +else() + message(STATUS "NCCL GIN: disabled -- NCCL not found. " + "Set -DNCCL_ROOT= to enable inter-node support.") +endif() + +mark_as_advanced(NCCL_INCLUDE_DIR NCCL_LIBRARY) diff --git a/src/cuda/CMakeLists.txt b/src/cuda/CMakeLists.txt index afb498d0ad..37602e6094 100644 --- a/src/cuda/CMakeLists.txt +++ b/src/cuda/CMakeLists.txt @@ -158,6 +158,33 @@ list(APPEND TILE_LANG_SRCS ${TILE_LANG_CUDA_SRCS}) list(APPEND TILE_LANG_INCLUDES ${CUDAToolkit_INCLUDE_DIRS}) link_directories(${CUDAToolkit_LIBRARY_DIR} ${CUDAToolkit_LIBRARY_DIR}/stubs) +# Inter-node support (NCCL Device API / GIN). Optional: when absent, the +# nccl_gin.h include is not emitted by codegen and distributed kernels keep +# working on the intra-node NVLink paths. +option(TILELANG_USE_NCCL_GIN + "Enable inter-node put/signal via the NCCL Device API (GIN) when available" ON) + +set(TILELANG_NCCL_GIN_ENABLED OFF) +if(TILELANG_USE_NCCL_GIN) + # This file is include()d from the top-level CMakeLists.txt, so + # CMAKE_CURRENT_SOURCE_DIR is the repo root (same assumption as the + # file(GLOB src/cuda/*.cc) calls below). + include(${CMAKE_CURRENT_SOURCE_DIR}/cmake/FindNCCLGin.cmake) + if(NCCLGin_FOUND) + set(TILELANG_NCCL_GIN_ENABLED ON) + list(APPEND TILE_LANG_INCLUDES ${NCCL_INCLUDE_DIR}) + # codegen_cuda.cc reads this to decide whether to emit the GIN include. + # tilelang_objs does not exist yet at include() time, so set the definition + # directory-wide (same approach as CUDA_MAJOR_VERSION above). + add_compile_definitions(TL_ENABLE_NCCL_GIN=1) + # Recorded for tilelang/env.py, which must pass the same include dir and + # define to nvcc at JIT time -- the host build enabling GIN is not enough. + set(TILELANG_NCCL_INCLUDE_DIR "${NCCL_INCLUDE_DIR}") + set(TILELANG_NCCL_LIBRARY "${NCCL_LIBRARY}") + set(TILELANG_NCCL_VERSION "${NCCL_VERSION_STRING}") + endif() +endif() + # Register stubs for linking and install if(TILELANG_USE_CUDA_STUBS) list(APPEND TILELANG_ACTIVE_BACKEND_STUB_LINK cuda_stub) diff --git a/tilelang/contrib/nvcc.py b/tilelang/contrib/nvcc.py index 41ac058072..0a49c3e38f 100644 --- a/tilelang/contrib/nvcc.py +++ b/tilelang/contrib/nvcc.py @@ -8,7 +8,13 @@ import subprocess import warnings import contextlib -from tilelang.env import CUDA_HOME, CUTLASS_INCLUDE_DIR, TILELANG_TEMPLATE_PATH, env +from tilelang.env import ( + CUDA_HOME, + CUTLASS_INCLUDE_DIR, + NCCL_INCLUDE_DIR, + TILELANG_TEMPLATE_PATH, + env, +) import shutil import tempfile import tvm_ffi @@ -180,6 +186,13 @@ def default_compile_options(compile_flags: list[str] | None = None) -> list[str] options.append(f"-I{CUTLASS_INCLUDE_DIR}") except Exception: pass + try: + # NCCL device headers plus the define codegen keys the GIN include off. + if NCCL_INCLUDE_DIR: + options.append(f"-I{NCCL_INCLUDE_DIR}") + options.append("-DTL_ENABLE_NCCL_GIN=1") + except Exception: + pass try: if CUDA_HOME: options.append(f"-I{os.path.join(CUDA_HOME, 'include')}") diff --git a/tilelang/env.py b/tilelang/env.py index 637ae48884..e37a9dc5b8 100644 --- a/tilelang/env.py +++ b/tilelang/env.py @@ -305,6 +305,9 @@ class Environment: # External library include paths CUTLASS_INCLUDE_DIR = EnvVar("TL_CUTLASS_PATH", None) COMPOSABLE_KERNEL_INCLUDE_DIR = EnvVar("TL_COMPOSABLE_KERNEL_PATH", None) + # NCCL headers, needed at JIT time by the inter-node (GIN) device path. + # Autodetected below when unset; set to "" to force GIN off. + NCCL_INCLUDE_DIR = EnvVar("TL_NCCL_PATH", None) # TVM integration TVM_PYTHON_PATH = EnvVar("TVM_IMPORT_PYTHON_PATH", None) @@ -511,7 +514,31 @@ def prepend_pythonpath(path): else: logger.warning(TL_TEMPLATE_NOT_FOUND_MESSAGE) +# Initialize NCCL include path for the inter-node (GIN) device path. +# +# Only a GIN-capable tree counts: nccl_device/gin.h first ships in NCCL 2.28.7, +# and many environments pin an older wheel whose nccl.h alone would pass. When +# nothing suitable is found this stays None and the GIN header is never included, +# leaving intra-node kernels unaffected. +if os.environ.get("TL_NCCL_PATH", None) is None: + _nccl_inc_candidates = [] + try: + import sysconfig as _sysconfig + + _nccl_inc_candidates.append( + os.path.join(_sysconfig.get_paths()["purelib"], "nvidia", "nccl", "include") + ) + except Exception: + pass + if env.CUDA_HOME: + _nccl_inc_candidates.append(os.path.join(env.CUDA_HOME, "include")) + for _cand in _nccl_inc_candidates: + if os.path.exists(os.path.join(_cand, "nccl_device", "gin.h")): + os.environ["TL_NCCL_PATH"] = env.NCCL_INCLUDE_DIR = _cand + break + # Export static variables after initialization. +NCCL_INCLUDE_DIR = env.NCCL_INCLUDE_DIR CUTLASS_INCLUDE_DIR = env.CUTLASS_INCLUDE_DIR COMPOSABLE_KERNEL_INCLUDE_DIR = env.COMPOSABLE_KERNEL_INCLUDE_DIR TILELANG_TEMPLATE_PATH = env.TILELANG_TEMPLATE_PATH diff --git a/tilelang/jit/adapter/libgen.py b/tilelang/jit/adapter/libgen.py index 238775e35f..bcf8545b95 100644 --- a/tilelang/jit/adapter/libgen.py +++ b/tilelang/jit/adapter/libgen.py @@ -112,6 +112,12 @@ def compile_lib(self, timeout: float = None): command += [ "-I" + CUTLASS_INCLUDE_DIR, ] + # Inter-node (GIN) device path. Codegen only emits the nccl_gin.h + # include when TL_ENABLE_NCCL_GIN is defined, so pass both together. + from tilelang.env import NCCL_INCLUDE_DIR + + if NCCL_INCLUDE_DIR: + command += ["-I" + NCCL_INCLUDE_DIR, "-DTL_ENABLE_NCCL_GIN=1"] elif is_hip_target(target): from tilelang.env import COMPOSABLE_KERNEL_INCLUDE_DIR, TILELANG_HIP_SAVE_TEMP_FILES diff --git a/tilelang/jit/adapter/nvrtc/libgen.py b/tilelang/jit/adapter/nvrtc/libgen.py index 032105221c..cc26e0ea10 100644 --- a/tilelang/jit/adapter/nvrtc/libgen.py +++ b/tilelang/jit/adapter/nvrtc/libgen.py @@ -222,6 +222,15 @@ def compile_lib(self, timeout: float | None = None): if __CUDACC_VER_MAJOR__ < 13: options += [f"-I{arch_include}/cuda/std"] + # This path builds its own option list rather than going through + # default_compile_options, so the GIN include and define have to be + # added here too. Without them nccl_gin.h compiles to nothing and a + # kernel calling tl::gin:: fails with "must be a class or namespace". + from tilelang.env import NCCL_INCLUDE_DIR + + if NCCL_INCLUDE_DIR: + options += [f"-I{NCCL_INCLUDE_DIR}", "-DTL_ENABLE_NCCL_GIN=1"] + if self.compile_flags: options += [item for flag in self.compile_flags for item in flag.split() if item not in options] From 20ec569253df72cb0af67d765be3e38e96d25367 Mon Sep 17 00:00:00 2001 From: Rachmanino <18805904201@163.com> Date: Mon, 3 Aug 2026 23:13:15 +0800 Subject: [PATCH 02/30] [Refactor] Give the distributed metadata table a named layout The table is the contract between the host allocator, the device headers and the host-side TMA encoder in runtime.cc, and all three indexed it with raw integers. A one-sided edit does not fail: cached kernels keep reading the old slot while the allocator publishes the new one, which corrupts remote addressing and hangs with no error. `runtime.cc` was already wrong this way, reading `meta_data[2 + rank]` after the header grew past index 2. `meta_layout.h` now declares every offset once as `TL_META_*` and is free of CUDA constructs so plain host translation units can include it. It documents why the GIN context count is *not* in the table: the devcomm may grant fewer contexts than requested, so the count has to be read back on the device, and publishing the host's request would let a kernel index past the end. Multi-node makes the global/local distinction load-bearing, so the table now carries node_rank, num_nodes, local_rank and local_world_size, and the peer pointer array is indexed by *local* rank -- inter-node peers have no local virtual address and can never appear there. `get_remote_base_ptr` returns 0 for a non-local rank rather than reading a neighbour's slot, and `runtime.cc` reduces a global destination to its node-local rank before indexing. `get_rank`/`get_num_ranks` stay as aliases so single-node kernels are unaffected. --- src/cuda/runtime.cc | 35 +++++++---- .../cuda/distributed/distributed.h | 46 +++++++++++++- .../cuda/distributed/meta_layout.h | 60 +++++++++++++++++++ 3 files changed, 128 insertions(+), 13 deletions(-) create mode 100644 src/tl_templates/cuda/distributed/meta_layout.h diff --git a/src/cuda/runtime.cc b/src/cuda/runtime.cc index 738fee20ee..cd55f16212 100644 --- a/src/cuda/runtime.cc +++ b/src/cuda/runtime.cc @@ -9,6 +9,7 @@ #include #include "cuda/stubs/cuda.h" +#include "tl_templates/cuda/distributed/meta_layout.h" #include #include #include @@ -220,21 +221,35 @@ static bool IsPackedAlign16TensorMapType(CUtensorMapDataType type) { } static void *RemapSymmetricRemoteAddress(void *local_address, int64_t dst_pe) { - ICHECK_GE(remote_tensormap_meta_data.size(), 3U) + ICHECK_GT(remote_tensormap_meta_data.size(), + static_cast(TL_META_PEER_BASE)) << "Distributed meta_data is not initialized. Call " "kernel.initialize(allocator=...) before using remote TMA " "descriptors."; ICHECK_GE(dst_pe, 0); - uint64_t local_rank = remote_tensormap_meta_data[0]; - if (dst_pe >= static_cast(remote_tensormap_meta_data[1])) { - // Remote descriptor lowering may materialize a fixed peer descriptor set - // and select among them in device code. Descriptors for peers outside the - // current process group are never selected, but still need a valid dummy - // address during host-side CUtensorMap encoding. - dst_pe = static_cast(local_rank); + // Peer base pointers are indexed by local rank: only node-local peers are + // mappable, so a global rank must be reduced to its node-local rank. + uint64_t local_world_size = + remote_tensormap_meta_data[TL_META_LOCAL_WORLD_SIZE]; + ICHECK_GT(local_world_size, 0U) << "Distributed meta_data reports an empty " + "node-local world size."; + uint64_t local_rank = remote_tensormap_meta_data[TL_META_LOCAL_RANK]; + uint64_t node_rank = remote_tensormap_meta_data[TL_META_NODE_RANK]; + // Remote descriptor lowering may materialize a fixed peer descriptor set and + // select among them in device code. Descriptors for peers outside this node + // are never selected, but still need a valid dummy address during host-side + // CUtensorMap encoding. + uint64_t dst = static_cast(dst_pe); + uint64_t dst_local_rank = (dst / local_world_size == node_rank) + ? dst % local_world_size + : local_rank; + if (dst_local_rank >= local_world_size) { + dst_local_rank = local_rank; } - uint64_t local_base = remote_tensormap_meta_data[2 + local_rank]; - uint64_t remote_base = remote_tensormap_meta_data[2 + dst_pe]; + uint64_t local_base = + remote_tensormap_meta_data[TL_META_PEER_BASE + local_rank]; + uint64_t remote_base = + remote_tensormap_meta_data[TL_META_PEER_BASE + dst_local_rank]; uint64_t local_ptr = reinterpret_cast(local_address); ICHECK_GE(local_ptr, local_base) << "Remote TMA descriptor base pointer is outside the symmetric " diff --git a/src/tl_templates/cuda/distributed/distributed.h b/src/tl_templates/cuda/distributed/distributed.h index a1fd350ad7..69e2af7d97 100644 --- a/src/tl_templates/cuda/distributed/distributed.h +++ b/src/tl_templates/cuda/distributed/distributed.h @@ -1,6 +1,7 @@ #pragma once #include "../common.h" +#include "meta_layout.h" #define TL_ENABLE_DISTRIBUTED_METADATA 1 extern "C" { @@ -8,12 +9,51 @@ __constant__ uint64_t meta_data[1024]; } namespace tl { -TL_DEVICE uint64_t get_rank() { return meta_data[0]; } +// See meta_layout.h for the table layout shared with the host runtime. -TL_DEVICE uint64_t get_num_ranks() { return meta_data[1]; } +TL_DEVICE uint64_t get_global_rank() { return meta_data[TL_META_GLOBAL_RANK]; } +TL_DEVICE uint64_t get_global_world_size() { + return meta_data[TL_META_GLOBAL_WORLD_SIZE]; +} + +TL_DEVICE uint64_t get_node_rank() { return meta_data[TL_META_NODE_RANK]; } + +TL_DEVICE uint64_t get_num_nodes() { return meta_data[TL_META_NUM_NODES]; } + +TL_DEVICE uint64_t get_local_rank() { return meta_data[TL_META_LOCAL_RANK]; } + +TL_DEVICE uint64_t get_local_world_size() { + return meta_data[TL_META_LOCAL_WORLD_SIZE]; +} + +// Backward compatibility aliases +TL_DEVICE uint64_t get_rank() { return get_global_rank(); } + +TL_DEVICE uint64_t get_num_ranks() { return get_global_world_size(); } + +// Check if target rank is on same node +TL_DEVICE bool is_local_peer(uint64_t target_global_rank) { + uint64_t target_node = target_global_rank / get_local_world_size(); + return target_node == get_node_rank(); +} + +// Get intra-node base pointer (local peers only) +TL_DEVICE uint64_t get_local_peer_base_ptr(uint64_t local_rank) { + return meta_data[TL_META_PEER_BASE + local_rank]; +} + +// Get remote base pointer by global rank. Single-node ranks are their own local +// ranks, so this indexes the peer array directly. Returns 0 for an inter-node +// rank, whose memory is not mappable here: route those through GIN instead. TL_DEVICE uint64_t get_remote_base_ptr(uint64_t rank) { - return meta_data[2 + rank]; + if (get_num_nodes() == 1) { + return get_local_peer_base_ptr(rank); + } + if (!is_local_peer(rank)) { + return 0; + } + return get_local_peer_base_ptr(rank % get_local_world_size()); } template TL_DEVICE uint64_t get_uintptr_t(dtype_t *ptr) { diff --git a/src/tl_templates/cuda/distributed/meta_layout.h b/src/tl_templates/cuda/distributed/meta_layout.h new file mode 100644 index 0000000000..d8f5ee2285 --- /dev/null +++ b/src/tl_templates/cuda/distributed/meta_layout.h @@ -0,0 +1,60 @@ +#pragma once + +// Layout of the distributed metadata table. +// +// The table is produced on the host by BaseAllocator._init_table +// (tilelang/distributed/allocator.py), copied into the device __constant__ +// `meta_data` symbol by __tilescale_init_table +// (src/runtime/tilescale_cuda_module.cc), and additionally cached on the host +// for remote TMA descriptor encoding by SetRemoteTensorMapMetaData +// (src/cuda/runtime.cc). +// +// All three readers must agree, so the offsets live here and nowhere else. +// This header is intentionally free of CUDA constructs so it can be included +// from plain host translation units. +// +// [0] global_rank +// [1] global_world_size +// [2] node_rank +// [3] num_nodes +// [4] local_rank +// [5] local_world_size +// [6] global NCCL device-comm handle (0 when unused) +// [7] inter-node NCCL device-comm handle (0 when unused) +// [8] ncclWindow_t for the whole allocator arena (0 when unregistered) +// [9] local arena base address, for pointer -> window offset conversion +// [10 .. 10 + local_world_size) intra-node peer base pointers, by local rank +// +// The GIN context count is deliberately absent. The devcomm may grant fewer +// contexts than the allocator requested, and only the device can read the +// granted number back (tl::gin::context_span in nccl_gin.h). Publishing the +// host's request here would let a kernel index past the end of the context +// array, which hangs rather than failing. +// +// Peer base pointers are node-local by construction: inter-node peers are not +// mappable into this address space, so their bases are never present here. +// +// Inter-node peers are reached instead through the arena window. GIN addresses +// memory as an (ncclWindow_t, byte offset) pair rather than a raw pointer, +// because a remote rank's allocation has no local virtual address. The window +// handle returned by ncclCommWindowRegister is not per-peer: one local handle +// plus a peer index names any rank's bytes, so a single arena registration +// covers the whole communicator. Since the arena is symmetric, a local pointer +// converts to the offset valid on every rank by subtracting ARENA_BASE -- the +// same subtraction the intra-node path already performs, just paired with a +// window handle instead of a peer base pointer. + +#define TL_META_GLOBAL_RANK 0 +#define TL_META_GLOBAL_WORLD_SIZE 1 +#define TL_META_NODE_RANK 2 +#define TL_META_NUM_NODES 3 +#define TL_META_LOCAL_RANK 4 +#define TL_META_LOCAL_WORLD_SIZE 5 +#define TL_META_GLOBAL_DEV_COMM 6 +#define TL_META_INTERNODE_DEV_COMM 7 +#define TL_META_ARENA_WINDOW 8 +#define TL_META_ARENA_BASE 9 +#define TL_META_PEER_BASE 10 + +// Number of scalar header entries preceding the peer base pointer array. +#define TL_META_HEADER_SIZE TL_META_PEER_BASE From e182d91bb2974eb23b20d0e41ccd4882e636c505 Mon Sep 17 00:00:00 2001 From: Rachmanino <18805904201@163.com> Date: Mon, 3 Aug 2026 23:13:15 +0800 Subject: [PATCH 03/30] [Feat] Add GIN put/signal/wait device templates and the T.nccl_gin DSL GIN one-sided ops are callable from inside a kernel, unlike host NCCL collectives, which is what makes them the right fit for TileScale's kernel-side data plane. This exposes `T.nccl_gin.put` / `put_signal` / `signal` / `wait_signal` / `flush`, lowering through `src/op/nccl_gin.cc` to `tl::gin::` helpers in the device header. Two properties of GIN signals drive the design and are easy to get wrong: Signal state is per context. A put issued on sender context `i` increments the receiver's signal through context `i`, so a CTA spread over `C` contexts sees only `1/C` of the rank's arrivals. `wait_signal` therefore divides the caller's grid-wide target by `context_span()` *on the device*, because only the device knows how many contexts the devcomm actually granted -- `ncclDevCommRequirements.ginContextCount` is documented as a hint, and asking for 8 while 4 are granted means indexing past the end of the context array, which hangs. Requesting the count on the host and trusting it produced a reading of 332 GB/s, above the PCIe ceiling for one GPU's egress, because the wait was 2x too weak and the kernel returned early. Signals are cumulative and a wait does not consume them, so the expected count must be a runtime argument rather than a compile-time constant. A constant target is satisfied instantly on every launch after the first: correct once, then a silent no-op that benchmarks nothing. The codegen change adds `tl::gin::` to the prefixes that mark a kernel as distributed, so a kernel using only GIN still gets the distributed includes, and emits the `nccl_gin.h` include only when the host build found GIN. --- src/cuda/codegen/codegen_cuda.cc | 9 + src/op/nccl_gin.cc | 259 +++++++++++++++ src/op/nccl_gin.h | 187 +++++++++++ src/tl_templates/cuda/distributed/nccl_gin.h | 325 +++++++++++++++++++ tilelang/language/distributed/__init__.py | 7 + tilelang/language/distributed/nccl_gin.py | 164 ++++++++++ 6 files changed, 951 insertions(+) create mode 100644 src/op/nccl_gin.cc create mode 100644 src/op/nccl_gin.h create mode 100644 src/tl_templates/cuda/distributed/nccl_gin.h create mode 100644 tilelang/language/distributed/nccl_gin.py diff --git a/src/cuda/codegen/codegen_cuda.cc b/src/cuda/codegen/codegen_cuda.cc index 5e7ff9ac17..ae481d29e1 100644 --- a/src/cuda/codegen/codegen_cuda.cc +++ b/src/cuda/codegen/codegen_cuda.cc @@ -683,6 +683,12 @@ std::string CodeGenTileLangCUDA::Finish() { decl_stream << "#include \n"; decl_stream << "#include \n"; decl_stream << "#include \n"; +#ifdef TL_ENABLE_NCCL_GIN + // Inter-node put/signal via the NCCL Device API. The header is a no-op + // unless TL_ENABLE_NCCL_GIN is also defined for the JIT compile, which + // tilelang/env.py forwards when the build found a GIN-capable NCCL. + decl_stream << "#include \n"; +#endif } if (need_multimem_h_) { decl_stream << "#include \n"; @@ -2001,6 +2007,9 @@ void CodeGenTileLangCUDA::PrintCallExtern(Type ret_type, String global_symbol, "tl::cp_warp", "tl::cp_block", "tl::st<", "tl::ld<", "tl::remote_load", "tl::remote_store", + // GIN helpers live under the distributed includes, so an inter-node + // kernel that only calls tl::gin:: must still pull them in. + "tl::gin::", }; for (const char *prefix : kDistributedPrefixes) { size_t prefix_len = std::strlen(prefix); diff --git a/src/op/nccl_gin.cc b/src/op/nccl_gin.cc new file mode 100644 index 0000000000..51a398e145 --- /dev/null +++ b/src/op/nccl_gin.cc @@ -0,0 +1,259 @@ +/*! + * \file tl/op/nccl_gin.cc + * \brief Lowering for the NCCL GIN inter-node operators. + */ + +#include "nccl_gin.h" + +#include +#include +#include + +#include + +#include "builtin.h" +#include "distributed.h" +#include "distributed_utils.h" +#include "operator.h" + +namespace tvm { +namespace tl { + +using namespace tirx; + +// Map the DSL scope name onto an NCCL coop *type*. These are tag types in the +// device API, not enum values, and they appear here as a template argument, so +// this must name the type -- `ncclCoopCta()` would be a value and fail to +// substitute. The device wrapper default-constructs it internally. +static std::string CoopType(const std::string &scope) { + if (scope == "thread") { + return "ncclCoopThread"; + } + if (scope == "warp") { + return "ncclCoopWarp"; + } + if (scope == "block") { + return "ncclCoopCta"; + } + LOG(FATAL) << "invalid GIN cooperation scope: " << scope; + return ""; +} + +// Unlike the intra-node copies, the transfer size is a runtime argument to +// ncclGin::put rather than a template parameter, so a dynamic size is fine here +// and no constant-folding check is needed. + +// `size` counts elements at the DSL surface, matching T.put_block, whose +// cp_block takes N elements of a typed pointer. ncclGin::put takes bytes, so +// the conversion happens here rather than being pushed onto the user -- a GIN op +// whose size meant something different from the intra-node op it sits beside +// would be silently wrong by a factor of the element width. +static PrimExpr CopySizeInBytes(const PrimExpr ©_size, const Buffer &buffer) { + const int bits = buffer->dtype.bits() * buffer->dtype.lanes(); + ICHECK(bits % 8 == 0) << "GIN put requires a byte-addressable element type, got " + << buffer->dtype; + return cast(DataType::UInt(64), copy_size) * + make_const(DataType::UInt(64), bits / 8); +} + +GinPutOp::GinPutOp(Array args, Map annotations) { + ObjectPtr node = tvm::ffi::make_object(); + node->src_addr = args[0]; + node->dst_addr = args[1]; + ICHECK(node->src_addr.as() && + node->src_addr.as()->op.same_as(builtin::address_of())) + << "GIN put src must be address_of(...)"; + ICHECK(node->dst_addr.as() && + node->dst_addr.as()->op.same_as(builtin::address_of())) + << "GIN put dst must be address_of(...)"; + + const auto *src_load = + node->src_addr.as()->args[0].as(); + const auto *dst_load = + node->dst_addr.as()->args[0].as(); + ICHECK(src_load && dst_load) << "address_of must wrap BufferLoad nodes"; + + node->src_buffer = src_load->buffer; + node->dst_buffer = dst_load->buffer; + node->src_indices = src_load->indices; + node->dst_indices = dst_load->indices; + + // `size` is an element count, so the two sides must agree on what an element + // is; otherwise the byte count computed from the source would under- or + // over-write the destination. + ICHECK_EQ(node->src_buffer->dtype, node->dst_buffer->dtype) + << "GIN put requires matching src/dst dtypes, got " << node->src_buffer->dtype + << " and " << node->dst_buffer->dtype; + + node->copy_size = args[2]; + node->peer = args[3]; + node->signal_id = args[4].as().value()->value; + node->with_signal = bool(args[5].as().value()->value); + node->scope = args[6].as().value()->value; + data_ = std::move(node); +} + +Stmt GinPutOpNode::Lower(const LowerArgs &T, arith::Analyzer *analyzer) const { + (void)analyzer; + Array new_args; + std::stringstream ss; + + // Both offsets are computed device-side by tl::gin::arena_offset, which + // subtracts the arena base published in the metadata table. Doing it here + // instead would require the arena base as a compile-time value, which it is + // not -- it differs per rank and is only known once the allocator has run. + ss << (with_signal ? "tl::gin::put_signal_addr<" : "tl::gin::put_addr<") + << CoopType(scope) << ">"; + new_args.push_back(StringImm(ss.str())); + + // Peer is a global rank: GIN puts go through the communicator-wide team, whose + // rank space is global, unlike the node-local peer index the IPC path uses. + new_args.push_back(peer); + new_args.push_back(MakeRemappedAddress(T, dst_buffer, dst_indices)); + new_args.push_back(MakeRemappedAddress(T, src_buffer, src_indices)); + new_args.push_back(CopySizeInBytes(copy_size, src_buffer)); + if (with_signal) { + new_args.push_back(IntImm(DataType::Int(32), signal_id)); + } + + return Evaluate( + Call(DataType::Handle(), builtin::call_extern(), new_args)); +} + +LayoutMap GinPutOpNode::InferLayout(const LayoutInferArgs &T, + InferLevel level) const { + (void)T; + (void)level; + return {}; +} + +TileOperator GinPutOpNode::Clone() const { + auto node = tvm::ffi::make_object(*this); + return GinPutOp(node); +} + +GinSignalOp::GinSignalOp(Array args, + Map annotations) { + ObjectPtr node = tvm::ffi::make_object(); + node->peer = args[0]; + node->signal_id = args[1].as().value()->value; + node->scope = args[2].as().value()->value; + data_ = std::move(node); +} + +Stmt GinSignalOpNode::Lower(const LowerArgs &T, + arith::Analyzer *analyzer) const { + (void)T; + (void)analyzer; + std::stringstream ss; + ss << "tl::gin::signal_peer<" << CoopType(scope) << ">"; + Array new_args{StringImm(ss.str()), peer, + IntImm(DataType::Int(32), signal_id)}; + return Evaluate(Call(DataType::Handle(), builtin::call_extern(), new_args)); +} + +LayoutMap GinSignalOpNode::InferLayout(const LayoutInferArgs &T, + InferLevel level) const { + (void)T; + (void)level; + return {}; +} + +TileOperator GinSignalOpNode::Clone() const { + auto node = tvm::ffi::make_object(*this); + return GinSignalOp(node); +} + +GinWaitSignalOp::GinWaitSignalOp(Array args, + Map annotations) { + ObjectPtr node = + tvm::ffi::make_object(); + node->least = args[0]; + node->signal_id = args[1].as().value()->value; + node->scope = args[2].as().value()->value; + data_ = std::move(node); +} + +Stmt GinWaitSignalOpNode::Lower(const LowerArgs &T, + arith::Analyzer *analyzer) const { + (void)T; + (void)analyzer; + std::stringstream ss; + ss << "tl::gin::wait_signal<" << CoopType(scope) << ">"; + Array new_args{StringImm(ss.str()), + IntImm(DataType::Int(32), signal_id), + cast(DataType::UInt(64), least)}; + return Evaluate(Call(DataType::Handle(), builtin::call_extern(), new_args)); +} + +LayoutMap GinWaitSignalOpNode::InferLayout(const LayoutInferArgs &T, + InferLevel level) const { + (void)T; + (void)level; + return {}; +} + +TileOperator GinWaitSignalOpNode::Clone() const { + auto node = tvm::ffi::make_object(*this); + return GinWaitSignalOp(node); +} + +GinFlushOp::GinFlushOp(Array args, + Map annotations) { + ObjectPtr node = tvm::ffi::make_object(); + node->scope = args[0].as().value()->value; + data_ = std::move(node); +} + +Stmt GinFlushOpNode::Lower(const LowerArgs &T, arith::Analyzer *analyzer) const { + (void)T; + (void)analyzer; + std::stringstream ss; + ss << "tl::gin::flush<" << CoopType(scope) << ">"; + Array new_args{StringImm(ss.str())}; + return Evaluate(Call(DataType::Handle(), builtin::call_extern(), new_args)); +} + +LayoutMap GinFlushOpNode::InferLayout(const LayoutInferArgs &T, + InferLevel level) const { + (void)T; + (void)level; + return {}; +} + +TileOperator GinFlushOpNode::Clone() const { + auto node = tvm::ffi::make_object(*this); + return GinFlushOp(node); +} + +// kUpdateState, matching the intra-node put/get: these mutate memory (locally or +// remotely) and must not be reordered or elided. +TIR_REGISTER_TL_TILE_OP(GinPutOp, gin_put) + .set_num_inputs(7) + .set_attr("TCallEffectKind", + Integer(CallEffectKind::kUpdateState)); + +TIR_REGISTER_TL_TILE_OP(GinSignalOp, gin_signal) + .set_num_inputs(3) + .set_attr("TCallEffectKind", + Integer(CallEffectKind::kUpdateState)); + +TIR_REGISTER_TL_TILE_OP(GinWaitSignalOp, gin_wait_signal) + .set_num_inputs(3) + .set_attr("TCallEffectKind", + Integer(CallEffectKind::kUpdateState)); + +TIR_REGISTER_TL_TILE_OP(GinFlushOp, gin_flush) + .set_num_inputs(1) + .set_attr("TCallEffectKind", + Integer(CallEffectKind::kUpdateState)); + +TVM_FFI_STATIC_INIT_BLOCK() { + GinPutOpNode::RegisterReflection(); + GinSignalOpNode::RegisterReflection(); + GinWaitSignalOpNode::RegisterReflection(); + GinFlushOpNode::RegisterReflection(); +} + +} // namespace tl +} // namespace tvm diff --git a/src/op/nccl_gin.h b/src/op/nccl_gin.h new file mode 100644 index 0000000000..0b61847c05 --- /dev/null +++ b/src/op/nccl_gin.h @@ -0,0 +1,187 @@ +/*! + * \file tl/op/nccl_gin.h + * \brief Inter-node put/signal operators backed by the NCCL GIN device API. + * + * These are kept apart from the put/get in remote_copy.h because they use a + * different addressing model rather than a different transport for the same one. + * remote_copy computes `peer_base + (local_ptr - local_base)`, which needs the + * peer's memory mapped locally; GIN names memory as an (ncclWindow_t, offset) + * pair, which is what allows it to reach a rank on another node whose allocation + * has no local virtual address. Sharing an op between the two would mean a + * single node carrying both meanings of "peer" and both address computations. + */ + +#ifndef TVM_TL_OP_NCCL_GIN_H_ +#define TVM_TL_OP_NCCL_GIN_H_ + +#include +#include + +#include "../layout/layout.h" +#include "operator.h" + +namespace tvm { +namespace tl { + +using namespace tirx; + +/*! + * \brief One-sided inter-node put, optionally incrementing a remote signal. + */ +class GinPutOpNode : public TileOperatorNode { +public: + PrimExpr src_addr; ///< address_of the local source buffer element + PrimExpr dst_addr; ///< address_of the destination buffer element + PrimExpr copy_size; ///< Bytes to transfer + PrimExpr peer; ///< Destination *global* rank within the world team + int signal_id; ///< Signal to increment on arrival, when with_signal + bool with_signal; ///< Whether completion increments a remote signal + std::string scope; ///< Cooperation scope: {thread, warp, block} + Buffer src_buffer; ///< Source buffer, for arena offset computation + Buffer dst_buffer; ///< Destination buffer + Array src_indices; ///< Source indices + Array dst_indices; ///< Destination indices + + TVM_FFI_DECLARE_OBJECT_INFO_FINAL("tl.GinPutOp", GinPutOpNode, + TileOperatorNode); + + static void RegisterReflection() { + namespace refl = tvm::ffi::reflection; + refl::ObjectDef() + .def_ro("src_addr", &GinPutOpNode::src_addr) + .def_ro("dst_addr", &GinPutOpNode::dst_addr) + .def_ro("copy_size", &GinPutOpNode::copy_size) + .def_ro("peer", &GinPutOpNode::peer) + .def_ro("signal_id", &GinPutOpNode::signal_id) + .def_ro("with_signal", &GinPutOpNode::with_signal) + .def_ro("scope", &GinPutOpNode::scope) + .def_ro("src_buffer", &GinPutOpNode::src_buffer) + .def_ro("dst_buffer", &GinPutOpNode::dst_buffer) + .def_ro("src_indices", &GinPutOpNode::src_indices) + .def_ro("dst_indices", &GinPutOpNode::dst_indices); + } + + Stmt Lower(const LowerArgs &T, arith::Analyzer *analyzer) const override; + LayoutMap InferLayout(const LayoutInferArgs &T, + InferLevel level) const override; + static const Op &Get(); + TileOperator Clone() const override; +}; + +class GinPutOp : public TileOperator { +public: + TVM_FFI_DEFINE_OBJECT_REF_METHODS_NULLABLE(GinPutOp, TileOperator, + GinPutOpNode); + TVM_DLL GinPutOp(Array args, Map annotations = + Map()); + static const Op &Get(); +}; + +/*! + * \brief Increment a signal on a peer without moving payload. + */ +class GinSignalOpNode : public TileOperatorNode { +public: + PrimExpr peer; ///< Destination global rank + int signal_id; ///< Signal to increment + std::string scope; ///< Cooperation scope + + TVM_FFI_DECLARE_OBJECT_INFO_FINAL("tl.GinSignalOp", GinSignalOpNode, + TileOperatorNode); + + static void RegisterReflection() { + namespace refl = tvm::ffi::reflection; + refl::ObjectDef() + .def_ro("peer", &GinSignalOpNode::peer) + .def_ro("signal_id", &GinSignalOpNode::signal_id) + .def_ro("scope", &GinSignalOpNode::scope); + } + + Stmt Lower(const LowerArgs &T, arith::Analyzer *analyzer) const override; + LayoutMap InferLayout(const LayoutInferArgs &T, + InferLevel level) const override; + static const Op &Get(); + TileOperator Clone() const override; +}; + +class GinSignalOp : public TileOperator { +public: + TVM_FFI_DEFINE_OBJECT_REF_METHODS_NULLABLE(GinSignalOp, TileOperator, + GinSignalOpNode); + TVM_DLL GinSignalOp(Array args, Map annotations = + Map()); + static const Op &Get(); +}; + +/*! + * \brief Block until a signal reaches a cumulative threshold. + */ +class GinWaitSignalOpNode : public TileOperatorNode { +public: + PrimExpr least; ///< Cumulative count to wait for + int signal_id; ///< Signal to observe + std::string scope; ///< Cooperation scope + + TVM_FFI_DECLARE_OBJECT_INFO_FINAL("tl.GinWaitSignalOp", GinWaitSignalOpNode, + TileOperatorNode); + + static void RegisterReflection() { + namespace refl = tvm::ffi::reflection; + refl::ObjectDef() + .def_ro("least", &GinWaitSignalOpNode::least) + .def_ro("signal_id", &GinWaitSignalOpNode::signal_id) + .def_ro("scope", &GinWaitSignalOpNode::scope); + } + + Stmt Lower(const LowerArgs &T, arith::Analyzer *analyzer) const override; + LayoutMap InferLayout(const LayoutInferArgs &T, + InferLevel level) const override; + static const Op &Get(); + TileOperator Clone() const override; +}; + +class GinWaitSignalOp : public TileOperator { +public: + TVM_FFI_DEFINE_OBJECT_REF_METHODS_NULLABLE(GinWaitSignalOp, TileOperator, + GinWaitSignalOpNode); + TVM_DLL GinWaitSignalOp(Array args, + Map annotations = + Map()); + static const Op &Get(); +}; + +/*! + * \brief Make put source buffers reusable. Implies nothing about remote arrival. + */ +class GinFlushOpNode : public TileOperatorNode { +public: + std::string scope; ///< Cooperation scope + + TVM_FFI_DECLARE_OBJECT_INFO_FINAL("tl.GinFlushOp", GinFlushOpNode, + TileOperatorNode); + + static void RegisterReflection() { + namespace refl = tvm::ffi::reflection; + refl::ObjectDef().def_ro("scope", &GinFlushOpNode::scope); + } + + Stmt Lower(const LowerArgs &T, arith::Analyzer *analyzer) const override; + LayoutMap InferLayout(const LayoutInferArgs &T, + InferLevel level) const override; + static const Op &Get(); + TileOperator Clone() const override; +}; + +class GinFlushOp : public TileOperator { +public: + TVM_FFI_DEFINE_OBJECT_REF_METHODS_NULLABLE(GinFlushOp, TileOperator, + GinFlushOpNode); + TVM_DLL GinFlushOp(Array args, Map annotations = + Map()); + static const Op &Get(); +}; + +} // namespace tl +} // namespace tvm + +#endif // TVM_TL_OP_NCCL_GIN_H_ diff --git a/src/tl_templates/cuda/distributed/nccl_gin.h b/src/tl_templates/cuda/distributed/nccl_gin.h new file mode 100644 index 0000000000..04789326ea --- /dev/null +++ b/src/tl_templates/cuda/distributed/nccl_gin.h @@ -0,0 +1,325 @@ +#pragma once + +// Inter-node data movement via the NCCL GPU-Initiated Networking (GIN) device +// API. +// +// GIN is a device-side one-sided put/signal interface: a thread on one GPU +// writes into a peer GPU's registered memory without host involvement and +// without that memory being mapped into the local address space. This is the +// mechanism TileScale uses for peers that CUDA IPC/VMM cannot reach, i.e. peers +// on a different node. +// +// Addressing model, and how it differs from the intra-node path +// ------------------------------------------------------------- +// The intra-node path in distributed.h addresses a peer as +// peer_base + (local_ptr - local_base) +// which requires the peer's allocation to be mapped locally. GIN instead +// addresses memory as an (ncclWindow_t, offset) pair, where the window is a +// registration of a host-allocated buffer created by ncclCommWindowRegister. +// The window handle is symmetric across the communicator, so the same handle +// plus a byte offset names the corresponding bytes on every rank. +// +// TileScale registers its symmetric allocator arena as one window, so an +// existing local pointer converts to a GIN address by subtracting the arena +// base. That keeps the symmetric-arena invariant already relied on by the +// intra-node path and by remote TMA descriptor encoding. +// +// Requirements +// ------------ +// NCCL >= 2.28.7 for the Device API used here (ncclDevCommCreate, nccl_device +// headers). The device comm must be created with ginForceEnable and nonzero +// ginSignalCount/ginCounterCount, and the buffers involved must live in +// registered windows. +// +// This header is only included in generated code when NCCL GIN support was +// detected at build time; see TL_ENABLE_NCCL_GIN. + +// distributed.h defines the __constant__ meta_data table this header reads, and +// pulls in meta_layout.h for the offset names. +#include "distributed.h" + +#if defined(TL_ENABLE_NCCL_GIN) + +#include +#include + +#if !defined(NCCL_VERSION_CODE) || (NCCL_VERSION_CODE < 22807) +#error "TL_ENABLE_NCCL_GIN requires NCCL >= 2.28.7 for the Device API" +#endif + +namespace tl { +namespace gin { + +// The device comm is passed to kernels by value in a __grid_constant__ param by +// NCCL's own examples; TileScale instead publishes the handle through the +// distributed metadata table so existing kernel signatures are unchanged. +// meta_data holds a pointer to a device-resident ncclDevComm. +TL_DEVICE ncclDevComm const *dev_comm() { + return reinterpret_cast( + meta_data[TL_META_GLOBAL_DEV_COMM]); +} + +TL_DEVICE bool available() { return meta_data[TL_META_GLOBAL_DEV_COMM] != 0; } + +// The arena window registered for GIN addressing. Zero when GIN is unavailable +// or the backend is cudaMalloc (which cannot be registered as an NCCL window). +// ncclWindow_t is `struct ncclWindow_vidmem *`, so the handle stored as an +// integer in the table has to be reinterpreted, not static_cast. +TL_DEVICE ncclWindow_t arena_window() { + return reinterpret_cast(meta_data[TL_META_ARENA_WINDOW]); +} + +// Arena base address; subtract from a local pointer to get a GIN window offset. +TL_DEVICE uint64_t arena_base() { return meta_data[TL_META_ARENA_BASE]; } + +// Convert a local arena pointer to the GIN offset valid on every rank. +// The arena is symmetric, so the same offset names the corresponding bytes on +// any peer — the identical subtraction the intra-node peer-pointer path already +// performs, just paired with a window handle instead of a peer base. +TL_DEVICE size_t arena_offset(void const *ptr) { + return reinterpret_cast(ptr) - arena_base(); +} + +// One GIN context is one network channel with its own QP per peer, so pinning +// every CTA to context 0 serializes the whole grid onto one QP and leaves the +// node's other NICs idle. Spreading CTAs across contexts is what would turn one +// channel's bandwidth into the fabric's -- and is the main throughput headroom +// left, since a single QP measures ~23 GB/s where one NIC's line rate is 50. +// +// Signal state is per context, and that is what makes spreading work at all. +// A put issued on sender context i increments the receiver's signal *through +// context i*, so a CTA sitting on one of C contexts observes only 1/C of the +// rank's arrivals. `wait_signal` therefore divides the caller's grid-wide target +// by context_span(); see the comment there. +// +// Measured 2026-08-03, allgather, 64 MB shard, bf16, two ranks: +// contexts 1 -> 24.3 GB/s 2 -> 44.0 GB/s 4 -> 44.5 GB/s +// So a single QP is the throughput wall and two contexts nearly double it, +// saturating by two on this fabric. Throughput is flat in chunk count from 1 to +// 4, confirming the QP rather than message size is the limit. +// +// This took three wrong turns worth recording, because each looked convincing: +// 1. Rotating on `ginContextCount` from ncclDevCommRequirements. That field is +// documented as only a hint -- "the actual context count in the devcomm may +// not match" -- and the request here (8) is NOT what gets granted (4). +// Indexing past the granted count hangs. +// 2. Blaming the kernel cache. Real hazard -- the cache key hashes the script, +// args, target and compile flags but not this header, so editing it leaves +// stale binaries behind -- but not this bug; TL_GIN_CONTEXTS is a -D flag +// and so is in the key. +// 3. "Disproving" per-context signals from a 332 GB/s reading. The reading was +// real and impossible -- above the ~63 GB/s PCIe Gen5 x16 cap on one GPU's +// egress -- but the cause was dividing by the *requested* 8 while the device +// had clamped to 4, so the wait was 2x too weak and the kernel returned +// early. The conclusion drawn from it was wrong. +// +// The lesson each time: get the number off the device (TL_GIN_DEBUG=1) instead of +// inferring it, and sanity-check any bandwidth against `nvidia-smi topo -m` +// before believing it. An under-target is silent -- the host's correctness check +// syncs and runs a reference collective first, which gives in-flight RDMA writes +// time to land before the comparison reads the buffer. +// +// TL_GIN_CONTEXTS is a -D compile flag, so it is part of the kernel cache key -- +// which also stops a stale entry from silently answering for a different policy. +// 0, 1 pin to context 0 +// n > 1 spread CTAs over min(n, granted) contexts +#ifndef TL_GIN_CONTEXTS +#define TL_GIN_CONTEXTS 0 +#endif + +// Splits the context choice between issuing and waiting, to tell apart the two +// candidate reasons a multi-context run hangs: puts on a non-zero context never +// delivering, versus a wait on a non-zero context never observing a signal that +// did arrive. With TL_GIN_WAIT_CTX0=1 the puts spread over contexts while every +// wait stays on context 0. If that passes on the honest target, signals are +// communicator-wide and only the waiting side is context-sensitive. +#ifndef TL_GIN_WAIT_CTX0 +#define TL_GIN_WAIT_CTX0 0 +#endif + +// TL_GIN_DEBUG=1 prints what the devcomm actually granted. Direct evidence beats +// inference here: every theory about why multiple contexts hang has hinged on how +// many contexts exist and which one a CTA ends up on, and nothing else exposes +// those. One line per CTA, from thread 0 only. +#ifndef TL_GIN_DEBUG +#define TL_GIN_DEBUG 0 +#endif + +// How many contexts this kernel actually spreads over: the requested count +// clamped to what the devcomm granted. Measured granted = 4 on these nodes even +// though the allocator asks for 8, which is why the clamp matters -- and why the +// host must never assume its own request is the number in play. +TL_DEVICE uint32_t context_span() { + ncclGin probe(*dev_comm(), 0); + uint32_t want = TL_GIN_CONTEXTS < 1 ? 1u : static_cast(TL_GIN_CONTEXTS); + uint32_t const granted = probe.nContexts; + if (granted < want) { + want = granted; + } + return want < 1 ? 1u : want; +} + +TL_DEVICE ncclGin make_gin() { + uint32_t const want = context_span(); + uint32_t const use = static_cast(blockIdx.x) % want; +#if TL_GIN_DEBUG + if (threadIdx.x == 0) { + printf("[gin] block=%u want=%u use=%u\n", blockIdx.x, want, use); + } +#endif + return ncclGin(*dev_comm(), static_cast(use)); +} + +// The gin used for waits. See TL_GIN_WAIT_CTX0. +TL_DEVICE ncclGin make_gin_for_wait() { +#if TL_GIN_WAIT_CTX0 + return ncclGin(*dev_comm(), 0); +#else + return make_gin(); +#endif +} + +// A put whose completion increments `signal` on the destination rank. The signal +// becomes visible to the peer only after this put's payload, and the payloads of +// preceding puts to that peer on this context, have settled -- so a peer that +// waits on the signal observes the data. +// +// dst_offset/src_offset are byte offsets into the registered windows. `peer` is +// a rank within `team`. +template +TL_DEVICE void put_signal(Coop coop, ncclTeam team, int peer, + ncclWindow_t dst_window, size_t dst_offset, + ncclWindow_t src_window, size_t src_offset, + size_t bytes, ncclGinSignal_t signal) { + make_gin() + .put(team, peer, dst_window, dst_offset, src_window, src_offset, bytes, + ncclGin_SignalInc{signal}, ncclGin_None{}, coop); +} + +// A put with no remote notification. Pair with a later signal() or a barrier +// when the peer needs to know the data arrived. +template +TL_DEVICE void put(Coop coop, ncclTeam team, int peer, ncclWindow_t dst_window, + size_t dst_offset, ncclWindow_t src_window, + size_t src_offset, size_t bytes) { + make_gin() + .put(team, peer, dst_window, dst_offset, src_window, src_offset, bytes, + ncclGin_None{}, ncclGin_None{}, coop); +} + +// Increment `signal` on `peer` without moving payload. Ordered after this +// context's preceding puts to that peer. +template +TL_DEVICE void signal_peer(Coop coop, ncclTeam team, int peer, + ncclGinSignal_t signal) { + make_gin() + .signal(team, peer, ncclGin_SignalInc{signal}, coop); +} + +// Block until `signal` has been incremented at least `least` times in total. +// Signals are cumulative and compared with rolling arithmetic, so callers track +// an expected running total rather than resetting between phases. +// +// `least` is the GRID-WIDE arrival count -- what the whole rank receives per +// launch. Signal state is per context, though: a put issued on sender context i +// increments the receiver's signal through context i, so a CTA sitting on one of +// `context_span()` contexts only ever observes its own share. Scaling here rather +// than on the host keeps the honest number in the caller and puts the division +// where the granted context count is actually known. +// +// Getting this wrong is silent in one direction: too small a target lets the wait +// return before the payload lands, and the host's correctness check still passes +// because it syncs and runs a reference collective first, which gives the +// in-flight writes time to arrive. The tell is bandwidth above what the hardware +// can carry -- 332 GB/s appeared this way, against a ~63 GB/s PCIe Gen5 x16 cap. +template +TL_DEVICE void wait_signal(Coop coop, ncclGinSignal_t signal, uint64_t least) { + uint64_t const span = static_cast(context_span()); + make_gin_for_wait().waitSignal(coop, signal, least / span); +} + +// Wait for one more increment than last observed, advancing the signal's shadow. +// This is the convenient form for a consumer draining a stream of arrivals. +template +TL_DEVICE void wait_signal_next(Coop coop, ncclGinSignal_t signal) { + make_gin_for_wait() + .waitSignalMeetShadow(coop, signal); +} + +// Make source buffers from this coop's puts safe to overwrite. This does NOT +// imply the data has landed remotely; use a signal for that. +template TL_DEVICE void flush(Coop coop) { + make_gin().flush(coop); +} + +// Reset a signal to zero along with its shadow. Must not race with concurrent +// increments to the same signal. +TL_DEVICE void reset_signal(ncclGinSignal_t signal) { + make_gin().resetSignal(signal); +} + +// The communicator-wide team, whose ranks are global ranks. +TL_DEVICE ncclTeam world_team() { return ncclTeamWorld(*dev_comm()); } + +// Convenience wrappers for whole-CTA and single-thread issue, which are the two +// shapes generated code uses today. +TL_DEVICE void put_signal_cta(int peer, ncclWindow_t dst_window, + size_t dst_offset, ncclWindow_t src_window, + size_t src_offset, size_t bytes, + ncclGinSignal_t signal) { + put_signal(ncclCoopCta(), world_team(), peer, dst_window, dst_offset, + src_window, src_offset, bytes, signal); +} + +TL_DEVICE void wait_signal_cta(ncclGinSignal_t signal, uint64_t least) { + wait_signal(ncclCoopCta(), signal, least); +} + +// --------------------------------------------------------------------------- +// Entry points used by generated code. +// +// These take the coop as a *template* parameter and the buffers as ordinary +// pointers, which is what the lowering in src/op/nccl_gin.cc can express: it +// emits a call_extern whose callee is a string, so the scope has to be baked +// into the name rather than passed as a constructed object, and it has no way to +// name an ncclWindow_t or compute a window offset at compile time. +// +// The window and the base used for the offsets are both read from the metadata +// table here, on the device. Both pointers must be inside the allocator arena -- +// a shared-memory or fragment buffer has no window and would produce a wild +// offset rather than an error. +// --------------------------------------------------------------------------- + +template +TL_DEVICE void put_addr(int peer, void *dst, void const *src, size_t bytes) { + ncclWindow_t window = arena_window(); + put(Coop(), world_team(), peer, window, arena_offset(dst), window, + arena_offset(src), bytes); +} + +template +TL_DEVICE void put_signal_addr(int peer, void *dst, void const *src, + size_t bytes, int signal) { + ncclWindow_t window = arena_window(); + put_signal(Coop(), world_team(), peer, window, arena_offset(dst), window, + arena_offset(src), bytes, + static_cast(signal)); +} + +template +TL_DEVICE void signal_peer(int peer, int signal) { + signal_peer(Coop(), world_team(), peer, + static_cast(signal)); +} + +template +TL_DEVICE void wait_signal(int signal, uint64_t least) { + wait_signal(Coop(), static_cast(signal), least); +} + +template TL_DEVICE void flush() { flush(Coop()); } + +} // namespace gin +} // namespace tl + +#endif // TL_ENABLE_NCCL_GIN diff --git a/tilelang/language/distributed/__init__.py b/tilelang/language/distributed/__init__.py index f87da8303b..40e7e4ead0 100644 --- a/tilelang/language/distributed/__init__.py +++ b/tilelang/language/distributed/__init__.py @@ -37,6 +37,12 @@ multimem_signal_add, ) +# Exposed as a namespace (T.nccl_gin.put) rather than flattened, because these +# names would otherwise collide with the intra-node put/get and read as though +# they were interchangeable with them. They are not: GIN addresses windows, not +# peer pointers, and takes global rather than local ranks. +from . import nccl_gin # noqa: F401 + __all__ = [ "get_rank", "get_num_ranks", @@ -63,4 +69,5 @@ "multimem_tma_store", "multimem_signal", "multimem_signal_add", + "nccl_gin", ] diff --git a/tilelang/language/distributed/nccl_gin.py b/tilelang/language/distributed/nccl_gin.py new file mode 100644 index 0000000000..36710488a5 --- /dev/null +++ b/tilelang/language/distributed/nccl_gin.py @@ -0,0 +1,164 @@ +"""Inter-node communication primitives backed by the NCCL GIN device API. + +These ops address memory as ``(window, offset)`` rather than as a peer pointer, +which is what makes them work across nodes: a remote node's allocation has no +local virtual address, so the intra-node ``peer_base + offset`` arithmetic in +:mod:`comm` cannot name it. TileScale registers the whole allocator arena as one +NCCL window, so a local arena pointer converts to a remote address by +subtracting the arena base -- the same subtraction the intra-node path performs, +paired with a window handle instead of a peer base. + +Both buffers must live in the allocator arena. A local ``T.alloc_shared`` or +``T.alloc_fragment`` buffer is not in a registered window and cannot be a GIN +source or destination. + +``peer`` is a **global** rank, since GIN puts are issued against the +communicator-wide team. Passing a peer on the local node is allowed and works, +but it routes through the network stack rather than NVLink; prefer +:func:`~tilelang.language.distributed.comm.put_block` for intra-node traffic. + +Requires NCCL >= 2.28.7 and a devcomm published by the allocator. When GIN was +not compiled in, the generated kernel will not contain these calls at all -- +build-time detection gates the device header. +""" + +from __future__ import annotations + +from tvm import tirx +from tvm.tirx import PrimExpr, IntImm, address_of + +# GIN 2.28.9 exposes no remote-read operation: the device class has put, +# putValue, signal, and the signal/counter waits, but nothing that pulls bytes +# from a peer. A `get` therefore cannot be one-sided the way intra-node +# `get_block` is -- it needs the data's owner to issue a put. Expressing that as +# a `get` would hide a required remote-side call behind a local-looking op, so it +# is deliberately absent. Model a pull as the owner putting plus a signal. + +__all__ = [ + "put", + "put_signal", + "signal", + "wait_signal", + "flush", +] + + +def _coop(scope: str) -> str: + """Validate a cooperation scope and return it. + + The scope becomes an ``ncclCoop*`` template argument, so an unknown value + would surface as an nvcc template error in generated code rather than here. + """ + if scope not in ("thread", "warp", "block"): + raise ValueError(f"scope must be one of 'thread', 'warp', or 'block', got {scope!r}") + return scope + + +def put( + src: PrimExpr, + dst: PrimExpr, + size: PrimExpr, + peer: PrimExpr | IntImm, + scope: str = "block", +): + """Write ``size`` elements from local ``src`` into ``dst`` on ``peer``. + + ``size`` counts elements, not bytes, matching + :func:`~tilelang.language.distributed.comm.put_block`. ``src`` and ``dst`` + must have the same dtype. + + One-sided and asynchronous: the call returns before the data has landed. + Nothing tells the peer the write happened -- pair with :func:`signal`, or use + :func:`put_signal` to fuse the notification into the put. Reusing ``src`` + requires a :func:`flush` first. + + ``dst`` is indexed with the *peer's* view of the buffer, which is the same + index the local rank would use because the arena is symmetric. + """ + # Validate before touching the buffers so a bad scope reports the scope + # rather than whatever address_of makes of the arguments. + coop = _coop(scope) + return tirx.call_intrin( + "handle", + tirx.op.Op.get("tl.tileop.gin_put"), + address_of(src), + address_of(dst), + size, + peer, + 0, # signal id, unused without a remote action + 0, # no remote signal + coop, + ) + + +def put_signal( + src: PrimExpr, + dst: PrimExpr, + size: PrimExpr, + peer: PrimExpr | IntImm, + signal_id: int = 0, + scope: str = "block", +): + """Like :func:`put`, and increment ``signal_id`` on ``peer`` once it lands. + + The increment is ordered after this put's payload and after any preceding + puts to the same peer on the same context, so a peer released by + :func:`wait_signal` is guaranteed to observe the bytes. This ordering is why + a fused put+signal is preferred over a separate :func:`signal`. + """ + coop = _coop(scope) + return tirx.call_intrin( + "handle", + tirx.op.Op.get("tl.tileop.gin_put"), + address_of(src), + address_of(dst), + size, + peer, + int(signal_id), + 1, # increment signal on arrival + coop, + ) + + +def signal(peer: PrimExpr | IntImm, signal_id: int = 0, scope: str = "block"): + """Increment ``signal_id`` on ``peer`` without moving payload. + + Ordered after this context's preceding puts to that peer, so it can act as a + completion marker for a batch of :func:`put` calls. + """ + return tirx.call_intrin( + "handle", + tirx.op.Op.get("tl.tileop.gin_signal"), + peer, + int(signal_id), + _coop(scope), + ) + + +def wait_signal(least: PrimExpr, signal_id: int = 0, scope: str = "block"): + """Block until ``signal_id`` has been incremented at least ``least`` times. + + Signals are cumulative running totals compared with rolling arithmetic, not + flags -- they are not consumed by a wait. A kernel that waits repeatedly + tracks an increasing expected total rather than resetting between phases. + """ + return tirx.call_intrin( + "handle", + tirx.op.Op.get("tl.tileop.gin_wait_signal"), + least, + int(signal_id), + _coop(scope), + ) + + +def flush(scope: str = "block"): + """Wait until this coop's put source buffers are safe to overwrite. + + This says nothing about remote visibility; only a signal does. Use it before + rewriting a send buffer, not to establish that a peer can read the data. + """ + return tirx.call_intrin( + "handle", + tirx.op.Op.get("tl.tileop.gin_flush"), + _coop(scope), + ) From 504bf27f7b30cf1e41e7e3fb75f816499d2774f2 Mon Sep 17 00:00:00 2001 From: Rachmanino <18805904201@163.com> Date: Mon, 3 Aug 2026 23:13:16 +0800 Subject: [PATCH 04/30] [Feat] Register the allocator arena as one NCCL window and create a GIN devcomm GIN addresses memory as an `(ncclWindow_t, byte offset)` pair, not a raw pointer, because a remote rank's allocation has no local virtual address. The window handle is not per-peer: one local handle plus a peer index names any rank's bytes. So the whole allocator arena is registered once, collectively, at allocator init rather than per tensor -- `ncclCommWindowRegister` is a collective, and per-tensor registration would turn every `allocate_tensor` into a world-wide barrier plus a handle table. Because the arena is symmetric, `local_ptr - arena_base` is the offset valid on every rank, the same subtraction the intra-node peer-pointer path already performs. Window registration rejects `cudaMalloc` memory: measured on NCCL 2.28.9, a plain `cudaMalloc` pointer fails with "invalid argument" under both `NCCL_WIN_COLL_SYMMETRIC` and `flags=0`, while a VMM-mapped arena registers cleanly. The allocator therefore raises rather than silently degrading when a window is requested on the cudaMalloc backend, and `shared_memory.cc` exposes the VMM capability probe the allocator needs to check that up front. Ordering is the one non-obvious constraint. `ncclDevCommCreate` needs a communicator that still supports symmetric memory, and a torch `ProcessGroupNCCL` communicator loses that after its first collective -- the call then *segfaults* rather than returning an error, so it cannot be attempted and recovered from. Allocator construction runs several collectives before GIN setup, so reusing the caller's group crashes every time. `_init_arena_window` instead makes a private `dist.new_group(backend="nccl", device_id=...)`; `device_id` makes the communicator eager, so its pointer is valid immediately instead of being created by the very collective that would invalidate it. The devcomm is created on that group *before* window registration, which must then use the same comm, and teardown runs in reverse. `init_dist` now returns node topology alongside the group so the allocator can tell node-local peers from inter-node ones, and `get_allocator` takes it as `node_info`. Single-node callers are unaffected: `num_nodes == 1` takes the existing IPC/VMM path untouched and no window is registered unless `TILESCALE_USE_GIN` asks for one. --- src/shared_memory/shared_memory.cc | 60 ++- tilelang/__init__.py | 2 + tilelang/distributed/__init__.py | 3 + tilelang/distributed/allocator.py | 411 ++++++++++++++++-- tilelang/distributed/host.py | 114 ++++- tilelang/distributed/nccl_window.py | 395 +++++++++++++++++ .../distributed/shared_memory/__init__.py | 2 + tilelang/engine/lower.py | 9 + 8 files changed, 951 insertions(+), 45 deletions(-) create mode 100644 tilelang/distributed/nccl_window.py diff --git a/src/shared_memory/shared_memory.cc b/src/shared_memory/shared_memory.cc index de6ad03c38..f2f3cff013 100644 --- a/src/shared_memory/shared_memory.cc +++ b/src/shared_memory/shared_memory.cc @@ -509,6 +509,22 @@ static bool can_create_multicast_object(int device_count) { // ---------- VMM malloc/free ---------- +// Fabric handles are what the intra-node peer-mapping path needs, but creating +// one requires an IMEX channel: without /dev/nvidia-caps-imex-channels, +// cuMemCreate(FABRIC) fails with CUDA_ERROR_NOT_PERMITTED even though the device +// reports CU_DEVICE_ATTRIBUTE_HANDLE_TYPE_FABRIC_SUPPORTED. A POSIX-FD handle +// type still yields the driver-level (VMM) allocation that +// ncclCommWindowRegister requires, so GIN works there while cross-process peer +// mapping does not. Probed once: the answer cannot change within a process. +static CUmemAllocationHandleType vmm_handle_type(CUdevice device) { + static const CUmemAllocationHandleType cached = [device] { + return can_create_fabric_allocation(device) + ? CU_MEM_HANDLE_TYPE_FABRIC + : CU_MEM_HANDLE_TYPE_POSIX_FILE_DESCRIPTOR; + }(); + return cached; +} + static int64_t vmm_malloc_impl(int64_t size_raw) { const size_t requested_size = checked_positive_size(size_raw, "size"); @@ -518,7 +534,7 @@ static int64_t vmm_malloc_impl(int64_t size_raw) { CUmemAllocationProp prop = {}; prop.type = CU_MEM_ALLOCATION_TYPE_PINNED; prop.location.type = CU_MEM_LOCATION_TYPE_DEVICE; - prop.requestedHandleTypes = CU_MEM_HANDLE_TYPE_FABRIC; + prop.requestedHandleTypes = vmm_handle_type(device); prop.location.id = device; size_t granularity = 0; @@ -685,6 +701,46 @@ static bool supports_vmm_fabric_impl() { return true; } +// Whether a VMM (driver-level) allocation can be created at all, by any handle +// type. This is a weaker question than supports_vmm_fabric: NCCL window +// registration for GIN needs only a VMM-backed arena, while mapping a peer's +// arena into this process needs an *exportable* fabric handle. A node without an +// IMEX channel answers true here and false there. +static bool supports_vmm_impl() { + if (!SharedMemoryDriverAPI::Get()->HasVMM()) { + return false; + } + + int device_count = 0; + cudaError_t err = cudaGetDeviceCount(&device_count); + if (err != cudaSuccess || device_count == 0) + return false; + + CUdevice device; + if (cuCtxGetDevice(&device) != CUDA_SUCCESS) + return false; + + CUmemAllocationProp prop = {}; + prop.type = CU_MEM_ALLOCATION_TYPE_PINNED; + prop.location.type = CU_MEM_LOCATION_TYPE_DEVICE; + prop.requestedHandleTypes = vmm_handle_type(device); + prop.location.id = device; + + size_t granularity = 0; + if (cuMemGetAllocationGranularity(&granularity, &prop, + CU_MEM_ALLOC_GRANULARITY_MINIMUM) != + CUDA_SUCCESS || + granularity == 0) { + return false; + } + + CUmemGenericAllocationHandle handle; + if (cuMemCreate(&handle, granularity, &prop, 0) != CUDA_SUCCESS) { + return false; + } + return cuMemRelease(handle) == CUDA_SUCCESS; +} + static bool supports_multicast_impl() { if (!SharedMemoryDriverAPI::Get()->HasMulticast()) { return false; @@ -986,6 +1042,8 @@ TVM_FFI_STATIC_INIT_BLOCK() { sync_ipc_handles_impl); // Support detection + refl::GlobalDef().def("tl.shared_memory.supports_vmm", + supports_vmm_impl); refl::GlobalDef().def("tl.shared_memory.supports_vmm_fabric", supports_vmm_fabric_impl); refl::GlobalDef().def("tl.shared_memory.supports_multicast", diff --git a/tilelang/__init__.py b/tilelang/__init__.py index 5337f26ff7..c52be6c5fc 100644 --- a/tilelang/__init__.py +++ b/tilelang/__init__.py @@ -226,6 +226,7 @@ def get_allocator( group=None, use_vmm: bool | None = None, mcast_size: int | None = None, + node_info=None, ): """Create a distributed allocator without importing its CUDA helpers eagerly.""" from .distributed.allocator import get_allocator as _get_allocator @@ -239,6 +240,7 @@ def get_allocator( group=group, use_vmm=use_vmm, mcast_size=mcast_size, + node_info=node_info, ) diff --git a/tilelang/distributed/__init__.py b/tilelang/distributed/__init__.py index 7cdab9f885..198941c92f 100644 --- a/tilelang/distributed/__init__.py +++ b/tilelang/distributed/__init__.py @@ -21,6 +21,7 @@ "_create_tensor": (".shared_memory", "_create_tensor"), "_create_ipc_handle": (".shared_memory", "_create_ipc_handle"), "_sync_ipc_handles": (".shared_memory", "_sync_ipc_handles"), + "_supports_vmm": (".shared_memory", "_supports_vmm"), "_supports_vmm_fabric": (".shared_memory", "_supports_vmm_fabric"), "_vmm_malloc": (".shared_memory", "_vmm_malloc"), "_vmm_free": (".shared_memory", "_vmm_free"), @@ -29,6 +30,8 @@ "_close_vmm_handle": (".shared_memory", "_close_vmm_handle"), "_sync_vmm_handles": (".shared_memory", "_sync_vmm_handles"), "_supports_multicast": (".shared_memory", "_supports_multicast"), + "nccl_supports_device_api": (".nccl_window", "supports_device_api"), + "nccl_version": (".nccl_window", "nccl_version"), } __all__ = list(_EXPORTS) diff --git a/tilelang/distributed/allocator.py b/tilelang/distributed/allocator.py index 2ee23946a6..5af20885cc 100644 --- a/tilelang/distributed/allocator.py +++ b/tilelang/distributed/allocator.py @@ -7,6 +7,7 @@ import operator import threading import warnings +from typing import TYPE_CHECKING import torch import torch.distributed as dist @@ -20,6 +21,7 @@ _create_vmm_handle, _open_vmm_handle, _close_vmm_handle, + _supports_vmm, _supports_vmm_fabric, _supports_multicast, _mc_create, @@ -34,8 +36,26 @@ ) from tilelang.utils.target import parse_device +if TYPE_CHECKING: + from tilelang.distributed.host import NodeTopology + __all__ = ["BaseAllocator", "get_allocator"] +# Distributed metadata table layout. Keep in sync with +# src/tl_templates/cuda/distributed/meta_layout.h, which documents the fields +# and is shared by the device helpers and the host remote-TMA remapping. +_META_GLOBAL_RANK = 0 +_META_GLOBAL_WORLD_SIZE = 1 +_META_NODE_RANK = 2 +_META_NUM_NODES = 3 +_META_LOCAL_RANK = 4 +_META_LOCAL_WORLD_SIZE = 5 +_META_GLOBAL_DEV_COMM = 6 +_META_INTERNODE_DEV_COMM = 7 +_META_ARENA_WINDOW = 8 +_META_ARENA_BASE = 9 +_META_PEER_BASE = 10 + _dtype_to_str = { torch.float32: "float32", torch.float16: "float16", @@ -120,14 +140,46 @@ def _parse_bool_env(name: str, value: str) -> bool: ) -def _resolve_use_vmm(use_vmm: bool | None, is_distributed: bool = False) -> bool: - """Resolve whether to use VMM based on env var and hardware support.""" +def _resolve_use_vmm( + use_vmm: bool | None, + is_distributed: bool = False, + needs_window: bool = False, +) -> bool: + """Resolve whether to use VMM based on env var and hardware support. + + Two different capabilities are involved. Mapping a peer's arena into this + process needs an *exportable fabric* handle, which requires an IMEX channel. + Registering the arena as an NCCL window for GIN needs only a driver-level + (VMM) allocation, which a POSIX-FD handle type also provides. So on a node + with no IMEX channel, fabric is unavailable but GIN is still reachable -- + hence the weaker ``_supports_vmm`` check when a window is wanted. + """ env_val = os.environ.get("TILESCALE_USE_VMM", None) if env_val is not None: return _parse_bool_env("TILESCALE_USE_VMM", env_val) if use_vmm is not None: return use_vmm - return is_distributed and _supports_vmm_fabric() + if is_distributed and _supports_vmm_fabric(): + return True + # A cudaMalloc arena cannot be registered as an NCCL window, so GIN would be + # dead on arrival; prefer VMM whenever it can be allocated at all. + return bool(needs_window) and _supports_vmm() + + +def _resolve_register_window(is_multi_node: bool) -> tuple[bool, bool]: + """Resolve whether to register the arena as an NCCL window for GIN. + + Returns ``(requested, required)``. ``required`` marks an explicit opt-in via + ``TILESCALE_USE_GIN``, where a failure to register must raise rather than + silently fall back -- inter-node traffic would otherwise be quietly dead. + Auto mode only attempts registration for a genuine multi-node job, so + single-node runs never take on an NCCL Device API dependency. + """ + env_val = os.environ.get("TILESCALE_USE_GIN", None) + if env_val is not None: + requested = _parse_bool_env("TILESCALE_USE_GIN", env_val) + return requested, requested + return is_multi_node, False class BaseAllocator: @@ -144,6 +196,7 @@ def __init__( align: int = 256, use_vmm: bool | None = None, mcast_size: int | None = None, + node_info: NodeTopology | None = None, ) -> None: # Keep potentially failing local parsing inside the first collective # stage below. Otherwise one rank can exit while its peers wait forever @@ -159,6 +212,8 @@ def __init__( self._local_rank = local_rank self._num_local_ranks = num_local_ranks self._group = group + self._node_info = node_info + self._is_multi_node = (node_info is not None and node_info.num_nodes > 1) self._align = align self._lock = threading.RLock() self._mcast_size_requested = mcast_size @@ -186,6 +241,23 @@ def __init__( self._use_multicast = False self._group_size = 1 self._group_root_global_rank = 0 + # NCCL/GIN arena window state. Zero means "no window registered", which + # is the normal single-node case and keeps the intra-node peer-pointer + # path free of any NCCL Device API dependency. + self._arena_window = 0 + self._arena_window_comm = 0 + self._arena_window_size = 0 + # The GIN devcomm, created only once the window registers. Holding the + # object alive is what keeps its device pointer valid. + self._dev_comm = None + # GIN gets its own process group. ncclDevCommCreate needs a communicator + # that still supports symmetric memory, and a torch communicator loses + # that after its first collective -- the call then *segfaults* rather + # than failing, so it cannot be attempted and recovered from. Allocator + # construction runs several collectives on the caller's group before GIN + # setup, so a private group is the only safe option. + self._gin_group = None + self._register_window, self._require_window = _resolve_register_window(self._is_multi_node) if self._is_distributed: if self._group is None: @@ -193,8 +265,17 @@ def __init__( if not dist.is_initialized(): raise RuntimeError("torch.distributed must be initialized before creating a distributed allocator") - self._group_size = dist.get_world_size(self._group) - group_rank = dist.get_rank(self._group) + # For multi-node, use node-local group for IPC/VMM operations + # For single-node, use the provided group as-is + if self._is_multi_node: + self._allocator_group = self._node_info.node_local_group + self._global_group = self._group + else: + self._allocator_group = self._group + self._global_group = self._group + + self._group_size = dist.get_world_size(self._allocator_group) + group_rank = dist.get_rank(self._allocator_group) try: if self._is_distributed: @@ -214,6 +295,7 @@ def __init__( self._collective_stage("allocate base storage", self._alloc_base) if self._mcast_size_requested is not None: self._init_multicast_buffer() + self._init_arena_window() self._init_table() else: self._prepare_local_configuration() @@ -245,7 +327,11 @@ def positive_integer(value, name: str) -> int: if self._mcast_size_requested is not None: self._mcast_size_requested = positive_integer(self._mcast_size_requested, "mcast_size") self._device = parse_device(self._device_request) - self._use_vmm = _resolve_use_vmm(self._use_vmm_requested, self._is_distributed) + self._use_vmm = _resolve_use_vmm( + self._use_vmm_requested, + self._is_distributed, + needs_window=self._register_window, + ) def _validate_distributed_configuration(self, group_rank: int) -> None: """Collectively validate invariants before any allocation is created.""" @@ -259,16 +345,19 @@ def _validate_distributed_configuration(self, group_rank: int) -> None: "device": self._device, } configurations = [None] * self._group_size - dist.all_gather_object(configurations, local_config, group=self._group) + dist.all_gather_object(configurations, local_config, group=self._allocator_group) failures = [] reference = configurations[0] invariant_keys = ("size", "align", "use_vmm", "mcast_size", "num_local_ranks") for rank, config in enumerate(configurations): + # The allocator group is node-local, so a rank's index within it is + # its local rank in both the single-node and multi-node cases. if config["local_rank"] != rank: failures.append(f"rank {rank} reports local_rank={config['local_rank']!r}") - if config["device"] != rank: - failures.append(f"rank {rank} reports device={config['device']!r}") + if config["device"] != config["local_rank"]: + failures.append(f"rank {rank} reports device={config['device']!r}, expected local_rank={config['local_rank']!r}") + for key in invariant_keys: if config[key] != reference[key]: failures.append(f"rank {rank} reports {key}={config[key]!r}, expected {reference[key]!r}") @@ -285,9 +374,12 @@ def _validate_distributed_configuration(self, group_rank: int) -> None: self._device_ids = [config["device"] for config in configurations] def _resolve_group_root(self) -> None: + # For multi-node, resolve global rank of node-local group root + # For single-node, resolve as before + group_to_resolve = self._allocator_group if self._is_multi_node else self._group if hasattr(dist, "get_global_rank"): - self._group_root_global_rank = dist.get_global_rank(self._group, 0) - elif self._group is not dist.group.WORLD: + self._group_root_global_rank = dist.get_global_rank(group_to_resolve, 0) + elif group_to_resolve is not dist.group.WORLD: raise RuntimeError("this PyTorch version cannot resolve the global rank of a subgroup") def _prepare_multicast(self) -> None: @@ -307,6 +399,23 @@ def _rollback_failed_initialization(self) -> RuntimeError | None: f"to avoid dangling peer mappings ({rollback_error})" ) + # Any one of the three is worth a rollback, and they are created in + # order (group, devcomm, window), so a failure at any point leaves a + # different subset behind. _free_arena_window releases whichever exist. + if self._arena_window or self._dev_comm is not None or self._gin_group is not None: + try: + self._collective_stage( + "rollback arena NCCL window", + self._free_arena_window, + group=self._global_group, + ) + except Exception as rollback_error: # noqa: BLE001 + return RuntimeError( + "distributed allocator initialization failed and deregistration of the " + "arena NCCL window also failed; owned allocations were intentionally " + f"retained to avoid freeing memory behind a live window ({rollback_error})" + ) + try: self._collective_stage("rollback owned allocations", self._free_local_allocations) except Exception as rollback_error: # noqa: BLE001 @@ -324,8 +433,16 @@ def _rollback_failed_initialization(self) -> RuntimeError | None: ) return None - def _collective_stage(self, stage: str, operation): - """Run a local operation and make every rank observe any exception.""" + def _collective_stage(self, stage: str, operation, group: dist.ProcessGroup | None = None): + """Run a local operation and make every rank observe any exception. + + ``group`` defaults to the node-local allocator group. Stages whose + collective spans nodes (NCCL window registration) must pass the global + group, or a failure on one node leaves the other node waiting forever. + """ + if group is None: + group = self._allocator_group + group_size = dist.get_world_size(group) local_exception = None result = None try: @@ -336,8 +453,8 @@ def _collective_stage(self, stage: str, operation): local_status = None if local_exception is not None: local_status = f"{type(local_exception).__name__}: {local_exception}" - statuses = [None] * self._group_size - dist.all_gather_object(statuses, local_status, group=self._group) + statuses = [None] * group_size + dist.all_gather_object(statuses, local_status, group=group) failures = [f"rank {rank}: {status}" for rank, status in enumerate(statuses) if status is not None] if failures: error = RuntimeError(f"distributed allocator stage '{stage}' failed ({'; '.join(failures)})") @@ -363,6 +480,104 @@ def _alloc_base(self): raise RuntimeError(f"cudaMalloc failed: {rc} {msg.decode() if msg else ''}") self._ptr.value = self._base_ptr.value + def _init_arena_window(self): + """Register the whole arena as one NCCL window for GIN inter-node access. + + Registration is collective over the *global* group and 4096-byte + aligned, so it happens once here rather than per tensor. One handle plus + a peer index names any rank's bytes, and the arena is symmetric, so + ``local_ptr - arena_base`` is the offset valid on every rank. + """ + if not self._register_window: + return + + from tilelang.distributed import nccl_window as _win + + def check_support(): + if not _win.supports_device_api(): + raise RuntimeError(_win.unavailable_reason()) + # Measured on node071 with NCCL 2.28.9: ncclCommWindowRegister + # rejects a plain cudaMalloc pointer with "invalid argument" under + # both NCCL_WIN_COLL_SYMMETRIC and flags=0, while a VMM-mapped arena + # registers cleanly. The NIC needs the driver-level allocation that + # only the VMM path (or ncclMemAlloc) produces. + if not self._use_vmm: + raise RuntimeError( + "NCCL window registration requires a VMM-backed arena; the cudaMalloc " + "backend cannot be registered. Enable VMM (unset TILESCALE_USE_VMM=0) " + "for inter-node GIN support") + + try: + self._collective_stage( + "check NCCL Device API support", check_support, group=self._global_group) + except Exception as exc: # noqa: BLE001 - optional unless explicitly requested + if self._require_window: + raise RuntimeError( + "TILESCALE_USE_GIN requested NCCL window registration, but the NCCL " + f"Device API is unavailable: {exc}" + ) from exc + warnings.warn( + "inter-node run without GIN: the arena could not be registered as an NCCL " + f"window ({exc}); inter-node primitives will be unavailable", + RuntimeWarning, + stacklevel=2, + ) + self._register_window = False + return + + def make_gin_group(): + # device_id makes the new communicator eager, so _comm_ptr() is + # non-null right away; without it the comm is created lazily on + # first use, and the first use would be the collective that + # invalidates it for ncclDevCommCreate. + self._gin_group = dist.new_group( + ranks=dist.get_process_group_ranks(self._global_group), + backend="nccl", + device_id=torch.device("cuda", self._device), + ) + + # new_group is collective over WORLD and every rank must reach it, so it + # runs as its own stage rather than inside the registration closure. + self._collective_stage("create GIN process group", make_gin_group, group=self._global_group) + + def create_dev_comm(): + # Before any collective touches this group -- see make_gin_group. + comm_ptr = _win.get_comm_ptr(self._gin_group) + if not comm_ptr: + raise RuntimeError( + "could not obtain the raw ncclComm_t for the GIN process group; " + "torch.distributed must expose ProcessGroupNCCL._comm_ptr()" + ) + self._arena_window_comm = comm_ptr + self._dev_comm = _win.create_dev_comm(comm_ptr) + + self._collective_stage("create GIN devcomm", create_dev_comm, group=self._global_group) + + def register(): + base = self._base_ptr.value + if base is None or base == 0: + raise RuntimeError("arena base pointer is null; cannot register an NCCL window") + if base % _win.NCCL_WIN_REQUIRED_ALIGNMENT: + raise RuntimeError( + f"arena base {base:#x} is not {_win.NCCL_WIN_REQUIRED_ALIGNMENT}-byte aligned; " + "NCCL window registration requires NCCL_WIN_REQUIRED_ALIGNMENT" + ) + size = _align_up(self.size, _win.NCCL_WIN_REQUIRED_ALIGNMENT) + if size > self.size: + raise RuntimeError( + f"arena size {self.size} is not a multiple of " + f"{_win.NCCL_WIN_REQUIRED_ALIGNMENT}; registering {size} bytes would run " + "past the allocation" + ) + # The same communicator the devcomm was created from: a window + # handle is only meaningful to the devcomm that shares its comm. + self._arena_window_size = size + self._arena_window = _win.register_window( + self._arena_window_comm, base, size, _win.NCCL_WIN_COLL_SYMMETRIC) + + # Spans nodes, so it must be gated on the global group. + self._collective_stage("register arena NCCL window", register, group=self._global_group) + def _init_multicast_buffer(self): """Create multicast object and map, following multi-process fabric pattern.""" num_devices = self._num_local_ranks @@ -384,8 +599,12 @@ def create_and_export(): mcast_fabric_bytes = self._collective_stage("create multicast object", create_and_export) def broadcast_handle(): + # The multicast object spans one node's devices, and each node's + # local rank 0 creates its own. Broadcasting over the global group + # would have ranks on different nodes pass different `src` values to + # the same collective, so this must stay node-local. obj_list = [mcast_fabric_bytes] - dist.broadcast_object_list(obj_list, src=self._group_root_global_rank, group=self._group) + dist.broadcast_object_list(obj_list, src=self._group_root_global_rank, group=self._allocator_group) return obj_list[0] mcast_fabric_bytes = self._collective_stage("broadcast multicast object", broadcast_handle) @@ -495,10 +714,22 @@ def close(self): lambda: torch.cuda.synchronize(self._device), ) self._collective_stage("release imported mappings", self._free_remote_mappings) + # Deregistration is collective over the global group and must + # precede freeing the arena the window points at. Destroying the + # devcomm and the GIN group are collective too, so this stage runs + # when any of the three exists -- gating on the window alone would + # skip resources left behind by a partially failed init. + if self._arena_window or self._dev_comm is not None or self._gin_group is not None: + self._collective_stage( + "deregister arena NCCL window", + self._free_arena_window, + group=self._global_group, + ) self._collective_stage("release owned allocations", self._free_local_allocations) else: self._set_device("allocator close") self._free_remote_mappings() + self._free_arena_window() self._free_local_allocations() self._closed = True @@ -516,6 +747,7 @@ def _free(self): """Best-effort non-collective teardown used by failed construction/destruction.""" self._set_device("allocator cleanup") self._free_remote_mappings() + self._free_arena_window() self._free_local_allocations() def _free_remote_mappings(self): @@ -545,6 +777,61 @@ def _free_remote_mappings(self): self._peer_ptr_values = [] self._buffer_ptrs = None + def _free_dev_comm(self): + """Destroy the GIN devcomm. Collective, and must precede window teardown.""" + dev_comm = getattr(self, "_dev_comm", None) + if dev_comm is None: + return + # Clear first so a failed destroy is not retried against a handle NCCL may + # already have released, and so kernels cannot read a stale pointer. + self._dev_comm = None + if self._table is not None and len(self._table) > _META_INTERNODE_DEV_COMM: + self._table[_META_GLOBAL_DEV_COMM] = 0 + self._table[_META_INTERNODE_DEV_COMM] = 0 + + from tilelang.distributed import nccl_window as _win + + _win.destroy_dev_comm(dev_comm) + + def _free_arena_window(self): + """Tear down the whole GIN stack: devcomm, then window, then its group. + + Ordered innermost-first, since each resource references the one after it. + No step short-circuits the rest: a partially failed init can leave any + subset of the three behind, and each one is skipped individually rather + than by returning early. + """ + # The devcomm references the communicator the window belongs to, so it has + # to go first regardless of whether a window was ever registered. + self._free_dev_comm() + + window = getattr(self, "_arena_window", 0) + if window: + comm_ptr = getattr(self, "_arena_window_comm", 0) + # Clear first: a failed deregistration must not be retried against a + # handle NCCL may already have destroyed. + self._arena_window = 0 + self._arena_window_comm = 0 + self._arena_window_size = 0 + if self._table is not None and len(self._table) > _META_ARENA_WINDOW: + self._table[_META_ARENA_WINDOW] = 0 + + from tilelang.distributed import nccl_window as _win + + _win.deregister_window(comm_ptr, window) + + # Last: the communicator both of the above were created from. + self._free_gin_group() + + def _free_gin_group(self): + """Destroy the private GIN process group, after its devcomm and window.""" + gin_group = getattr(self, "_gin_group", None) + if gin_group is None: + return + self._gin_group = None + # Collective, and only safe once nothing references the communicator. + dist.destroy_process_group(gin_group) + def _free_local_allocations(self): if getattr(self, "_mcast_phys_ptr", 0) and self._mcast_phys_ptr: mcast_phys_ptr = self._mcast_phys_ptr @@ -572,17 +859,27 @@ def _init_table(self): # Synchronize handles (VMM or IPC) handles = [None] * self._group_size + # With a single rank on the node there is no peer to map, so the export is + # pure overhead -- and it is not always possible: a VMM arena allocated + # with a POSIX-FD handle type (the fallback when no IMEX channel exists) + # cannot be exported as a fabric handle. Skipping keeps a GIN-only run + # working on nodes where intra-node peer mapping is unavailable. + skip_peer_handles = self._group_size == 1 + def create_handle(): if self._use_vmm: return _create_vmm_handle(self._base_ptr.value) return _create_ipc_handle(self._base_ptr.value) - local_handle = self._collective_stage("export allocation handles", create_handle) - local_handle = self._collective_stage( - "serialize allocation handle", - lambda: bytes(local_handle), - ) - dist.all_gather_object(handles, local_handle, group=self._group) + if skip_peer_handles: + handles = [b""] + else: + local_handle = self._collective_stage("export allocation handles", create_handle) + local_handle = self._collective_stage( + "serialize allocation handle", + lambda: bytes(local_handle), + ) + dist.all_gather_object(handles, local_handle, group=self._allocator_group) def allocate_peer_pointer_table(): self._buffer_ptrs = torch.empty(self._group_size, dtype=torch.uint64, device=f"cuda:{self._device}") @@ -596,6 +893,8 @@ def import_handles(): for peer_rank, handle in enumerate(handles): if peer_rank == self._local_rank: self._peer_ptr_values[peer_rank] = self._base_ptr.value + elif skip_peer_handles: + continue elif self._use_vmm: self._peer_ptr_values[peer_rank] = _open_vmm_handle(handle) else: @@ -607,11 +906,42 @@ def import_handles(): self._collective_stage("import allocation handles", import_handles) def finalize_pointer_table(): - self._table_size = 2 + self._group_size + # Layout is defined once in + # src/tl_templates/cuda/distributed/meta_layout.h and must match the + # device helpers in distributed.h and the host-side remote TMA + # remapping in src/cuda/runtime.cc. + if self._is_multi_node: + global_rank = dist.get_rank(self._global_group) + global_world_size = dist.get_world_size(self._global_group) + node_rank = self._node_info.node_rank + num_nodes = self._node_info.num_nodes + else: + global_rank = self._local_rank + global_world_size = self._num_local_ranks + node_rank = 0 + num_nodes = 1 + local_rank = self._local_rank + local_world_size = self._num_local_ranks + + self._table_size = _META_PEER_BASE + local_world_size self._table = torch.empty(self._table_size, dtype=torch.uint64, device="cpu") - self._table[0] = self._local_rank - self._table[1] = self._group_size - self._table[2:] = self._buffer_ptrs + self._table[_META_GLOBAL_RANK] = global_rank + self._table[_META_GLOBAL_WORLD_SIZE] = global_world_size + self._table[_META_NODE_RANK] = node_rank + self._table[_META_NUM_NODES] = num_nodes + self._table[_META_LOCAL_RANK] = local_rank + self._table[_META_LOCAL_WORLD_SIZE] = local_world_size + # Device pointer to the GIN devcomm, or zero when GIN is unavailable; + # tl::gin::available() tests exactly this slot. The devcomm covers the + # whole global communicator, so the same handle serves both slots -- + # the inter-node entry is kept distinct for a future rail-local comm. + dev_comm_ptr = self._dev_comm.device_ptr if self._dev_comm is not None else 0 + self._table[_META_GLOBAL_DEV_COMM] = dev_comm_ptr + self._table[_META_INTERNODE_DEV_COMM] = dev_comm_ptr + # Zero window means no GIN; kernels must check before using it. + self._table[_META_ARENA_WINDOW] = self._arena_window + self._table[_META_ARENA_BASE] = self._base_ptr.value or 0 + self._table[_META_PEER_BASE:] = self._buffer_ptrs self._collective_stage("finalize peer pointer table", finalize_pointer_table) @@ -697,6 +1027,33 @@ def table(self) -> torch.Tensor: def table_size(self) -> int: return self._table_size + @property + def arena_window(self) -> int: + """``ncclWindow_t`` for the whole arena, or 0 when GIN is unavailable.""" + return self._arena_window + + @property + def arena_base(self) -> int: + """Arena base address; subtract from a local pointer to get a window offset.""" + return int(self._base_ptr.value) if self._base_ptr and self._base_ptr.value else 0 + + def window_offset(self, ptr: int) -> int: + """Convert a local arena pointer to the offset valid on every rank. + + The arena is symmetric, so the same offset names the corresponding bytes + on any peer -- the identical subtraction the intra-node peer-pointer path + performs, just paired with a window handle instead of a peer base. + """ + base = self.arena_base + if not base: + raise RuntimeError("allocator has no arena base; cannot compute a window offset") + ptr = operator.index(ptr) + if not base <= ptr < base + self.size: + raise ValueError( + f"pointer {ptr:#x} is outside the arena [{base:#x}, {base + self.size:#x})" + ) + return ptr - base + def __enter__(self): return self @@ -734,6 +1091,7 @@ def get_allocator( group: dist.ProcessGroup | None = None, use_vmm: bool | None = None, mcast_size: int | None = None, + node_info: NodeTopology | None = None, ) -> BaseAllocator: return BaseAllocator( size, @@ -744,4 +1102,5 @@ def get_allocator( group=group, use_vmm=use_vmm, mcast_size=mcast_size, + node_info=node_info, ) diff --git a/tilelang/distributed/host.py b/tilelang/distributed/host.py index 2a9f393624..d93456565b 100644 --- a/tilelang/distributed/host.py +++ b/tilelang/distributed/host.py @@ -3,6 +3,7 @@ import inspect import os import operator +from dataclasses import dataclass from functools import lru_cache import torch @@ -21,6 +22,22 @@ ) from exc +@dataclass +class NodeTopology: + """Multi-node topology information for distributed execution. + + Attributes: + node_rank: Rank of this node (0 to num_nodes-1) + num_nodes: Total number of nodes in the distributed job + local_world_size: Number of GPUs per node + node_local_group: NCCL process group containing only ranks on this node + """ + node_rank: int + num_nodes: int + local_world_size: int + node_local_group: dist.ProcessGroup + + def CUDA_CHECK(err): if isinstance(err, cuda.CUresult): if err != cuda.CUresult.CUDA_SUCCESS: @@ -32,11 +49,31 @@ def CUDA_CHECK(err): raise RuntimeError(f"Unknown error type: {err}") -def init_dist(local_rank: int, num_local_ranks: int, master_port: int | None = None): - """Initialize the currently supported single-node NCCL process group. +def init_dist( + local_rank: int, + num_local_ranks: int, + master_port: int | None = None, + return_node_info: bool = False, +): + """Initialize an NCCL process group with single-node or multi-node support. + + Args: + local_rank: Local rank on this node (0 to num_local_ranks-1) + num_local_ranks: Number of ranks per node + master_port: Optional master port (defaults to TILESCALE_MASTER_PORT or MASTER_PORT) + return_node_info: When True, also return the :class:`NodeTopology` describing + the multi-node layout. Defaults to False so that existing single-node + callers keep the historical three-value return. - Single-node ``torchrun`` variables are accepted. Multi-node groups are - rejected until TileScale creates a node-local allocator group. + Returns: + ``(rank, world_size, global_group)`` by default, or + ``(rank, world_size, global_group, node_info)`` when ``return_node_info`` + is True. ``node_info`` is None for a single-node launch. + + Note: + A multi-node launch requires ``return_node_info=True``; the resulting + ``node_info`` must be forwarded to :func:`tilelang.get_allocator` so the + allocator restricts IPC/VMM handle exchange to node-local peers. """ os.environ.setdefault("NCCL_IB_DISABLE", "1") os.environ.setdefault("NCCL_DEBUG", "ERROR") @@ -44,7 +81,9 @@ def init_dist(local_rank: int, num_local_ranks: int, master_port: int | None = N if not 0 <= local_rank < num_local_ranks: raise ValueError(f"local_rank must be in [0, {num_local_ranks}), got {local_rank}") + # Detect topology from environment if "LOCAL_WORLD_SIZE" in os.environ: + # torchrun style variables launcher_local_world_size = int(os.environ["LOCAL_WORLD_SIZE"]) launcher_world_size = int(os.environ.get("WORLD_SIZE", launcher_local_world_size)) launcher_rank = int(os.environ.get("RANK", local_rank)) @@ -60,24 +99,26 @@ def init_dist(local_rank: int, num_local_ranks: int, master_port: int | None = N "local_rank must match torchrun RANK and LOCAL_RANK for a single-node launch: " f"local_rank={local_rank}, RANK={launcher_rank}, LOCAL_RANK={launcher_local_rank}" ) + global_rank = launcher_rank + global_world_size = launcher_world_size else: + # Manual environment variables (NNODES, NODE_RANK) num_nodes = int(os.environ.get("NNODES", "1")) node_rank = int(os.environ.get("NODE_RANK", "0")) - legacy_world_size = int(os.environ.get("WORLD_SIZE", "1")) - legacy_rank = int(os.environ.get("RANK", "0")) - if legacy_world_size != 1 or legacy_rank != 0: - raise NotImplementedError( - "Ambiguous WORLD_SIZE/RANK launcher environment without LOCAL_WORLD_SIZE; " - "TileScale supports local spawn or single-node torchrun only." - ) + global_rank = int(os.environ.get("RANK", local_rank)) + global_world_size = int(os.environ.get("WORLD_SIZE", num_local_ranks)) - if num_nodes != 1 or node_rank != 0: - raise NotImplementedError( - "TileScale currently supports only single-node process groups; multi-node launch requires a node-local allocator group." - ) + # Validate consistency + if global_world_size != num_nodes * num_local_ranks: + raise ValueError( + f"Inconsistent configuration: WORLD_SIZE ({global_world_size}) != " + f"NNODES ({num_nodes}) * num_local_ranks ({num_local_ranks})" + ) + # Set device torch.cuda.set_device(local_rank) + # Initialize global process group ip = os.getenv("MASTER_ADDR", "127.0.0.1") port = master_port if master_port is not None else int(os.getenv("TILESCALE_MASTER_PORT", os.getenv("MASTER_PORT", "8361"))) @@ -85,14 +126,51 @@ def init_dist(local_rank: int, num_local_ranks: int, master_port: int | None = N params = { "backend": "nccl", "init_method": f"tcp://{ip}:{port}", - "world_size": num_local_ranks, - "rank": local_rank, + "world_size": global_world_size, + "rank": global_rank, } if "device_id" in sig.parameters: params["device_id"] = torch.device(f"cuda:{local_rank}") + # Opt-in shorter group timeout. NCCL's 10-minute default means a mismatched + # or stuck setup collective busy-waits at 100% GPU utilisation and gets + # killed by the launcher's outer timeout before the watchdog ever reports, + # which leaves no diagnostic at all. Unset -> unchanged default behaviour. + _pg_timeout = os.getenv("TL_PG_TIMEOUT_SEC") + if _pg_timeout and "timeout" in sig.parameters: + import datetime + + params["timeout"] = datetime.timedelta(seconds=int(_pg_timeout)) dist.init_process_group(**params) - return dist.get_rank(), dist.get_world_size(), dist.group.WORLD + if num_nodes == 1: + if return_node_info: + return dist.get_rank(), dist.get_world_size(), dist.group.WORLD, None + return dist.get_rank(), dist.get_world_size(), dist.group.WORLD + + if not return_node_info: + raise ValueError( + f"a multi-node launch was detected (num_nodes={num_nodes}), which requires " + "init_dist(..., return_node_info=True) so the node-local topology can be " + "forwarded to the allocator via get_allocator(node_info=...)" + ) + + # Every rank must build the same subgroups in the same order: dist.new_group + # is collective over the default group, so ranks cannot create only their own. + node_local_group = None + for node in range(num_nodes): + group = dist.new_group( + ranks=[node * num_local_ranks + i for i in range(num_local_ranks)]) + if node == node_rank: + node_local_group = group + + node_info = NodeTopology( + node_rank=node_rank, + num_nodes=num_nodes, + local_world_size=num_local_ranks, + node_local_group=node_local_group, + ) + + return global_rank, global_world_size, dist.group.WORLD, node_info @lru_cache diff --git a/tilelang/distributed/nccl_window.py b/tilelang/distributed/nccl_window.py new file mode 100644 index 0000000000..35d1e18864 --- /dev/null +++ b/tilelang/distributed/nccl_window.py @@ -0,0 +1,395 @@ +"""NCCL window registration for GIN (GPU-Initiated Networking). + +GIN addresses remote memory as an ``(ncclWindow_t, byte offset)`` pair rather +than a raw pointer, because a remote rank's allocation has no local virtual +address. ``ncclCommWindowRegister`` is collective and returns one local handle +that is *not* per-peer: that handle plus a peer index names any rank's bytes. + +TileScale therefore registers the whole allocator arena once, at allocator init, +instead of per tensor. Registration is collective and 4096-byte aligned, so a +per-tensor scheme would turn every ``allocate_tensor`` into a world-wide barrier. +Since the arena is symmetric, ``local_ptr - arena_base`` yields the offset valid +on every rank -- the same subtraction the intra-node path already performs. + +Requires the NCCL Device API (>= 2.28.7). Everything here degrades to "no +window" when that is unavailable, so single-node paths keep working untouched. +""" + +from __future__ import annotations + +import ctypes +import os + +import torch +import torch.distributed as dist + +__all__ = [ + "NCCL_WIN_DEFAULT", + "NCCL_WIN_COLL_SYMMETRIC", + "NCCL_WIN_REQUIRED_ALIGNMENT", + "DEV_COMM_STORAGE_BYTES", + "DevComm", + "GIN_CONTEXT_COUNT", + "GIN_SIGNAL_COUNT", + "GIN_COUNTER_COUNT", + "nccl_version", + "supports_device_api", + "unavailable_reason", + "get_comm_ptr", + "register_window", + "deregister_window", + "create_dev_comm", + "destroy_dev_comm", +] + +NCCL_WIN_DEFAULT = 0x00 +NCCL_WIN_COLL_SYMMETRIC = 0x01 +# NCCL_WIN_REQUIRED_ALIGNMENT from nccl.h; both base and size must respect it. +NCCL_WIN_REQUIRED_ALIGNMENT = 4096 + +# The Device API (nccl_device.h, ncclDevCommCreate, GIN) first ships in 2.28.7. +# Encoded as NCCL_VERSION(x, y, z) = x * 10000 + y * 100 + z. +_MIN_DEVICE_API_VERSION = 2 * 10000 + 28 * 100 + 7 + +# sizeof(ncclDevComm), measured from the 2.28.9 headers on node071. The struct is +# opaque here: kernels only ever receive a pointer to it, so its size and 8-byte +# alignment are all the host side needs. Over-allocating a page is deliberate -- +# a later NCCL may grow the struct, and a short buffer would be a silent +# out-of-bounds write inside ncclDevCommCreate rather than an error. +_DEV_COMM_SIZEOF = 200 +DEV_COMM_STORAGE_BYTES = 4096 + +# GIN resource counts requested when creating the devcomm. +# +# Contexts are independent network channels; kernels rotate over them +# (blockIdx.x % count) so concurrent CTAs do not serialize on one channel. +# Signals and counters are guaranteed to start at id 0, so a kernel can address +# signal i directly for i < GIN_SIGNAL_COUNT. +GIN_CONTEXT_COUNT = 8 +GIN_SIGNAL_COUNT = 32 +GIN_COUNTER_COUNT = 32 + + +class _DevCommRequirements(ctypes.Structure): + """Mirror of ``struct ncclDevCommRequirements`` from nccl_device/core.h. + + Field order and types are taken from the 2.28.9 header; the layout was + verified against ``offsetof``/``sizeof`` on node071 (56 bytes, 8-byte + aligned, ``bool`` members padded to the following 4-byte field). ctypes + reproduces that natural layout, so no explicit padding is declared. + """ + + _fields_ = [ + ("resourceRequirementsList", ctypes.c_void_p), + ("teamRequirementsList", ctypes.c_void_p), + ("lsaMultimem", ctypes.c_bool), + ("barrierCount", ctypes.c_int), + ("lsaBarrierCount", ctypes.c_int), + ("railGinBarrierCount", ctypes.c_int), + ("lsaLLA2ABlockCount", ctypes.c_int), + ("lsaLLA2ASlotCount", ctypes.c_int), + ("ginForceEnable", ctypes.c_bool), + ("ginContextCount", ctypes.c_int), + ("ginSignalCount", ctypes.c_int), + ("ginCounterCount", ctypes.c_int), + ] + + +# Guard against a silently different layout: ctypes computing a size other than +# the measured 56 would mean requirements land in the wrong fields, which NCCL +# would read as a garbage resource request rather than reject. +assert ctypes.sizeof(_DevCommRequirements) == 56, ( + f"ncclDevCommRequirements mirror is {ctypes.sizeof(_DevCommRequirements)} bytes, expected 56") + +_lib = None +_lib_error: str | None = None + + +def _load_libnccl(): + """Return the libnccl already mapped into this process, or None. + + dlopen by soname returns the existing mapping, so this resolves to the same + copy torch links against rather than loading a second one. A second copy + would hand back handles from a different NCCL state than the communicator + the process group owns. + """ + global _lib, _lib_error + if _lib is not None or _lib_error is not None: + return _lib + + candidates = [] + override = os.environ.get("TILESCALE_NCCL_LIB") + if override: + candidates.append(override) + candidates += ["libnccl.so.2", "libnccl.so"] + + errors = [] + for name in candidates: + try: + lib = ctypes.CDLL(name) + except OSError as exc: + errors.append(f"{name}: {exc}") + continue + _lib = lib + return _lib + + _lib_error = "; ".join(errors) + return None + + +def nccl_version() -> int | None: + """Return the integer NCCL version, or None when libnccl is unavailable.""" + lib = _load_libnccl() + if lib is None: + return None + try: + fn = lib.ncclGetVersion + except AttributeError: + return None + fn.restype = ctypes.c_int + fn.argtypes = [ctypes.POINTER(ctypes.c_int)] + version = ctypes.c_int(0) + if fn(ctypes.byref(version)) != 0: + return None + return int(version.value) + + +def supports_device_api() -> bool: + """True when the loaded NCCL exposes the Device API used by GIN.""" + lib = _load_libnccl() + if lib is None: + return False + version = nccl_version() + if version is None or version < _MIN_DEVICE_API_VERSION: + return False + # Version alone is not proof: some builds omit the device symbols. + return all(hasattr(lib, sym) for sym in ("ncclCommWindowRegister", "ncclDevCommCreate")) + + +def unavailable_reason() -> str: + """Human-readable explanation for why windows cannot be registered.""" + lib = _load_libnccl() + if lib is None: + return f"libnccl could not be loaded ({_lib_error})" + version = nccl_version() + if version is None: + return "ncclGetVersion failed; cannot confirm Device API support" + if version < _MIN_DEVICE_API_VERSION: + return ( + f"NCCL {version // 10000}.{version // 100 % 100}.{version % 100} predates the " + "Device API; GIN needs >= 2.28.7 (set TILESCALE_NCCL_LIB to a newer libnccl)" + ) + missing = [s for s in ("ncclCommWindowRegister", "ncclDevCommCreate") if not hasattr(lib, s)] + if missing: + return f"libnccl is missing Device API symbols: {', '.join(missing)}" + return "NCCL Device API is available" + + +def get_comm_ptr(group: dist.ProcessGroup) -> int: + """Return the raw ``ncclComm_t`` backing ``group``, or 0 if unobtainable. + + NCCL communicators are created lazily, so the pointer only exists after the + group has run at least one collective on the current device. The caller is + expected to have done so (allocator init is collective throughout). + """ + try: + backend = group._get_backend(torch.device("cuda", torch.cuda.current_device())) + except Exception: # noqa: BLE001 - non-NCCL backend or unsupported torch + return 0 + + comm_ptr = getattr(backend, "_comm_ptr", None) + if comm_ptr is None: + return 0 + try: + value = comm_ptr() + except Exception: # noqa: BLE001 - comm not yet initialized + return 0 + return int(value) if value else 0 + + +def _check(lib, rc: int, what: str) -> None: + if rc == 0: + return + detail = "" + try: + lib.ncclGetErrorString.restype = ctypes.c_char_p + lib.ncclGetErrorString.argtypes = [ctypes.c_int] + msg = lib.ncclGetErrorString(rc) + if msg: + detail = f" ({msg.decode()})" + except Exception: # noqa: BLE001 - error string is best effort + pass + raise RuntimeError(f"{what} failed: rc={rc}{detail}") + + +def register_window(comm_ptr: int, base_ptr: int, size: int, flags: int = NCCL_WIN_COLL_SYMMETRIC) -> int: + """Register ``[base_ptr, base_ptr + size)`` as an NCCL window. + + Collective over the communicator: every rank must call this at the same + point in its init sequence with a matching size. ``NCCL_WIN_COLL_SYMMETRIC`` + asserts matching layouts across ranks, and a mismatch hangs rather than + erroring, which is why the arena registration lives in a single place. + + Returns the ``ncclWindow_t`` as an integer handle. + """ + lib = _load_libnccl() + if lib is None: + raise RuntimeError(f"cannot register an NCCL window: {unavailable_reason()}") + if not comm_ptr: + raise ValueError("comm_ptr must be a non-null ncclComm_t") + if base_ptr % NCCL_WIN_REQUIRED_ALIGNMENT: + raise ValueError( + f"window base {base_ptr:#x} is not {NCCL_WIN_REQUIRED_ALIGNMENT}-byte aligned " + "as NCCL_WIN_REQUIRED_ALIGNMENT demands" + ) + if size <= 0 or size % NCCL_WIN_REQUIRED_ALIGNMENT: + raise ValueError(f"window size {size} must be positive and a multiple of {NCCL_WIN_REQUIRED_ALIGNMENT}") + + fn = lib.ncclCommWindowRegister + fn.restype = ctypes.c_int + fn.argtypes = [ + ctypes.c_void_p, + ctypes.c_void_p, + ctypes.c_size_t, + ctypes.POINTER(ctypes.c_void_p), + ctypes.c_int, + ] + win = ctypes.c_void_p(0) + _check( + lib, + fn(ctypes.c_void_p(comm_ptr), ctypes.c_void_p(base_ptr), ctypes.c_size_t(size), ctypes.byref(win), ctypes.c_int(flags)), + "ncclCommWindowRegister", + ) + if not win.value: + raise RuntimeError("ncclCommWindowRegister returned success but a null window handle") + return int(win.value) + + +class DevComm: + """A live GIN devcomm: the host struct NCCL owns, plus its device copy. + + Both halves must be kept: ``ncclDevCommDestroy`` takes the *host* pointer it + was created with, while kernels dereference the *device* pointer published in + the metadata table. ``device_ptr`` is only valid while this object is alive. + """ + + __slots__ = ( + "_host_buf", + "_storage", + "device_ptr", + "_comm_ptr", + "_destroyed", + "context_count", + ) + + def __init__(self, comm_ptr: int, host_buf, storage: "torch.Tensor", context_count: int): + self._comm_ptr = comm_ptr + self._host_buf = host_buf + self._storage = storage + self.device_ptr = int(storage.data_ptr()) + # Carried so callers publish the count this devcomm was actually built + # with. Reading the module default instead would silently lie to kernels + # whenever a caller overrides context_count. + self.context_count = context_count + self._destroyed = False + + def destroy(self) -> None: + """Release the devcomm. Collective; idempotent.""" + if self._destroyed: + return + self._destroyed = True + lib = _load_libnccl() + if lib is None or not self._comm_ptr: + return + fn = lib.ncclDevCommDestroy + fn.restype = ctypes.c_int + fn.argtypes = [ctypes.c_void_p, ctypes.c_void_p] + _check( + lib, + fn(ctypes.c_void_p(self._comm_ptr), ctypes.byref(self._host_buf)), + "ncclDevCommDestroy", + ) + + +def create_dev_comm( + comm_ptr: int, + *, + context_count: int = GIN_CONTEXT_COUNT, + signal_count: int = GIN_SIGNAL_COUNT, + counter_count: int = GIN_COUNTER_COUNT, + rail_gin_barrier_count: int = 1, +) -> DevComm: + """Create a GIN-enabled ``ncclDevComm``. + + Collective over the communicator: every rank must call this with matching + requirements, and a mismatch hangs rather than erroring. + + ``ncclDevCommCreate`` writes the devcomm into caller-provided **host** + storage. Measured on node071 with NCCL 2.28.9: passing a ``cudaMalloc`` + pointer segfaults inside the call, while a host struct succeeds. NCCL's own + examples then pass the result to kernels by value in a ``__grid_constant__`` + parameter, so the struct is trivially copyable. + + TileScale instead publishes a pointer through the metadata table, so kernel + signatures stay unchanged. That pointer is dereferenced on the device, so the + host struct is created first and then copied into a device tensor. + + ``ginForceEnable`` is set because GIN is otherwise enabled only when NCCL + decides the topology warrants it; TileScale needs it deterministically, since + a kernel compiled for the GIN path cannot fall back at runtime. + + The returned :class:`DevComm` owns both buffers and must outlive every kernel + that uses it: dropping it frees the memory its device pointer names. + """ + lib = _load_libnccl() + if lib is None: + raise RuntimeError(f"cannot create an NCCL devcomm: {unavailable_reason()}") + if not comm_ptr: + raise ValueError("comm_ptr must be a non-null ncclComm_t") + if signal_count <= 0 or counter_count <= 0 or context_count <= 0: + raise ValueError("GIN context, signal, and counter counts must all be positive") + + reqs = _DevCommRequirements() + ctypes.memset(ctypes.byref(reqs), 0, ctypes.sizeof(reqs)) + reqs.ginForceEnable = True + reqs.ginContextCount = context_count + reqs.ginSignalCount = signal_count + reqs.ginCounterCount = counter_count + reqs.railGinBarrierCount = rail_gin_barrier_count + + host_buf = (ctypes.c_uint8 * DEV_COMM_STORAGE_BYTES)() + + fn = lib.ncclDevCommCreate + fn.restype = ctypes.c_int + fn.argtypes = [ctypes.c_void_p, ctypes.POINTER(_DevCommRequirements), ctypes.c_void_p] + _check( + lib, + fn(ctypes.c_void_p(comm_ptr), ctypes.byref(reqs), ctypes.byref(host_buf)), + "ncclDevCommCreate", + ) + + # uint8 so numel() == bytes. Only the struct itself is copied; the rest of the + # padding stays zero. + storage = torch.zeros(DEV_COMM_STORAGE_BYTES, dtype=torch.uint8, device="cuda") + staged = torch.frombuffer( + memoryview(host_buf)[:_DEV_COMM_SIZEOF], dtype=torch.uint8).clone() + storage[:_DEV_COMM_SIZEOF].copy_(staged) + torch.cuda.synchronize() + return DevComm(comm_ptr, host_buf, storage, context_count) + + +def destroy_dev_comm(dev_comm: DevComm | None) -> None: + """Release a devcomm. Collective, and must precede freeing its storage.""" + if dev_comm is not None: + dev_comm.destroy() + + +def deregister_window(comm_ptr: int, window: int) -> None: + """Release a window handle. Collective, and must precede freeing the arena.""" + lib = _load_libnccl() + if lib is None or not window or not comm_ptr: + return + fn = lib.ncclCommWindowDeregister + fn.restype = ctypes.c_int + fn.argtypes = [ctypes.c_void_p, ctypes.c_void_p] + _check(lib, fn(ctypes.c_void_p(comm_ptr), ctypes.c_void_p(window)), "ncclCommWindowDeregister") diff --git a/tilelang/distributed/shared_memory/__init__.py b/tilelang/distributed/shared_memory/__init__.py index b9614fc1bd..7c1779feec 100644 --- a/tilelang/distributed/shared_memory/__init__.py +++ b/tilelang/distributed/shared_memory/__init__.py @@ -51,6 +51,7 @@ def _get_capability_global_func(name): _close_ipc_handle = _get_required_global_func("tl.shared_memory.close_ipc_handle") _sync_ipc_handles_raw = _get_required_global_func("tl.shared_memory.sync_ipc_handles") +_supports_vmm = _get_capability_global_func("tl.shared_memory.supports_vmm") _supports_vmm_fabric = _get_capability_global_func("tl.shared_memory.supports_vmm_fabric") _supports_multicast = _get_capability_global_func("tl.shared_memory.supports_multicast") @@ -285,6 +286,7 @@ def create_host_device_tensor(shape, dtype): "_close_ipc_handle", "_sync_ipc_handles", "create_host_device_tensor", + "_supports_vmm", "_supports_vmm_fabric", "_vmm_malloc", "_vmm_free", diff --git a/tilelang/engine/lower.py b/tilelang/engine/lower.py index 21e80b475b..e41d446053 100644 --- a/tilelang/engine/lower.py +++ b/tilelang/engine/lower.py @@ -119,6 +119,15 @@ def tilelang_callback_cuda_compile(code, target, pass_config=None): "-I" + TILELANG_TEMPLATE_PATH, "-I" + CUTLASS_INCLUDE_DIR, ] + # Inter-node (GIN) device path. codegen emits the nccl_gin.h include guarded + # by TL_ENABLE_NCCL_GIN, so the define and the include dir have to travel + # together -- with only one of them the header's body vanishes and a kernel + # calling tl::gin:: fails with "must be a class or namespace name". + from tilelang.env import NCCL_INCLUDE_DIR + + if NCCL_INCLUDE_DIR: + options += ["-I" + NCCL_INCLUDE_DIR, "-DTL_ENABLE_NCCL_GIN=1"] + # Merge extra device compiler flags from pass config, if provided extra_flags = cfg.get(PassConfigKey.TL_DEVICE_COMPILE_FLAGS, None) if extra_flags: From 58e557e373cdedbd63cdb82a348fb3a4dfdc1e8d Mon Sep 17 00:00:00 2001 From: Rachmanino <18805904201@163.com> Date: Mon, 3 Aug 2026 23:13:16 +0800 Subject: [PATCH 05/30] [Example] Add inter-node allgather, allreduce and reduce-scatter over GIN Three collectives built on `T.nccl_gin`, sharing `internode_common.py` and launched by `run_internode.sh`. All three verified on two physical nodes on 2026-08-03 against their torch references (`all_gather_into_tensor`, `all_reduce`, `reduce_scatter_tensor`) with `LOCAL_WORLD_SIZE=1 NNODES=2`, so every put crossed the RoCE fabric. At 64 MB shards, bf16, on two idle nodes they all beat torch NCCL: tilescale torch ratio allgather 47.6 39.8 1.20x allreduce 47.2 45.1 1.05x reduce_scatter 46.2 42.8 1.08x 46-48 GB/s is 93-95% of one 400 Gbps NIC's 50 GB/s, i.e. the collectives are at line rate, which is why little else moves them. The shape that gets there, modelled on Triton-distributed's inter-node sender: size the grid by peers and channels, not by payload, so one CTA issues one large put. The first version launched a CTA per 8192-element block, turning an 8 MB shard into 512 separate 16 KB puts -- latency-bound, since RDMA is only bandwidth-bound once messages are large. That plus a `T.Parallel` local copy and a per-launch signal target is ~25x the first implementation. `sweep_internode.sh` and the `--tune` mode exist because the interesting knob is not stable across conditions. `--chunks` matters for the reduce variants, whose CTAs also carry the reduction (allreduce 41.5 -> 47.2 GB/s from 4 to 64) and is flat for allgather. `--gin-contexts` is insurance rather than a win: on idle NICs one context already reaches line rate, but on a NIC shared with another tenant one context collapsed to 23 GB/s while 2-4 held ~44. Tuning on contended hardware therefore misattributes the win, so the tuner verifies every config and re-times torch before and after each sweep. `--algo oneshot` for allreduce is kept behind a flag: at W=2 it sends the same bytes as two-shot in half the phases and still loses (30.3 vs 43.4 GB/s), because it reduces over the full buffer instead of a shard. The balance shifts with rank count and dtype. --- .../internode/example_internode_allgather.py | 229 ++++++++++ .../internode/example_internode_allreduce.py | 310 ++++++++++++++ .../example_internode_reduce_scatter.py | 212 ++++++++++ .../distributed/internode/internode_common.py | 398 ++++++++++++++++++ .../distributed/internode/run_internode.sh | 122 ++++++ .../distributed/internode/sweep_internode.sh | 93 ++++ 6 files changed, 1364 insertions(+) create mode 100644 examples/distributed/internode/example_internode_allgather.py create mode 100644 examples/distributed/internode/example_internode_allreduce.py create mode 100644 examples/distributed/internode/example_internode_reduce_scatter.py create mode 100644 examples/distributed/internode/internode_common.py create mode 100755 examples/distributed/internode/run_internode.sh create mode 100755 examples/distributed/internode/sweep_internode.sh diff --git a/examples/distributed/internode/example_internode_allgather.py b/examples/distributed/internode/example_internode_allgather.py new file mode 100644 index 0000000000..218102a65e --- /dev/null +++ b/examples/distributed/internode/example_internode_allgather.py @@ -0,0 +1,229 @@ +"""Inter-node allgather over NCCL GIN. + +Each rank holds one shard and ends up with every rank's shard concatenated in +rank order. A rank pushes its own shard directly into slot ``rank`` of every +other rank's output, then waits for the slots it does not own. + +Shape of the work, and why +-------------------------- +One CTA issues one large put. The grid is ``(world_size - 1) * chunks``, so it is +sized by *peers and channels*, not by payload -- 8 CTAs for a 2-node run, not +512. This mirrors Triton-distributed's inter-node sender, which launches +``grid = (n_nodes - 1,)`` with ``num_warps=32`` and hands each block a single +``putmem_signal_block`` covering a whole shard. + +The earlier version here inverted that: it launched one CTA per 8192-element +block, so an 8 MB shard became 512 separate 16 KB puts with 512 signal +increments. RDMA at 400 Gbps is bandwidth-bound only once messages are large; +512 small messages on one queue pair is latency-bound, which is what held this +kernel to ~2 GB/s against torch NCCL's ~30 GB/s. + +``chunks`` is not a tiling parameter. One GIN context is one QP per peer, hence +one NIC; splitting a peer's transfer into ``chunks`` pieces lets ``make_gin()`` +place each piece on a different context and use several NICs at once. Splitting +beyond the granted context count only shrinks messages for no extra +parallelism, so ``--chunks`` and ``--gin-contexts`` want to match. + +Signals are cumulative +---------------------- +GIN signals are running totals that a wait does not consume, so ``least`` must be +a per-launch target rather than a constant. With a constant, the second launch +finds the counter already at the target and returns *without waiting for any +data* -- correct on the first call only, and it silently turns a benchmark into a +measurement of nothing. ``signal_target`` is therefore passed in and advanced by +the host, the same way Triton-distributed threads its ``signal_target`` through +``NVSHMEM_SIGNAL_SET``. + +Launch: see run_internode.sh. +""" + +# NOTE: no `from __future__ import annotations` here. T.prim_func resolves the +# parameter annotations at runtime via get_type_hints, and PEP 563 would turn +# `T.Tensor((shard_numel,), dtype)` into a string evaluated against module +# globals -- where the closure locals shard_numel/dtype do not exist. +import argparse + +import torch +import torch.distributed as dist + +import tilelang.language as T +from tilelang.distributed.bench import do_bench + +from internode_common import ( + SIGNAL_DATA, + per_launch_signals, + report_tuning, + tune_grid, + Context, + TL_DTYPES, + TORCH_DTYPES, + add_common_args, + check, + prepare_env, + report, +) + + +def allgather_kernel(shard_numel: int, chunks: int, threads: int, world_size: int, dtype: str, + signal_id: int = SIGNAL_DATA): + """One CTA per (peer, chunk); each issues a single large put. + + ``peer`` is selected by block index rather than looped inside the CTA so + every put is independent and can sit on its own context. + """ + peers = world_size - 1 + chunk_numel = shard_numel // chunks + blocks = peers * chunks + # Each block also copies one slice of the local shard into place, so the + # slices tile the shard exactly once across the grid. + copy_numel = shard_numel // blocks + + @T.prim_func + def main( + shard: T.Tensor((shard_numel,), dtype), + out: T.Tensor((world_size * shard_numel,), dtype), + rank: T.int32, + signal_target: T.int32, + ): + with T.Kernel(blocks, threads=threads) as bx: + peer_idx = bx // chunks + chunk_idx = bx % chunks + # Rotating the peer by rank keeps concurrent senders from all + # hammering the same destination first. + peer = (rank + peer_idx + 1) % world_size + # Destination is slot `rank` of the peer's output. The arena is + # symmetric, so the index written is the index the peer reads. + T.nccl_gin.put_signal( + src=shard[chunk_idx * chunk_numel], + dst=out[rank * shard_numel + chunk_idx * chunk_numel], + size=chunk_numel, + peer=peer, + signal_id=signal_id, + scope="block", + ) + # The local shard never crosses the network; copy it while the puts + # are in flight. T.Parallel spreads this across the CTA's threads -- + # a T.serial loop here would have every thread redundantly walk the + # whole slice. + for i in T.Parallel(copy_numel): + out[rank * shard_numel + bx * copy_numel + i] = shard[bx * copy_numel + i] + T.nccl_gin.wait_signal(least=signal_target, signal_id=signal_id, scope="block") + + return main + + +def main() -> int: + parser = add_common_args(argparse.ArgumentParser(description=__doc__)) + args = parser.parse_args() + + prepare_env() + ctx = Context() + + # --numel is the gathered size, so it must split evenly into shards. + if args.numel % ctx.world_size: + raise SystemExit(f"--numel {args.numel} must be divisible by world_size {ctx.world_size}") + shard_numel = args.numel // ctx.world_size + peers = ctx.world_size - 1 + # Only validate the single requested config here. Under --tune each config is + # checked in the loop and unusable ones are skipped rather than fatal. + if not args.tune: + blocks = peers * args.chunks + if shard_numel % args.chunks: + raise SystemExit(f"shard {shard_numel} must be a multiple of --chunks {args.chunks}") + if shard_numel % blocks: + raise SystemExit( + f"shard {shard_numel} must be a multiple of (world_size-1)*chunks = {blocks} " + "so the local copy tiles evenly" + ) + + torch_dtype = TORCH_DTYPES[args.dtype] + ctx.log( + f"allgather: world_size={ctx.world_size} nodes={ctx.num_nodes} " + f"numel={args.numel} shard={shard_numel} dtype={args.dtype} " + + ( + f"tune chunks={args.tune_chunks} contexts={args.tune_contexts} " + f"threads={args.tune_threads}" + if args.tune + else f"chunks={args.chunks} contexts={args.gin_contexts} threads={args.threads}" + ) + ) + + shard = ctx.tensor((shard_numel,), torch_dtype) + out = ctx.tensor((args.numel,), torch_dtype) + # Rank-dependent values so a slot filled by the wrong peer, or not at all, + # cannot compare equal by accident. + shard.copy_( + torch.arange(shard_numel, device=shard.device, dtype=torch.float32).to(torch_dtype) + + ctx.rank * 1000.0 + ) + out.zero_() + + # Golden result, once. Each rank sends its shard to world_size-1 peers; that + # egress is what the link has to carry. + ref = torch.empty_like(out) + dist.all_gather_into_tensor(ref, shard, group=ctx.group) + moved = shard.numel() * shard.element_size() * peers + ref_buf = torch.empty_like(out) + + def time_torch() -> float: + return do_bench( + lambda: dist.all_gather_into_tensor(ref_buf, shard, group=ctx.group), + warmup=args.warmup, rep=args.rep, group=ctx.group, + ) + + configs = tune_grid(args, signals_per_config=1) + torch_before = time_torch() if not args.no_bench else float("nan") + + rows = [] + failures = 0 + for cfg in configs: + blocks = peers * cfg["chunks"] + if shard_numel % cfg["chunks"] or shard_numel % blocks: + continue + kernel = ctx.compile( + allgather_kernel( + shard_numel, cfg["chunks"], cfg["threads"], ctx.world_size, + TL_DTYPES[args.dtype], signal_id=cfg["signals"][0], + ), + expect=("tl::gin::put_signal_addr", "tl::gin::wait_signal"), + gin_contexts=cfg["gin_contexts"], + wait_ctx0=args.wait_ctx0, + ) + # This config's slot starts at zero, so its own running total is all that + # matters -- see tune_grid for why the slots are not shared. + per_launch = per_launch_signals(peers, cfg["chunks"], args.signal_div) + target = [0] + + def launch(kernel=kernel, target=target, per_launch=per_launch): + target[0] += per_launch + kernel(shard, out, ctx.rank, target[0]) + + out.zero_() + torch.cuda.synchronize() + dist.barrier(ctx.group) + launch() + torch.cuda.synchronize() + bad = check(ctx, out, ref, f"allgather[c{cfg['chunks']}/x{cfg['gin_contexts']}]") + failures += bad + + ms = float("inf") + if bad == 0 and not args.no_bench: + dist.barrier(ctx.group) + ms = do_bench(launch, warmup=args.warmup, rep=args.rep, group=ctx.group) + rows.append({ + **cfg, "ok": bad == 0, "ms": ms, + "gbps": (moved / (ms * 1e-3) / 1e9) if ms not in (float("inf"), 0) else 0.0, + }) + dist.barrier(ctx.group) + + if not args.no_bench: + report_tuning(ctx, "allgather", rows, torch_before, time_torch(), moved) + + ctx.close() + if ctx.is_leader: + print("PASS" if failures == 0 else f"FAIL: {failures} rank(s) mismatched", flush=True) + return 1 if failures else 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/examples/distributed/internode/example_internode_allreduce.py b/examples/distributed/internode/example_internode_allreduce.py new file mode 100644 index 0000000000..c1fe107914 --- /dev/null +++ b/examples/distributed/internode/example_internode_allreduce.py @@ -0,0 +1,310 @@ +"""Inter-node allreduce over NCCL GIN, one-shot or two-shot. + +Every rank ends with the elementwise sum of all ranks' inputs, full length. +``--algo`` picks the algorithm, and which one wins is decided by rank count, not +by size: + +* ``twoshot`` (default) -- reduce-scatter then allgather, so each rank reduces one + shard and broadcasts it. Two dependent phases. Sends ``2 * (W-1) * N / W``. +* ``oneshot`` -- push the whole input to every peer, sum locally. One network + phase. Sends ``(W-1) * N`` per rank. + +Network volume alone says they tie at W=2 (the ratio is ``W / 2``) and that +one-shot should then win on having half the phases. **Measured, it loses**: 30.3 +GB/s against two-shot's 43.4 at chunks=8, 64 MB shard, two nodes. Volume parity is +not parity, because one-shot reduces over the *full* buffer rather than a shard -- +roughly twice the HBM traffic, plus a full-length local copy into scratch. Both +readings of this file's history were wrong in turn, first that one-shot only wins +when latency dominates, then that it must win at W=2; the profile decided it. + +``oneshot`` is kept because the balance shifts with rank count and dtype, and it is +the cheaper shape when the reduction is trivial relative to the transfer. + +Shape of the work +----------------- +``grid = chunks``, one CTA per chunk, owning that chunk end to end. Every +dependency stays inside one CTA, so no grid-wide barrier is needed and no +cooperative launch. See the allgather docstring for why the grid is sized by +chunks and channels rather than by payload. + +Two-shot's phases use different signal slots. Signals are cumulative running +totals, so if both phases counted into one slot the phase-2 wait could be +satisfied by phase-1 arrivals and read shards that had not been reduced yet. Both +slots advance by the same amount per launch, so one ``signal_target`` serves both. + +One-shot needs ``world_size * N`` of scratch against two-shot's ``N``, so it wants +a larger arena. +""" + +# NOTE: no `from __future__ import annotations` here. T.prim_func resolves the +# parameter annotations at runtime via get_type_hints, and PEP 563 would turn +# `T.Tensor((shard_numel,), dtype)` into a string evaluated against module +# globals -- where the closure locals shard_numel/dtype do not exist. +import argparse + +import torch +import torch.distributed as dist + +import tilelang.language as T +from tilelang.distributed.bench import do_bench + +from internode_common import ( + SIGNAL_DATA, + per_launch_signals, + report_tuning, + tune_grid, + fp32_sum, + SIGNAL_PHASE2, + Context, + TL_DTYPES, + TORCH_DTYPES, + add_common_args, + check, + prepare_env, + report, +) + + +def allreduce_oneshot_kernel(numel: int, chunks: int, threads: int, world_size: int, dtype: str, + signal_id: int = SIGNAL_DATA): + """Push the whole input to every peer, then reduce locally. One network phase. + + Volume, for ``W`` ranks and ``N`` elements: one-shot sends ``(W-1) * N`` per + rank, two-shot ``2 * (W-1) * N / W`` -- ratio ``W / 2``, so equal at W=2. + Equal bytes in one phase instead of two looks like a clear win, and it is not: + measured 30.3 GB/s against two-shot's 43.4 (chunks=8, 64 MB shard, two nodes). + The reduction here spans the whole buffer rather than a shard, so it costs + about twice the HBM traffic and adds a full-length local copy into scratch. + Not the default; see the module docstring. + + Each rank lands its contribution in slot ``rank`` of every peer's scratch, so + the reduction is a straight sum over the ``world_size`` slots. This rank's own + slot is filled by a local copy rather than a branch, since ``rank`` is a + runtime value and a trace-time branch on it is not available. + + Scratch is ``world_size * numel``: bigger than the two-shot's, and the reason + this needs a larger arena. + """ + peers = world_size - 1 + chunk_numel = numel // chunks + + @T.prim_func + def main( + inp: T.Tensor((numel,), dtype), + scratch: T.Tensor((world_size * numel,), dtype), + out: T.Tensor((numel,), dtype), + rank: T.int32, + signal_target: T.int32, + ): + with T.Kernel(chunks, threads=threads) as bx: + base = bx * chunk_numel + for step in range(peers): + peer = (rank + step + 1) % world_size + # Slot by *sender* rank, so on the peer this lands in the slot it + # reads for us. Symmetric arena, so the index matches. + T.nccl_gin.put_signal( + src=inp[base], + dst=scratch[rank * numel + base], + size=chunk_numel, + peer=peer, + signal_id=signal_id, + scope="block", + ) + # Our own contribution, no network hop. Reads inp, so it does not race + # the puts above. + for i in T.Parallel(chunk_numel): + scratch[rank * numel + base + i] = inp[base + i] + T.nccl_gin.wait_signal(least=signal_target, signal_id=signal_id, scope="block") + for i in T.Parallel(chunk_numel): + out[base + i] = T.cast( + fp32_sum( + world_size, + lambda s: T.cast(scratch[s * numel + base + i], "float32"), + ), + dtype, + ) + + return main + + +def allreduce_kernel(shard_numel: int, chunks: int, threads: int, world_size: int, dtype: str, + signal_id: int = SIGNAL_DATA, signal_id2: int = SIGNAL_PHASE2): + """Scatter-reduce into an owned chunk, then broadcast every owned chunk. + + ``out`` doubles as the phase-2 destination and as the phase-1 reduction + target: the reduced chunk is written to ``out[rank * shard + base]``, which is + exactly where the allgather wants it, so no copy sits between the phases. + """ + peers = world_size - 1 + chunk_numel = shard_numel // chunks + + @T.prim_func + def main( + inp: T.Tensor((world_size * shard_numel,), dtype), + scratch: T.Tensor((world_size * shard_numel,), dtype), + out: T.Tensor((world_size * shard_numel,), dtype), + rank: T.int32, + signal_target: T.int32, + ): + with T.Kernel(chunks, threads=threads) as bx: + base = bx * chunk_numel + # ---- phase 1: scatter-reduce ---- + for step in range(peers): + peer = (rank + step + 1) % world_size + T.nccl_gin.put_signal( + src=inp[peer * shard_numel + base], + dst=scratch[rank * shard_numel + base], + size=chunk_numel, + peer=peer, + signal_id=signal_id, + scope="block", + ) + for i in T.Parallel(chunk_numel): + scratch[rank * shard_numel + base + i] = inp[rank * shard_numel + base + i] + T.nccl_gin.wait_signal(least=signal_target, signal_id=signal_id, scope="block") + + # world_size is a Python int, so this unrolls into one fp32 + # expression. See fp32_sum for why it must not use an accumulator. + for i in T.Parallel(chunk_numel): + out[rank * shard_numel + base + i] = T.cast( + fp32_sum( + world_size, + lambda s: T.cast(scratch[s * shard_numel + base + i], "float32"), + ), + dtype, + ) + + # ---- phase 2: allgather the reduced chunks ---- + # Reading out[rank*shard + base] as the put source is safe: this CTA + # wrote exactly those bytes above, and a put by the same coop is + # ordered after the writes it depends on. + for step in range(peers): + peer = (rank + step + 1) % world_size + T.nccl_gin.put_signal( + src=out[rank * shard_numel + base], + dst=out[rank * shard_numel + base], + size=chunk_numel, + peer=peer, + signal_id=signal_id2, + scope="block", + ) + T.nccl_gin.wait_signal(least=signal_target, signal_id=signal_id2, scope="block") + + return main + + +def main() -> int: + parser = add_common_args(argparse.ArgumentParser(description=__doc__)) + # twoshot by measurement, not by theory. One-shot moves the same network bytes + # at W=2, but reduces over the full buffer instead of a shard -- ~2x the HBM + # traffic plus a full-length local copy -- and lost: 30.3 GB/s against + # two-shot's 43.4 at chunks=8. See allreduce_oneshot_kernel. + parser.add_argument("--algo", choices=("twoshot", "oneshot"), default="twoshot") + args = parser.parse_args() + + prepare_env() + ctx = Context() + + if args.numel % ctx.world_size: + raise SystemExit(f"--numel {args.numel} must be divisible by world_size {ctx.world_size}") + shard_numel = args.numel // ctx.world_size + peers = ctx.world_size - 1 + if not args.tune and shard_numel % args.chunks: + raise SystemExit(f"shard {shard_numel} must be a multiple of --chunks {args.chunks}") + + torch_dtype = TORCH_DTYPES[args.dtype] + itemsize = torch.empty((), dtype=torch_dtype).element_size() + ctx.log( + f"allreduce[{args.algo}]: world_size={ctx.world_size} nodes={ctx.num_nodes} " + f"numel={args.numel} shard={shard_numel} chunks={args.chunks} " + f"put={shard_numel // args.chunks * itemsize / 1024:.0f}KiB threads={args.threads} " + f"gin_contexts={args.gin_contexts} dtype={args.dtype}" + ) + + inp = ctx.tensor((args.numel,), torch_dtype) + # one-shot needs a slot per rank; two-shot only reduces its own shard. + scratch_numel = ctx.world_size * args.numel if args.algo == "oneshot" else args.numel + scratch = ctx.tensor((scratch_numel,), torch_dtype) + out = ctx.tensor((args.numel,), torch_dtype) + inp.copy_( + (torch.arange(args.numel, device=inp.device, dtype=torch.float32) % 7 + ctx.rank).to( + torch_dtype + ) + ) + scratch.zero_() + out.zero_() + + ref = inp.clone() + dist.all_reduce(ref, op=dist.ReduceOp.SUM, group=ctx.group) + # Both phases send one shard to each of world_size-1 peers. + moved = 2 * shard_numel * inp.element_size() * peers + ref_buf = inp.clone() + + def time_torch() -> float: + return do_bench( + lambda: dist.all_reduce(ref_buf, op=dist.ReduceOp.SUM, group=ctx.group), + warmup=args.warmup, rep=args.rep, group=ctx.group, + ) + + # two-shot burns two signal slots per config, one per phase. + configs = tune_grid(args, signals_per_config=1 if args.algo == "oneshot" else 2) + torch_before = time_torch() if not args.no_bench else float("nan") + + rows = [] + failures = 0 + for cfg in configs: + span = args.numel if args.algo == "oneshot" else shard_numel + if span % cfg["chunks"]: + continue + if args.algo == "oneshot": + func = allreduce_oneshot_kernel( + args.numel, cfg["chunks"], cfg["threads"], ctx.world_size, + TL_DTYPES[args.dtype], signal_id=cfg["signals"][0], + ) + else: + func = allreduce_kernel( + shard_numel, cfg["chunks"], cfg["threads"], ctx.world_size, + TL_DTYPES[args.dtype], signal_id=cfg["signals"][0], signal_id2=cfg["signals"][1], + ) + kernel = ctx.compile( + func, + expect=("tl::gin::put_signal_addr", "tl::gin::wait_signal"), + gin_contexts=cfg["gin_contexts"], + ) + per_launch = per_launch_signals(peers, cfg["chunks"], args.signal_div) + target = [0] + + def launch(kernel=kernel, target=target, per_launch=per_launch): + target[0] += per_launch + kernel(inp, scratch, out, ctx.rank, target[0]) + + out.zero_() + scratch.zero_() + torch.cuda.synchronize() + dist.barrier(ctx.group) + launch() + torch.cuda.synchronize() + bad = check(ctx, out, ref, f"allreduce[{args.algo}/c{cfg['chunks']}/x{cfg['gin_contexts']}]") + failures += bad + + ms = float("inf") + if bad == 0 and not args.no_bench: + dist.barrier(ctx.group) + ms = do_bench(launch, warmup=args.warmup, rep=args.rep, group=ctx.group) + rows.append({ + **cfg, "ok": bad == 0, "ms": ms, + "gbps": (moved / (ms * 1e-3) / 1e9) if ms not in (float("inf"), 0) else 0.0, + }) + dist.barrier(ctx.group) + + if not args.no_bench: + report_tuning(ctx, f"allreduce[{args.algo}]", rows, torch_before, time_torch(), moved) + + ctx.close() + if ctx.is_leader: + print("PASS" if failures == 0 else f"FAIL: {failures} rank(s) mismatched", flush=True) + return 1 if failures else 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/examples/distributed/internode/example_internode_reduce_scatter.py b/examples/distributed/internode/example_internode_reduce_scatter.py new file mode 100644 index 0000000000..0ff626f4e6 --- /dev/null +++ b/examples/distributed/internode/example_internode_reduce_scatter.py @@ -0,0 +1,212 @@ +"""Inter-node reduce-scatter over NCCL GIN. + +Every rank starts with a full-length input and ends with the elementwise sum of +all ranks' inputs, restricted to its own shard. The algorithm is the transpose of +the one-shot allgather: rank ``r`` pushes the slice destined for rank ``p`` into +slot ``r`` of ``p``'s scratch buffer, then each rank reduces the ``world_size`` +slices it has collected. + +The reduction happens on the receiving rank rather than in flight because GIN +2.28.9 has no remote atomic or reduce-on-put; the only device op that moves +payload is a plain put. Landing the contributions in separate slots and summing +locally costs one extra read of the scratch buffer, which is HBM-local and cheap +next to a network hop. + +Shape of the work +----------------- +``grid = chunks`` -- one CTA per chunk, each owning that chunk end to end: it +sends the chunk to every peer, waits, then reduces it. Sizing the grid by chunks +rather than by payload keeps messages large, which is what RDMA needs (see the +allgather docstring for the measurement that motivated this). + +The grid is *not* ``peers * chunks`` here, unlike the allgather. A CTA has to +reduce the whole chunk it owns, so splitting a chunk across CTAs by peer would +make the reduction depend on slices other CTAs received -- a grid-wide barrier +inside one kernel, which needs a cooperative launch. Looping over peers inside +the CTA keeps every dependency CTA-local. Puts to different peers land on +different QPs regardless, since a context has one QP per peer. + +Scratch is allocated from the arena like everything else -- a GIN destination has +to live in the registered window. +""" + +# NOTE: no `from __future__ import annotations` here. T.prim_func resolves the +# parameter annotations at runtime via get_type_hints, and PEP 563 would turn +# `T.Tensor((shard_numel,), dtype)` into a string evaluated against module +# globals -- where the closure locals shard_numel/dtype do not exist. +import argparse + +import torch +import torch.distributed as dist + +import tilelang.language as T +from tilelang.distributed.bench import do_bench + +from internode_common import ( + SIGNAL_DATA, + per_launch_signals, + report_tuning, + tune_grid, + fp32_sum, + Context, + TL_DTYPES, + TORCH_DTYPES, + add_common_args, + check, + prepare_env, + report, +) + + +def reduce_scatter_kernel(shard_numel: int, chunks: int, threads: int, world_size: int, dtype: str, + signal_id: int = SIGNAL_DATA): + """Scatter each peer's slice, then sum the slices that arrive for this rank. + + One launch does both phases. The wait sits between them, so the reduction + only reads slots whose payload has been signalled as landed. + """ + peers = world_size - 1 + chunk_numel = shard_numel // chunks + + @T.prim_func + def main( + inp: T.Tensor((world_size * shard_numel,), dtype), + scratch: T.Tensor((world_size * shard_numel,), dtype), + out: T.Tensor((shard_numel,), dtype), + rank: T.int32, + signal_target: T.int32, + ): + with T.Kernel(chunks, threads=threads) as bx: + base = bx * chunk_numel + # `peers` is a Python int, so this unrolls into independent puts. + for step in range(peers): + peer = (rank + step + 1) % world_size + # Send the part of our input that belongs to `peer`, into the + # slot `peer` reserves for us. Symmetric arena, so the index we + # write is the index the peer reads. + T.nccl_gin.put_signal( + src=inp[peer * shard_numel + base], + dst=scratch[rank * shard_numel + base], + size=chunk_numel, + peer=peer, + signal_id=signal_id, + scope="block", + ) + # Our own contribution needs no network hop. + for i in T.Parallel(chunk_numel): + scratch[rank * shard_numel + base + i] = inp[rank * shard_numel + base + i] + T.nccl_gin.wait_signal(least=signal_target, signal_id=signal_id, scope="block") + # world_size is a Python int, so this unrolls into one fp32 + # expression. See fp32_sum for why it must not use an accumulator. + for i in T.Parallel(chunk_numel): + out[base + i] = T.cast( + fp32_sum( + world_size, + lambda s: T.cast(scratch[s * shard_numel + base + i], "float32"), + ), + dtype, + ) + + return main + + +def main() -> int: + parser = add_common_args(argparse.ArgumentParser(description=__doc__)) + args = parser.parse_args() + + prepare_env() + ctx = Context() + + if args.numel % ctx.world_size: + raise SystemExit(f"--numel {args.numel} must be divisible by world_size {ctx.world_size}") + shard_numel = args.numel // ctx.world_size + peers = ctx.world_size - 1 + if not args.tune and shard_numel % args.chunks: + raise SystemExit(f"shard {shard_numel} must be a multiple of --chunks {args.chunks}") + + torch_dtype = TORCH_DTYPES[args.dtype] + itemsize = torch.empty((), dtype=torch_dtype).element_size() + ctx.log( + f"reduce_scatter: world_size={ctx.world_size} nodes={ctx.num_nodes} " + f"numel={args.numel} shard={shard_numel} chunks={args.chunks} " + f"put={shard_numel // args.chunks * itemsize / 1024:.0f}KiB threads={args.threads} " + f"gin_contexts={args.gin_contexts} dtype={args.dtype}" + ) + + inp = ctx.tensor((args.numel,), torch_dtype) + scratch = ctx.tensor((args.numel,), torch_dtype) + out = ctx.tensor((shard_numel,), torch_dtype) + # Small magnitudes: a bf16 sum of 16 ranks' worth of arange values would + # otherwise land outside any sensible tolerance. + inp.copy_( + (torch.arange(args.numel, device=inp.device, dtype=torch.float32) % 7 + ctx.rank).to( + torch_dtype + ) + ) + scratch.zero_() + out.zero_() + + ref = torch.empty_like(out) + dist.reduce_scatter_tensor(ref, inp, op=dist.ReduceOp.SUM, group=ctx.group) + moved = out.numel() * out.element_size() * peers + ref_buf = torch.empty_like(out) + + def time_torch() -> float: + return do_bench( + lambda: dist.reduce_scatter_tensor(ref_buf, inp, op=dist.ReduceOp.SUM, group=ctx.group), + warmup=args.warmup, rep=args.rep, group=ctx.group, + ) + + configs = tune_grid(args, signals_per_config=1) + torch_before = time_torch() if not args.no_bench else float("nan") + + rows = [] + failures = 0 + for cfg in configs: + if shard_numel % cfg["chunks"]: + continue + kernel = ctx.compile( + reduce_scatter_kernel( + shard_numel, cfg["chunks"], cfg["threads"], ctx.world_size, + TL_DTYPES[args.dtype], signal_id=cfg["signals"][0], + ), + expect=("tl::gin::put_signal_addr", "tl::gin::wait_signal"), + gin_contexts=cfg["gin_contexts"], + ) + per_launch = per_launch_signals(peers, cfg["chunks"], args.signal_div) + target = [0] + + def launch(kernel=kernel, target=target, per_launch=per_launch): + target[0] += per_launch + kernel(inp, scratch, out, ctx.rank, target[0]) + + out.zero_() + scratch.zero_() + torch.cuda.synchronize() + dist.barrier(ctx.group) + launch() + torch.cuda.synchronize() + bad = check(ctx, out, ref, f"reduce_scatter[c{cfg['chunks']}/x{cfg['gin_contexts']}]") + failures += bad + + ms = float("inf") + if bad == 0 and not args.no_bench: + dist.barrier(ctx.group) + ms = do_bench(launch, warmup=args.warmup, rep=args.rep, group=ctx.group) + rows.append({ + **cfg, "ok": bad == 0, "ms": ms, + "gbps": (moved / (ms * 1e-3) / 1e9) if ms not in (float("inf"), 0) else 0.0, + }) + dist.barrier(ctx.group) + + if not args.no_bench: + report_tuning(ctx, "reduce_scatter", rows, torch_before, time_torch(), moved) + + ctx.close() + if ctx.is_leader: + print("PASS" if failures == 0 else f"FAIL: {failures} rank(s) mismatched", flush=True) + return 1 if failures else 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/examples/distributed/internode/internode_common.py b/examples/distributed/internode/internode_common.py new file mode 100644 index 0000000000..f795d7f87b --- /dev/null +++ b/examples/distributed/internode/internode_common.py @@ -0,0 +1,398 @@ +"""Shared setup for the inter-node GIN collectives. + +Every example here follows the same shape: bring up a process group with the +network enabled, allocate symmetric buffers from the TileScale allocator (GIN can +only address the registered arena), run a kernel that moves bytes with +``T.nccl_gin``, check the result against torch, and time both. + +The pieces live here rather than in each example because the environment setup is +easy to get subtly wrong -- ``init_dist`` disables InfiniBand by default, and a +cudaMalloc arena cannot be registered as an NCCL window -- and a silent mistake +in either shows up as a passing test that moved no data over the fabric. +""" + +from __future__ import annotations + +import argparse +import functools +import operator +import os + +import torch +import torch.distributed as dist + +import tilelang +import tilelang.language as T + +# Signal slots. The reduce-scatter half of allreduce and its allgather half must +# not share a slot: signals are cumulative totals, so two phases counting into +# one slot cannot be told apart. +SIGNAL_DATA = 0 +SIGNAL_PHASE2 = 1 + + +def fp32_sum(count: int, term): + """Fold ``term(0) + ... + term(count-1)`` into one add-expression. + + ``term`` maps a source index to a PrimExpr; callers cast to float32 inside it, + because a bf16 running sum over 16 ranks loses enough low bits to fail a + tolerance check. + + Both the loop and the fold live here, outside any traced function, and that + placement is the point. Two shapes that look more natural both fail inside a + ``T.prim_func``: + + * ``acc = T.cast(...)`` then ``acc = acc + ...`` -- in a kernel body the eager + builder treats assignment of a PrimExpr as a TIR variable *declaration*, so + the reassignment emits a second variable (``acc_1``) and the enclosing + ``T.Parallel`` frame rejects it. + * ``fp32_sum([... for src in range(n)])`` -- the builder rewrites ``for`` + statements *and comprehension for-clauses* into TIR loops, so the + comprehension raises ``'ForFrame' object is not iterable``. + + Calling a plain Python helper sidesteps both: the fold runs at trace time and + only its result, a single expression, is emitted. That keeps the copy + vectorisable. + """ + return functools.reduce(operator.add, (term(k) for k in range(count))) + + +def prepare_env() -> None: + """Set the environment GIN needs, before ``init_dist`` is imported or run. + + ``init_dist`` sets ``NCCL_IB_DISABLE=1`` unless it is already set, which would + keep every transfer inside shared memory and make an "inter-node" benchmark + measure nothing. The VMM and GIN flags are hard requirements of window + registration; asserting them here turns a missing Device API into an error at + startup instead of a null devcomm read inside a kernel. + """ + os.environ["NCCL_IB_DISABLE"] = "0" + os.environ.setdefault("TILESCALE_USE_VMM", "1") + os.environ.setdefault("TILESCALE_USE_GIN", "1") + os.environ.setdefault("NCCL_DEBUG", "ERROR") + + +def add_common_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: + parser.add_argument( + "--numel", + type=int, + default=1 << 22, + help="elements in the collective's logical input, summed over ranks", + ) + parser.add_argument("--block", type=int, default=8192, help="elements per CTA") + # Chunks exist to spread a peer's transfer across GIN contexts (one context + # is one QP per peer, so one NIC), NOT to parallelise within a channel. + # Splitting further than the context count only shrinks messages. + parser.add_argument( + "--chunks", + type=int, + default=64, + # Tuned on two idle nodes, 64 MB shard: the reduce variants climb with + # chunk count and peak at 64 (allreduce 41.5 -> 47.2 GB/s from 4 to 64, + # reduce_scatter 44.0 -> 46.2 from 16 to 64; 128 is worse), because their + # CTAs also carry the reduction and want more of them. Allgather is flat + # from 2 to 64, so 64 is a safe shared default. + help="puts per peer; each becomes one CTA issuing one large put", + ) + parser.add_argument( + "--gin-contexts", + type=int, + default=4, + # Contexts are insurance against a busy NIC, not a win on an idle one. + # On two *idle* nodes one context already reaches line rate and extra + # contexts cost ~1% (allgather 47.6 GB/s at 1 vs 47.0 at 4). On a NIC + # shared with another tenant's job, one context collapsed to 23 GB/s while + # 2-4 held ~44. Defaulting to 4 trades that 1% for the contended case. + # The device clamps to what the devcomm granted (4 here, though the + # allocator asks for 8) and scales the wait target to match. + help="-DTL_GIN_CONTEXTS: spread CTAs over n GIN contexts (QPs); 1 pins to context 0", + ) + parser.add_argument( + "--signal-div", + type=int, + default=0, + help="DEBUG ONLY: divide the wait target; under-waits, so any result is invalid", + ) + parser.add_argument( + "--wait-ctx0", + action="store_true", + help="-DTL_GIN_WAIT_CTX0: puts spread over contexts, every wait on context 0", + ) + # 1024 threads matches triton-dist's num_warps=32 for its inter-node send + # blocks: one CTA cooperatively driving one large put. + parser.add_argument("--threads", type=int, default=1024) + parser.add_argument("--dtype", choices=("fp32", "bf16", "fp16"), default="bf16") + parser.add_argument("--warmup", type=int, default=20) + parser.add_argument("--rep", type=int, default=50) + parser.add_argument("--no-bench", action="store_true", help="check correctness only") + parser.add_argument("--print-source", action="store_true") + parser.add_argument( + "--tune", + action="store_true", + help="sweep the grid below in one process, verify each, report the best vs torch", + ) + parser.add_argument("--tune-chunks", default="2,4,8,16", help="--tune: chunk counts to try") + parser.add_argument("--tune-contexts", default="1,2,4", help="--tune: GIN context counts") + parser.add_argument("--tune-threads", default="1024", help="--tune: threads per CTA") + return parser + + +# GIN_SIGNAL_COUNT in nccl_window.py. Signals are "guaranteed to start at id=0", +# so ids [0, 32) are usable. +MAX_SIGNALS = 32 + + +def tune_grid(args, signals_per_config: int = 1): + """Configurations to try, each with its own signal slots. + + Why distinct slots per config rather than one shared slot: signals are + cumulative and nothing resets them, and ``wait_signal`` divides the target by + the device's ``context_span()``. A config that changes the context count would + therefore divide a total accumulated under a *different* span, so the target + would be wrong from that point on. Giving each config fresh slots means every + counter starts at zero under exactly one span, which makes ``--tune-contexts`` + safe to sweep in a single process. + """ + if not args.tune: + return [{ + "chunks": args.chunks, + "gin_contexts": args.gin_contexts, + "threads": args.threads, + "signals": tuple(range(signals_per_config)), + }] + + ints = lambda s: [int(x) for x in str(s).split(",") if x != ""] + grid = [] + for contexts in ints(args.tune_contexts): + for chunks in ints(args.tune_chunks): + for threads in ints(args.tune_threads): + # A context must back the same number of CTAs for one wait target + # to be right, so the span has to divide the chunk count. + if chunks % max(1, contexts): + continue + base = len(grid) * signals_per_config + if base + signals_per_config > MAX_SIGNALS: + break + grid.append({ + "chunks": chunks, + "gin_contexts": contexts, + "threads": threads, + "signals": tuple(range(base, base + signals_per_config)), + }) + if not grid: + raise SystemExit("--tune grid is empty; check --tune-chunks / --tune-contexts") + return grid + + +def report_tuning(ctx, name: str, rows, torch_ms_before: float, torch_ms_after: float, moved: int): + """Print every configuration, then the winner against torch. + + torch is timed twice, before and after the sweep. If the two disagree the + fabric moved under us and the ratios are not trustworthy -- worth knowing, + since another tenant's job on the same NICs shifts these by up to 2x. + """ + if not ctx.is_leader: + return + gbps = lambda ms: moved / (ms * 1e-3) / 1e9 + print(f"\n===== {name}: tuning results ({len(rows)} configs) =====", flush=True) + print(f"{'chunks':>7} {'ctx':>4} {'thr':>5} {'ms':>9} {'GB/s':>8} status", flush=True) + for r in sorted(rows, key=lambda r: -(r["gbps"] if r["ok"] else -1)): + print( + f"{r['chunks']:>7} {r['gin_contexts']:>4} {r['threads']:>5} " + f"{r['ms']:>9.3f} {r['gbps']:>8.1f} {'PASS' if r['ok'] else 'FAIL'}", + flush=True, + ) + good = [r for r in rows if r["ok"]] + if not good: + print("no configuration passed", flush=True) + return + best = max(good, key=lambda r: r["gbps"]) + t_before, t_after = gbps(torch_ms_before), gbps(torch_ms_after) + t_ms = min(torch_ms_before, torch_ms_after) # torch at its best + t_best = gbps(t_ms) + drift = abs(t_before - t_after) / max(t_before, t_after) + print( + f"\n{name}: BEST chunks={best['chunks']} contexts={best['gin_contexts']} " + f"threads={best['threads']}\n" + f" tilescale {best['ms']:.3f} ms {best['gbps']:.1f} GB/s\n" + f" torch {t_ms:.3f} ms {t_best:.1f} GB/s " + f"(measured {t_before:.1f} then {t_after:.1f} GB/s, drift {drift * 100:.0f}%)\n" + f" speedup {best['gbps'] / t_best:.2f}x vs torch's best of the two", + flush=True, + ) + if drift > 0.15: + print(" WARNING: torch drifted >15% across the sweep; fabric was not stable", flush=True) + + +def per_launch_signals(peers: int, chunks: int, signal_div: int = 0) -> int: + """Signal increments this rank actually receives per launch. + + Each of ``peers`` senders signals once per chunk, so the honest target is + ``peers * chunks``. This is the only value that makes ``wait_signal`` mean + "all my data has landed". + + ``signal_div`` divides it, and exists solely to investigate why multiple GIN + contexts hang (see nccl_gin.h). It makes the wait weaker than the data, so the + kernel can return before the payload arrives -- which shows up as bandwidth + above what the hardware can carry. Any number produced with signal_div set is + not a measurement of the collective. + """ + total = peers * chunks + if not signal_div: + return total + if total % signal_div: + raise SystemExit(f"--signal-div {signal_div} must divide peers*chunks = {total}") + return total // signal_div + + +TORCH_DTYPES = {"fp32": torch.float32, "bf16": torch.bfloat16, "fp16": torch.float16} +TL_DTYPES = {"fp32": "float32", "bf16": "bfloat16", "fp16": "float16"} + + +class Context: + """Process group, allocator and topology for one rank.""" + + def __init__(self, arena_bytes: int = 1 << 30): + from tilelang.distributed.host import init_dist + + self.local_rank = int(os.environ.get("LOCAL_RANK", 0)) + self.local_world_size = int( + os.environ.get("LOCAL_WORLD_SIZE", torch.cuda.device_count()) + ) + # Staged so a hang can be attributed. init_dist, allocator construction + # (which creates the GIN devcomm and registers the arena window) and + # compile are all collective; without markers they are one opaque block. + if os.environ.get("TL_STAGE_TRACE"): + print(f"[rank?] init_dist: enter local_rank={self.local_rank}", flush=True) + self.rank, self.world_size, self.group, self.node_info = init_dist( + self.local_rank, self.local_world_size, return_node_info=True + ) + self.trace("init_dist: done") + self.num_nodes = self.node_info.num_nodes if self.node_info is not None else 1 + self.trace( + f"allocator: enter bytes={arena_bytes} nodes={self.num_nodes} " + f"world={self.world_size} (devcomm + window register)" + ) + self.allocator = tilelang.get_allocator( + size=arena_bytes, + device="cuda", + is_distributed=True, + local_rank=self.local_rank, + num_local_ranks=self.local_world_size, + group=self.group, + node_info=self.node_info, + ) + self.trace("allocator: done (arena window live)") + + @property + def is_leader(self) -> bool: + return self.rank == 0 + + def tensor(self, shape, dtype: torch.dtype): + """Allocate from the arena. Required: only the arena is a GIN window.""" + return tilelang.tensor(shape, dtype, allocator=self.allocator) + + def log(self, msg: str) -> None: + if self.is_leader: + print(msg, flush=True) + + def trace(self, msg: str) -> None: + """Print from every rank, tagged. + + ``log`` is leader-only, so a non-leader rank cannot report progress at + all: it looks identical whether it hung in setup, hung in compile, or + ran fine. That ambiguity is what made the first two-node hang + undiagnosable. Enabled by TL_STAGE_TRACE=1 to keep normal runs quiet. + """ + if os.environ.get("TL_STAGE_TRACE"): + print(f"[rank{self.rank}] {msg}", flush=True) + + def compile(self, func, *, expect: tuple[str, ...] = (), gin_contexts: int | None = None, + wait_ctx0: bool = False): + # Every rank checks the tokens, not just the leader. compile_once makes + # this a collective, so a leader-only assertion aborts rank 0 while the + # others march on into close()'s barrier and hang until the outer + # timeout -- turning a clear assertion failure into a mystery stall. + # + # gin_contexts becomes -DTL_GIN_CONTEXTS=n. compile_flags is part of the + # cache key, so each setting gets its own cache entry -- which is also + # the only thing that keeps a sweep honest, since the key does not cover + # the device headers this define lives in. + flags = None if gin_contexts is None else [f"-DTL_GIN_CONTEXTS={int(gin_contexts)}"] + if wait_ctx0: + flags = (flags or []) + ["-DTL_GIN_WAIT_CTX0=1"] + if os.environ.get("TL_GIN_DEBUG"): + flags = (flags or []) + ["-DTL_GIN_DEBUG=1"] + self.trace(f"compile: enter (collective in compile_once) flags={flags}") + kernel = tilelang.compile( + func, compile_once=True, compile_group=self.group, compile_flags=flags + ) + self.trace("compile: lowered") + if expect: + source = kernel.get_kernel_source() + # Without this the kernel still compiles and silently moves nothing, + # which would read as a fast and correct-looking result. + for token in expect: + assert token in source, f"lowering did not emit {token!r}" + assert "nccl_gin.h" in source, "generated code is missing the GIN header" + kernel.initialize(allocator=self.allocator) + self.trace("compile: initialized") + return kernel + + def close(self) -> None: + # allocator.close() is collective, so every rank must reach it even if + # this rank's check failed. Guard the barrier: if a peer already died, + # blocking here forever converts its error into a timeout on this rank + # and buries the real message. + self.trace("close: barrier") + try: + dist.barrier(self.group) + except Exception as exc: # noqa: BLE001 - report and keep tearing down + print(f"[rank{self.rank}] close: barrier failed: {exc}", flush=True) + self.trace("close: allocator") + self.allocator.close() + dist.destroy_process_group() + self.trace("close: done") + + +def check(ctx: Context, got: torch.Tensor, want: torch.Tensor, name: str) -> int: + """Compare on every rank and aggregate, so one bad rank fails the run.""" + # bf16 accumulation order differs between a tree reduction and our linear + # one, so compare with a tolerance rather than exactly. + if got.dtype in (torch.bfloat16, torch.float16): + ok = torch.allclose(got.float(), want.float(), rtol=6e-2, atol=6e-2) + else: + ok = torch.allclose(got, want, rtol=1e-5, atol=1e-5) + if not ok: + diff = (got.float() - want.float()).abs() + bad = (diff > 6e-2).nonzero().flatten() + print( + f"[rank {ctx.rank}] {name} MISMATCH: {bad.numel()}/{got.numel()} differ, " + f"max |diff| {diff.max().item():.4g}, first at {bad[0].item() if bad.numel() else -1}", + flush=True, + ) + status = torch.tensor([0 if ok else 1], device=got.device, dtype=torch.int32) + dist.all_reduce(status, group=ctx.group) + failures = int(status.item()) + if failures == 0: + ctx.log(f"{name}: correct on all {ctx.world_size} ranks") + return failures + + +def report(ctx: Context, name: str, tl_ms: float, ref_ms: float, moved_bytes: int) -> None: + """Print both timings plus the bus bandwidth each implies. + + ``moved_bytes`` is the payload that has to cross a rank's network link, not + the buffer size, so the number is comparable between collectives with + different algorithmic volumes. + """ + if not ctx.is_leader: + return + tl_gbps = moved_bytes / (tl_ms * 1e-3) / 1e9 + ref_gbps = moved_bytes / (ref_ms * 1e-3) / 1e9 + speedup = ref_ms / tl_ms if tl_ms > 0 else float("nan") + print( + f"{name:<16} tilescale {tl_ms:8.3f} ms {tl_gbps:7.1f} GB/s | " + f"torch {ref_ms:8.3f} ms {ref_gbps:7.1f} GB/s | speedup {speedup:5.2f}x", + flush=True, + ) diff --git a/examples/distributed/internode/run_internode.sh b/examples/distributed/internode/run_internode.sh new file mode 100755 index 0000000000..046e600d25 --- /dev/null +++ b/examples/distributed/internode/run_internode.sh @@ -0,0 +1,122 @@ +#!/bin/bash +# Launch an inter-node collective example across two nodes, one GPU each. +# +# Usage: run_internode.sh