diff --git a/AGENTS.md b/AGENTS.md index 2db41f9..74ec946 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -56,6 +56,7 @@ Optional CMake switches are `TILEXR_BUILD_COLLECTIVES`, `TILEXR_BUILD_EP`, `TILE - Ring a UDMA doorbell only with `st_dev`, after the corresponding MTE3 WQE write has completed. Scalar stores must never be used for doorbells. - Never put `${ASCEND_HOME_PATH}/${ARCH}-linux/devlib` in runtime RPATH/RUNPATH; runtime must resolve the real driver HAL. - Never use Host wrappers that launch Ascend C kernels with `kernel<<<...>>>` syntax. Build and embed pure AICore binaries, register them with `rtDevBinaryRegister` and `rtFunctionRegister`, then launch the registered signature with `rtKernelLaunchWithFlagV2`. +- MoonEP source changes must preserve the API contract consumed by `tools/moonep/test_npu_e2e.py`; do not change that test's calls to accommodate an implementation change. Preserve the `Buffer` and `MoonEPCommPlan` names, method signatures and keyword arguments, return tuple arity and order, exposed plan fields, and async-event and zero-copy call contracts exercised by the test. - Host, simulator, and 910B fallback tests do not prove UDMA data-plane transfer. Keep validation claims scoped to the hardware actually exercised. diff --git a/docs/plans/2026-08-11-moonep-combine-v1-memory-single-launch.md b/docs/plans/2026-08-11-moonep-combine-v1-memory-single-launch.md new file mode 100644 index 0000000..fee5f38 --- /dev/null +++ b/docs/plans/2026-08-11-moonep-combine-v1-memory-single-launch.md @@ -0,0 +1,105 @@ +# MoonEP Combine V1 Memory Single-Launch Implementation Plan + +## Goal + +Implement the approved single-launch Combine V1 Memory baseline while keeping +the complete flow on Combine V2 by default. + +## Scope And Constraints + +- Follow `docs/specs/2026-08-11-moonep-combine-v1-memory-single-launch-design.md`. +- Preserve C++14, CANN 9.1, pure AICore binary registration, and peer-memory + cross-node semantics. +- Preserve existing unrelated Dispatch V2, Combine V2, PrefetchWeight, and + ReduceGrad changes in this worktree. +- Do not commit, push, or create a PR unless separately requested. +- Real hardware validation is limited to `141.61.49.223`. + +## Ordered Work + +### 1. Lock The Host/Layout Contract + +Objective: replace split-phase layout behavior with the single-launch peer +window layout and V2 block scheduling. + +Scope: + +- `src/include/tilexr_moonep.h` +- `src/moonep/combine/host/combine_layout.*` +- `src/moonep/combine/host/combine_host.*` +- `src/moonep/combine/host/combine_launch.*` +- `src/moonep/combine/common/*` +- focused tests under `tests/moonep/unit/` + +Acceptance: + +- flags other than NONE fail; +- the replacement V1 descriptor requires and forwards `dstLocal`; +- source, receive, control, failure, and chunk regions fit 512 MiB; +- smallest and multi-chunk shapes produce nonzero aligned chunks; +- supported rank sizes produce the V2 active core count; +- one V1 call delegates to one registered-kernel launch. + +### 2. Replace The V1 Kernel + +Objective: implement one-launch duplicate pre-reduction, V2 peer/core scheduling, +memory push, completion, receive reduction, and drain. + +Scope: + +- `src/moonep/combine/kernels/tilexr_moonep_combine_kernel.cpp` +- shared V2 scheduling helpers only when they are transport-independent +- kernel source guards and schedule/layout tests + +Acceptance: + +- old publish-only/consume-only and registered-UDMA branches are absent; +- duplicate groups are partitioned without overlapping writes; +- hidden duplicate rows are skipped after pre-reduction while weight routes are preserved; +- all active cores and steps cover every peer exactly once; +- payload MTE3 precedes done publication; +- output token partitions have no gaps or overlap; +- hidden and optional weights complete in the same kernel launch; +- target `bisheng` compile succeeds. + +### 3. Add The Runtime Switch + +Objective: default to V2 and allow explicit V1 Memory selection. + +Scope: + +- `integrations/moonep_torch/tilexr_moonep/runtime.py` +- `integrations/moonep_torch/tilexr_moonep/torch_api.py` +- `tools/moonep/benchmark.py` +- Python FFI, mode, and performance-report tests + +Acceptance: + +- unset/2 calls V2; +- 1 calls V1 exactly once with hidden and optional weights; +- invalid values fail before communicator creation; +- metadata and trace labels match the selected backend; +- no V1-only workspace activation is added. + +### 4. Local And Target Validation + +Objective: establish Host, compile, integration, and hardware evidence. + +Checks: + +- focused Python tests, then `python -m pytest tests/moonep/python -q`; +- focused CTest targets for V1 Host/layout/source and V2 regressions; +- target CANN 9.1 configure/build/install on the existing isolated deployment; +- 8-NPU V1 benchmark/reference/correctness; +- 8-NPU default V2 benchmark/reference/correctness; +- `git diff --check`, final diff inspection, and original-workspace preservation. + +The benchmark uses warmup 5 and 20 measured iterations. Performance results are +reported with the benchmark's own validity scope and are not generalized to +unvalidated multi-node transport. + +## Dependencies And Coordination + +The tasks are intentionally sequential. Host layout fields define the kernel +ABI; the kernel ABI defines the Python selection tests and target build. No task +is safely independent while these shared contracts are changing. diff --git a/docs/specs/2026-08-11-moonep-combine-v1-memory-single-launch-design.md b/docs/specs/2026-08-11-moonep-combine-v1-memory-single-launch-design.md new file mode 100644 index 0000000..da8789f --- /dev/null +++ b/docs/specs/2026-08-11-moonep-combine-v1-memory-single-launch-design.md @@ -0,0 +1,128 @@ +# MoonEP Combine V1 Memory Single-Launch Design + +## Goal + +Replace the existing MoonEP Combine V1 implementation with a single-launch +peer-memory baseline. The baseline keeps the saved duplicate metadata algorithm +but follows the Combine V2 rank/core/peer schedule and receive-side reduction +flow. Combine V2 remains the default complete-flow backend. + +## Scope + +- Direct AICore binary registration and `rtKernelLaunchWithFlagV2` invocation. +- Ascend 910A5 / Ascend950 with CANN 9.1. +- BF16 hidden input/output and optional FP32 route weights. +- Rank sizes supported by `MoonEpCombineV2RankSizeSupported`. +- One Host launch per Combine call, including hidden and optional weights. +- Memory transport through `CommArgs.peerMems[]` for same-node and cross-node + mappings. + +## Non-Goals + +- Preserve the old V1 publish-only or consume-only split-launch behavior. +- Preserve V1 status values or internal kernel ABI. +- Add a topology-only memory/UDMA selector. +- Change Combine V2, Dispatch V2, or the Planner output contract. +- Claim real multi-node validation while only one host is available. + +## Public And Runtime Contract + +`TileXRMoonEpCombineV1` remains the V1 entry point. Its descriptor is replaced, +without binary-compatibility guarantees, to carry the Planner V3 reverse route +map required by a V2-style push. + +- `args.flags` must be `TILEXR_MOONEP_FLAG_NONE`. +- `args.dstLocal` points to `NvS` `int32_t` entries in device memory. +- `dstLocal[expertRecvSlot] = srcRank * NvS + token * K + topk`; `-1` marks an + unused slot. `plan.dst` is not a valid substitute because it has the opposite + route direction. +- `plan.dst`, `plan.dupGroups`, `plan.dupLoffs`, and `plan.dupCounts` remain + required for the legacy duplicate contract and plan validation. +- Hidden tensors are BF16 `[NvS,H]` and `[S,H]`. +- Route weights are either both absent or FP32 `[NvS]` and `[S,K]`. +- A successful kernel stores `0` in `plan.status`; positive values identify + device validation, timeout, or peer failures. +- The call is asynchronous on the supplied stream and launches exactly one + registered AICore kernel. + +Python selects the backend once when the runtime is constructed: + +- unset or `TILEXR_MOONEP_COMBINE_VERSION=2`: Combine V2; +- `TILEXR_MOONEP_COMBINE_VERSION=1`: Combine V1 Memory; +- any other value: fail before communicator initialization. + +Benchmark metadata and trace labels report the selected implementation. + +## Host Layout + +The peer data window is limited by `TileXR::IPC_BUFF_MAX_SIZE`. Host layout uses +checked arithmetic to reserve: + +1. a source region containing `NvS` chunk rows and optional route weights; +2. a receive region containing `NvS` chunk rows and optional route weights; +3. magic-tagged per-epoch, per-rank, per-core completion and drain records; +4. per-core failure records. + +The largest 32-byte-aligned hidden chunk stride that fits all regions is chosen. +The kernel loops over `ceil(hiddenRowBytes / hiddenChunkBytes)` chunks. A zero +chunk, arithmetic overflow, or a layout beyond the peer window is rejected. + +`blockDim` is `MoonEpCombineV2ActiveCoreCount(rankSize)`, bounded by the device +vector-core count. Every Host layout field has an explicit byte unit and is +passed in the exact direct-launch ABI order. + +## Kernel Flow + +For every hidden chunk, within one kernel launch: + +1. Cores cooperatively copy local hidden rows into the source region. Route + weights are copied during the first chunk. +2. Duplicate groups are partitioned across active cores. Each core marks the + duplicate source slots, accumulates their BF16 hidden rows into the primary + row in FP32, and converts the primary row back to BF16 once. Route weights + are not pre-reduced because every TopK route retains its own weight. +3. A local magic-tagged barrier prevents any core from sending before all + duplicate groups are complete. +4. Each core follows `MoonEpCombineV2Peer(rank, step, core, rankSize)`. It scans + `dstLocal` as V2 does and MTE3-copies rows from its source region into the + target rank's receive region. Duplicate hidden rows are skipped after their + contribution has been folded into the primary; weight rows are all copied. +5. After payload MTE3 completion, the source core publishes a magic-tagged done + record to the target peer window. The receiver waits only for the source set + assigned by the V2 schedule. +6. Receiver cores partition output tokens exactly as V2. BF16 routes are loaded + from the local receive region, converted to FP32, summed, converted once to + BF16, and written directly to `hiddenSh`. +7. Optional FP32 route weights are copied by route without arithmetic. +8. A magic-tagged drain barrier completes before source/receive regions are + reused for the next chunk. + +GM-to-UB and UB-to-GM transfers use the CANN 9.1 `DataCopyPad` forms already +compiled by the neighboring kernels. Queue/event dependencies must order MTE2, +Vector, and MTE3 operations. The completion record must never become visible +before its payload MTE3 writes complete. + +## Failure Handling + +- Host validation returns without launching on invalid descriptors, flags, + rank topology, core count, peer mappings, or capacity. +- Device validation checks destination encoding and duplicate metadata bounds. +- Each core publishes a failure record; core 0 converges records and writes the + final plan status. +- All waits are bounded by the existing MoonEP wait-iteration contract. +- Magic-tagged epochs avoid clearing the complete peer window between calls. + +## Verification + +1. Host layout/schedule tests cover smallest, tail, multi-chunk, overflow, rank + sizes, core coverage, and control-region bounds. +2. Host/launch tests prove flags are rejected, duplicate pointers are forwarded, + blockDim follows V2, and one API call performs one kernel launch. +3. Source guards prove the split-launch branches and UDMA WQE/CQ paths are absent + from V1 and V2 schedule helpers are used. +4. Python tests cover default V2, explicit V1, invalid switch values, dynamic + metadata, and one V1 FFI call containing both hidden and weights. +5. The target CANN 9.1 build compiles the kernel and all focused CTest targets. +6. On `141.61.49.223`, run V1 benchmark/reference/correctness on one 8-NPU host, + then rerun the same three modes with the default V2 backend. +7. Multi-node behavior is limited to Host/mock protocol coverage. diff --git a/integrations/moonep_torch/tilexr_moonep/abi.py b/integrations/moonep_torch/tilexr_moonep/abi.py index 2ca7387..0c2ac63 100644 --- a/integrations/moonep_torch/tilexr_moonep/abi.py +++ b/integrations/moonep_torch/tilexr_moonep/abi.py @@ -104,6 +104,10 @@ class TileXRMoonEPDispatchArgsV1(ctypes.Structure): ] +class TileXRMoonEPDispatchArgsV2(ctypes.Structure): + _fields_ = TileXRMoonEPDispatchArgsV1._fields_ + + class TileXRMoonEPPrefetchWeightArgsV1(ctypes.Structure): _fields_ = [ ("structSize", ctypes.c_uint32), @@ -123,6 +127,7 @@ class TileXRMoonEPCombineArgsV1(ctypes.Structure): ("abiVersion", ctypes.c_uint32), ("comm", ctypes.c_void_p), ("plan", ctypes.POINTER(TileXRMoonEPPlanV1)), + ("dstLocal", ctypes.c_void_p), ("hiddenNvsh", ctypes.POINTER(TileXRMoonEPTensorV1)), ("routeWeightsNvs", ctypes.POINTER(TileXRMoonEPTensorV1)), ("hiddenSh", ctypes.POINTER(TileXRMoonEPTensorV1)), diff --git a/integrations/moonep_torch/tilexr_moonep/runtime.py b/integrations/moonep_torch/tilexr_moonep/runtime.py index 8e1cc29..6b1ac30 100644 --- a/integrations/moonep_torch/tilexr_moonep/runtime.py +++ b/integrations/moonep_torch/tilexr_moonep/runtime.py @@ -9,11 +9,11 @@ from .abi import ( TILEXR_MOONEP_ABI_VERSION, - TILEXR_MOONEP_FLAG_BUILD_DEDUP, TILEXR_MOONEP_FLAG_NONE, TILEXR_SUCCESS, TileXRMoonEPDType, - TileXRMoonEPDispatchArgsV1, + TileXRMoonEPCombineArgsV1, + TileXRMoonEPDispatchArgsV2, TileXRMoonEPPlanV1, TileXRMoonEPPlanningArgsV1, TileXRMoonEPPrefetchWeightArgsV1, @@ -190,6 +190,13 @@ def __init__( raise ValueError(f"rank must be in [0, {world_size}), got {rank}") self.rank = int(rank) self.world_size = int(world_size) + combine_version = os.environ.get("TILEXR_MOONEP_COMBINE_VERSION", "2").strip() + if combine_version not in ("1", "2"): + raise ValueError( + "TILEXR_MOONEP_COMBINE_VERSION must be 1 or 2, " + f"got {combine_version!r}" + ) + self.combine_version = int(combine_version) self.install_prefix = _resolve_install_prefix(install_prefix) self._closed = False self._comm = ctypes.c_void_p() @@ -220,15 +227,21 @@ def __init__( self.install_prefix, paths.get("moonep"), ) - combine_v2_path = _resolve_library( - ("libtilexr-moonep-combine-v2.so.2", "libtilexr-moonep-combine-v2.so"), - "TILEXR_MOONEP_COMBINE_V2_LIB", - self.install_prefix, - paths.get("combine_v2"), - ) + combine_v2_path = None + if self.combine_version == 2: + combine_v2_path = _resolve_library( + ("libtilexr-moonep-combine-v2.so.2", "libtilexr-moonep-combine-v2.so"), + "TILEXR_MOONEP_COMBINE_V2_LIB", + self.install_prefix, + paths.get("combine_v2"), + ) self._comm_lib = cdll_loader(comm_path, mode=ctypes.RTLD_GLOBAL) self._planner_lib = cdll_loader(planner_path, mode=ctypes.RTLD_GLOBAL) - self._combine_v2_lib = cdll_loader(combine_v2_path, mode=ctypes.RTLD_GLOBAL) + self._combine_v2_lib = ( + cdll_loader(combine_v2_path, mode=ctypes.RTLD_GLOBAL) + if combine_v2_path is not None + else None + ) self._moonep_lib = cdll_loader(moonep_path, mode=ctypes.RTLD_GLOBAL) self._configure_symbols() ret = self._comm_lib.TileXRCommInitRankWithSharedQpDomain( @@ -292,7 +305,7 @@ def _configure_symbols(self) -> None: ctypes.POINTER(ctypes.c_int64), ] self._moonep_lib.TileXRMoonEpPlanningGetWorkspaceSizeV1.restype = ctypes.c_int - self._moonep_lib.TileXRMoonEpDispatchGetWorkspaceSizeV1.argtypes = [ + self._moonep_lib.TileXRMoonEpDispatchGetWorkspaceSizeV2.argtypes = [ ctypes.c_void_p, ctypes.c_int64, ctypes.c_int64, @@ -301,7 +314,7 @@ def _configure_symbols(self) -> None: ctypes.POINTER(ctypes.c_uint64), ctypes.POINTER(ctypes.c_uint64), ] - self._moonep_lib.TileXRMoonEpDispatchGetWorkspaceSizeV1.restype = ctypes.c_int + self._moonep_lib.TileXRMoonEpDispatchGetWorkspaceSizeV2.restype = ctypes.c_int self._planner_lib.TileXRMoonEpPlannerGetDstLocalOffsetV3.argtypes = [ ctypes.c_void_p, ctypes.c_int64, @@ -312,45 +325,50 @@ def _configure_symbols(self) -> None: ctypes.POINTER(ctypes.c_uint64), ] self._planner_lib.TileXRMoonEpPlannerGetDstLocalOffsetV3.restype = ctypes.c_int - self._combine_v2_lib.TileXRMoonEpCombineGetWorkspaceSizeV2.argtypes = [ - ctypes.c_int64, - ctypes.c_int64, - ctypes.c_int64, - ctypes.c_int64, - ctypes.c_uint32, - ctypes.POINTER(ctypes.c_uint64), - ctypes.POINTER(ctypes.c_uint64), - ctypes.POINTER(ctypes.c_uint64), - ctypes.POINTER(ctypes.c_uint64), - ] - self._combine_v2_lib.TileXRMoonEpCombineGetWorkspaceSizeV2.restype = ctypes.c_int - self._combine_v2_lib.TileXRMoonEpCombineStageV2.argtypes = [ - ctypes.c_void_p, - ctypes.c_uint64, - ctypes.c_void_p, - ctypes.c_void_p, - ctypes.c_int64, - ctypes.c_int64, - ctypes.c_int64, - ctypes.c_int64, - ctypes.c_uint32, - ctypes.c_void_p, - ctypes.c_void_p, - ctypes.c_void_p, - ctypes.c_void_p, - ctypes.c_uint32, - ctypes.c_void_p, - ] - self._combine_v2_lib.TileXRMoonEpCombineStageV2.restype = ctypes.c_int + if self._combine_v2_lib is not None: + self._combine_v2_lib.TileXRMoonEpCombineGetWorkspaceSizeV2.argtypes = [ + ctypes.c_int64, + ctypes.c_int64, + ctypes.c_int64, + ctypes.c_int64, + ctypes.c_uint32, + ctypes.POINTER(ctypes.c_uint64), + ctypes.POINTER(ctypes.c_uint64), + ctypes.POINTER(ctypes.c_uint64), + ctypes.POINTER(ctypes.c_uint64), + ] + self._combine_v2_lib.TileXRMoonEpCombineGetWorkspaceSizeV2.restype = ctypes.c_int + self._combine_v2_lib.TileXRMoonEpCombineStageV2.argtypes = [ + ctypes.c_void_p, + ctypes.c_uint64, + ctypes.c_void_p, + ctypes.c_void_p, + ctypes.c_int64, + ctypes.c_int64, + ctypes.c_int64, + ctypes.c_int64, + ctypes.c_uint32, + ctypes.c_void_p, + ctypes.c_void_p, + ctypes.c_void_p, + ctypes.c_void_p, + ctypes.c_uint32, + ctypes.c_void_p, + ] + self._combine_v2_lib.TileXRMoonEpCombineStageV2.restype = ctypes.c_int symbols = ( ("TileXRMoonEpPlanningV1", TileXRMoonEPPlanningArgsV1), - ("TileXRMoonEpDispatchV1", TileXRMoonEPDispatchArgsV1), + ("TileXRMoonEpDispatchV2", TileXRMoonEPDispatchArgsV2), ("TileXRMoonEpPrefetchWeightV1", TileXRMoonEPPrefetchWeightArgsV1), ) for name, args_type in symbols: function = getattr(self._moonep_lib, name) function.argtypes = [ctypes.POINTER(args_type), ctypes.c_void_p] function.restype = ctypes.c_int + self._moonep_lib.TileXRMoonEpCombineV1.argtypes = [ + ctypes.POINTER(TileXRMoonEPCombineArgsV1), ctypes.c_void_p + ] + self._moonep_lib.TileXRMoonEpCombineV1.restype = ctypes.c_int self._moonep_lib.TileXRMoonEpReduceGradGetWorkspaceSizeV2.argtypes = [ ctypes.POINTER(TileXRMoonEPReduceGradWorkspaceQueryV2), ctypes.POINTER(TileXRMoonEPReduceGradWorkspaceInfoV2), @@ -492,6 +510,8 @@ def planning_dst_local_offset(self, context) -> int: return int(dst_local_offset.value) def _combine_v2_workspace_size(self, context, h: int, dtype: int) -> int: + if self._combine_v2_lib is None: + return 0 workspace_bytes = ctypes.c_uint64() profile_offset = ctypes.c_uint64() output_epoch0_offset = ctypes.c_uint64() @@ -513,13 +533,13 @@ def _combine_v2_workspace_size(self, context, h: int, dtype: int) -> int: def dispatch_workspace_size(self, context) -> tuple[int, int]: workspace_bytes = ctypes.c_uint64() workspace_alignment = ctypes.c_uint64() - # The V1 query has no NvS argument. Rounding NvS to complete top-k rows + # The V2 query has no NvS argument. Rounding NvS to complete top-k rows # produces a conservative layout when token padding makes NvS > S*K. query_s = max( int(context.tokens_per_rank), (int(context.nv_s) + int(context.topk) - 1) // int(context.topk), ) - ret = self._moonep_lib.TileXRMoonEpDispatchGetWorkspaceSizeV1( + ret = self._moonep_lib.TileXRMoonEpDispatchGetWorkspaceSizeV2( void_p(self.comm_ptr), ctypes.c_int64(query_s), ctypes.c_int64(context.topk), @@ -528,13 +548,16 @@ def dispatch_workspace_size(self, context) -> tuple[int, int]: ctypes.byref(workspace_bytes), ctypes.byref(workspace_alignment), ) - self._check("TileXRMoonEpDispatchGetWorkspaceSizeV1", ret) - hidden_combine_bytes = self._combine_v2_workspace_size( - context, int(context.hidden_size), dtype_code(context.dtype) - ) - weight_combine_bytes = self._combine_v2_workspace_size( - context, 1, int(TileXRMoonEPDType.FLOAT32) - ) + self._check("TileXRMoonEpDispatchGetWorkspaceSizeV2", ret) + hidden_combine_bytes = 0 + weight_combine_bytes = 0 + if self.combine_version == 2: + hidden_combine_bytes = self._combine_v2_workspace_size( + context, int(context.hidden_size), dtype_code(context.dtype) + ) + weight_combine_bytes = self._combine_v2_workspace_size( + context, 1, int(TileXRMoonEPDType.FLOAT32) + ) alignment = max(int(workspace_alignment.value), _UDMA_REGISTRATION_ALIGNMENT) required_bytes = max( int(workspace_bytes.value), hidden_combine_bytes, weight_combine_bytes @@ -643,10 +666,10 @@ def dispatch( raise ValueError( "route_weights and output_route_weights must both be provided or both be None" ) - if (registered_workspace is None) != (registered_workspace_bytes == 0): - raise ValueError( - "registered_workspace and registered_workspace_bytes must be provided together" - ) + if registered_workspace is None or int(registered_workspace_bytes) <= 0: + raise ValueError("Dispatch V2 requires the registered Dispatch workspace") + if build_dedup: + raise ValueError("Dispatch V2 does not support the legacy build_dedup flag") plan_v1 = self._plan_v1(context, plan) hidden_sh = make_tensor_v1(input_tensor) hidden_nvsh = make_tensor_v1(output_tensor) @@ -654,7 +677,7 @@ def dispatch( weights_nvs = ( make_tensor_v1(output_route_weights) if output_route_weights is not None else None ) - args = initialize_struct(TileXRMoonEPDispatchArgsV1()) + args = initialize_struct(TileXRMoonEPDispatchArgsV2()) args.comm = void_p(self.comm_ptr) args.plan = ctypes.pointer(plan_v1) args.hiddenSh = ctypes.pointer(hidden_sh) @@ -663,17 +686,13 @@ def dispatch( args.routeWeightsNvs = ( ctypes.pointer(weights_nvs) if weights_nvs is not None else None ) - args.flags = ( - TILEXR_MOONEP_FLAG_BUILD_DEDUP - if build_dedup and registered_workspace is None - else TILEXR_MOONEP_FLAG_NONE - ) + args.flags = TILEXR_MOONEP_FLAG_NONE args.registeredWorkspace = void_p(registered_workspace) args.registeredWorkspaceBytes = int(registered_workspace_bytes) - ret = self._moonep_lib.TileXRMoonEpDispatchV1( + ret = self._moonep_lib.TileXRMoonEpDispatchV2( ctypes.byref(args), void_p(stream_ptr) ) - self._check("TileXRMoonEpDispatchV1", ret) + self._check("TileXRMoonEpDispatchV2", ret) def prefetch_weight(self, context, plan, projections, stream_ptr: int) -> None: plan_v1 = self._plan_v1(context, plan) @@ -727,9 +746,7 @@ def combine( "route_weights and output_route_weights must both be provided or both be None" ) if int(flags) != TILEXR_MOONEP_FLAG_NONE: - raise ValueError("Combine V2 does not support V1 publish/consume flags") - if registered_workspace is None or int(registered_workspace_bytes) <= 0: - raise ValueError("Combine V2 requires the registered Dispatch workspace") + raise ValueError("MoonEP Combine does not support publish/consume flags") dst_local_offset = int(plan.dst_local_offset) planner_workspace_bytes = tensor_nbytes(plan.workspace) dst_local_bytes = int(context.nv_s) * ctypes.sizeof(ctypes.c_int32) @@ -744,6 +761,36 @@ def combine( f"workspace={planner_workspace_bytes}" ) dst_local_ptr = int(plan.workspace.data_ptr()) + dst_local_offset + if self.combine_version == 1: + plan_v1 = self._plan_v1(context, plan) + hidden_nvsh = make_tensor_v1(input_tensor) + hidden_sh = make_tensor_v1(output_tensor) + weights_nvs = make_tensor_v1(route_weights) if route_weights is not None else None + weights_sk = ( + make_tensor_v1(output_route_weights) + if output_route_weights is not None + else None + ) + args = initialize_struct(TileXRMoonEPCombineArgsV1()) + args.comm = void_p(self.comm_ptr) + args.plan = ctypes.pointer(plan_v1) + args.dstLocal = void_p(dst_local_ptr) + args.hiddenNvsh = ctypes.pointer(hidden_nvsh) + args.routeWeightsNvs = ( + ctypes.pointer(weights_nvs) if weights_nvs is not None else None + ) + args.hiddenSh = ctypes.pointer(hidden_sh) + args.routeWeightsSk = ( + ctypes.pointer(weights_sk) if weights_sk is not None else None + ) + args.flags = TILEXR_MOONEP_FLAG_NONE + ret = self._moonep_lib.TileXRMoonEpCombineV1( + ctypes.byref(args), void_p(stream_ptr) + ) + self._check("TileXRMoonEpCombineV1", ret) + return + if registered_workspace is None or int(registered_workspace_bytes) <= 0: + raise ValueError("Combine V2 requires the registered Dispatch workspace") ret = self._combine_v2_lib.TileXRMoonEpCombineStageV2( void_p(registered_workspace), ctypes.c_uint64(registered_workspace_bytes), diff --git a/integrations/moonep_torch/tilexr_moonep/torch_api.py b/integrations/moonep_torch/tilexr_moonep/torch_api.py index a1ce931..f422330 100644 --- a/integrations/moonep_torch/tilexr_moonep/torch_api.py +++ b/integrations/moonep_torch/tilexr_moonep/torch_api.py @@ -953,7 +953,8 @@ def combine( else None ) self._retain(plan, hidden_nvsh, hidden_sh, route_weights_nvs, route_weights_sk) - self._trace_stage("combine_v2_launch_begin") + combine_version = int(getattr(self.runtime, "combine_version", 2)) + self._trace_stage(f"combine_v{combine_version}_launch_begin") self.runtime.combine( c, plan, @@ -967,7 +968,9 @@ def combine( registered_workspace=self.context.dispatch_workspace[0], registered_workspace_bytes=self.context.dispatch_workspace[1], ) - self._trace_stage("combine_v2_launch_end") + self._trace_stage(f"combine_v{combine_version}_launch_end") + if combine_version == 1: + self._expect_status(plan, 0) event = self._record_event() if async_finish else None return hidden_sh, route_weights_sk, event diff --git a/src/include/tilexr_moonep.h b/src/include/tilexr_moonep.h index a64ad78..f89630b 100644 --- a/src/include/tilexr_moonep.h +++ b/src/include/tilexr_moonep.h @@ -115,6 +115,12 @@ typedef struct TileXRMoonEpDispatchArgsV1 { uint64_t registeredWorkspaceBytes; } TileXRMoonEpDispatchArgsV1; +/* + * V2 preserves the tensor and plan descriptors while requiring the registered + * workspace path. It never falls back to the legacy peer-memory kernel. + */ +typedef TileXRMoonEpDispatchArgsV1 TileXRMoonEpDispatchArgsV2; + typedef struct TileXRMoonEpPrefetchWeightArgsV1 { uint32_t structSize; uint32_t abiVersion; @@ -131,6 +137,7 @@ typedef struct TileXRMoonEpCombineArgsV1 { uint32_t abiVersion; TileXRCommPtr comm; const TileXRMoonEpPlanV1 *plan; + const int32_t *dstLocal; const TileXRMoonEpTensorV1 *hiddenNvsh; const TileXRMoonEpTensorV1 *routeWeightsNvs; TileXRMoonEpTensorV1 *hiddenSh; @@ -207,6 +214,12 @@ int TileXRMoonEpDispatchGetWorkspaceSizeV1(TileXRCommPtr comm, int64_t s, int TileXRMoonEpDispatchV1(const TileXRMoonEpDispatchArgsV1 *args, aclrtStream stream); +int TileXRMoonEpDispatchGetWorkspaceSizeV2(TileXRCommPtr comm, int64_t s, + int64_t k, int64_t h, uint32_t hiddenDtype, uint64_t *workspaceBytes, + uint64_t *workspaceAlignment); + +int TileXRMoonEpDispatchV2(const TileXRMoonEpDispatchArgsV2 *args, aclrtStream stream); + int TileXRMoonEpPrefetchWeightV1(const TileXRMoonEpPrefetchWeightArgsV1 *args, aclrtStream stream); diff --git a/src/moonep/combine/CMakeLists.txt b/src/moonep/combine/CMakeLists.txt index 9875a03..9ad3472 100644 --- a/src/moonep/combine/CMakeLists.txt +++ b/src/moonep/combine/CMakeLists.txt @@ -59,6 +59,7 @@ tilexr_add_moonep_kernel(tilexr_moonep_combine_kernel DEPENDS "${CMAKE_CURRENT_SOURCE_DIR}/common/combine_common.h" "${CMAKE_SOURCE_DIR}/src/moonep/common/moonep_peer_window.h" + "${CMAKE_SOURCE_DIR}/src/moonep/common/moonep_combine_schedule.h" ) add_library(tilexr-moonep-combine SHARED diff --git a/src/moonep/combine/host/combine_host.cpp b/src/moonep/combine/host/combine_host.cpp index 1096a49..18f139b 100644 --- a/src/moonep/combine/host/combine_host.cpp +++ b/src/moonep/combine/host/combine_host.cpp @@ -14,7 +14,8 @@ int TileXRMoonEpPrepareCombineLaunch(const TileXRMoonEpCombineArgsV1 *args, *context = CombineLaunchContext {}; if (args == nullptr || args->structSize < sizeof(*args) || args->abiVersion != TILEXR_MOONEP_ABI_VERSION_V1 || args->comm == nullptr || - args->plan == nullptr || args->hiddenNvsh == nullptr || args->hiddenSh == nullptr || + args->plan == nullptr || args->dstLocal == nullptr || + args->hiddenNvsh == nullptr || args->hiddenSh == nullptr || stream == nullptr) { return TILEXR_MOONEP_ERROR_INVALID_ARGUMENT; } @@ -40,6 +41,7 @@ int TileXRMoonEpPrepareCombineLaunch(const TileXRMoonEpCombineArgsV1 *args, } params->comm = args->comm; + params->dstLocal = args->dstLocal; params->dst = static_cast(args->plan->dst); params->dupGroups = static_cast(args->plan->dupGroups); params->dupLoffs = static_cast(args->plan->dupLoffs); diff --git a/src/moonep/combine/host/combine_host.h b/src/moonep/combine/host/combine_host.h index 912693d..28363e4 100644 --- a/src/moonep/combine/host/combine_host.h +++ b/src/moonep/combine/host/combine_host.h @@ -11,6 +11,7 @@ namespace TileXRMoonEp { struct CombineParams { TileXRCommPtr comm = nullptr; + const int32_t *dstLocal = nullptr; const int32_t *dst = nullptr; const int32_t *dupGroups = nullptr; const int32_t *dupLoffs = nullptr; diff --git a/src/moonep/combine/host/combine_launch.cpp b/src/moonep/combine/host/combine_launch.cpp index f3ca1fa..acbb777 100644 --- a/src/moonep/combine/host/combine_launch.cpp +++ b/src/moonep/combine/host/combine_launch.cpp @@ -22,6 +22,7 @@ int TileXRMoonEpLaunchCombineKernel( { struct CombineKernelArgs { GM_ADDR commArgs; + GM_ADDR dstLocal; GM_ADDR dst; GM_ADDR dupGroups; GM_ADDR dupLoffs; @@ -39,14 +40,22 @@ int TileXRMoonEpLaunchCombineKernel( int64_t hiddenChunkBytes; int64_t hiddenChunkStride; int64_t chunkCount; + uint64_t sourceHiddenOffset; + uint64_t receiveHiddenOffset; uint64_t hiddenPayloadBytes; - uint64_t routeWeightsOffset; + uint64_t sourceWeightsOffset; + uint64_t receiveWeightsOffset; uint64_t routeWeightsBytes; + uint64_t duplicateMaskOffset; + uint64_t doneOffset; + uint64_t coreStatusOffset; + uint64_t windowBytes; uint64_t waitIterations; uint64_t flags; int64_t magic; } args { static_cast(context.devArgs), + reinterpret_cast(const_cast(params.dstLocal)), reinterpret_cast(const_cast(params.dst)), reinterpret_cast(const_cast(params.dupGroups)), reinterpret_cast(const_cast(params.dupLoffs)), @@ -60,8 +69,12 @@ int TileXRMoonEpLaunchCombineKernel( static_cast(context.layout.hiddenRowBytes), static_cast(context.layout.hiddenChunkBytes), static_cast(context.layout.hiddenChunkStride), context.layout.chunkCount, - context.layout.hiddenPayloadBytes, context.layout.routeWeightsOffset, - context.layout.routeWeightsBytes, context.waitIterations, params.flags, + context.layout.sourceHiddenOffset, context.layout.receiveHiddenOffset, + context.layout.hiddenPayloadBytes, context.layout.sourceWeightsOffset, + context.layout.receiveWeightsOffset, context.layout.routeWeightsBytes, + context.layout.duplicateMaskOffset, context.layout.doneOffset, + context.layout.coreStatusOffset, context.layout.windowBytes, + context.waitIterations, params.flags, context.magic }; diff --git a/src/moonep/combine/host/combine_layout.cpp b/src/moonep/combine/host/combine_layout.cpp index eb3e100..9db0f2b 100644 --- a/src/moonep/combine/host/combine_layout.cpp +++ b/src/moonep/combine/host/combine_layout.cpp @@ -1,5 +1,8 @@ #include "combine_layout.h" +#include + +#include "moonep_combine_schedule.h" #include "comm_args.h" #include "moonep_peer_window.h" #include "moonep_stage_layout.h" @@ -31,6 +34,69 @@ bool WeightsValid(const TileXRMoonEpTensorV1 *input, TileXRMoonEpTensorV1 *outpu output->shape[1] == k; } +bool BuildRegions(uint64_t nvS, uint64_t hiddenStride, uint64_t weightsBytes, + uint64_t world, uint64_t activeCores, CombineLayout *layout) +{ + if (layout == nullptr || nvS == 0 || hiddenStride == 0) { + return false; + } + CombineLayout next = *layout; + uint64_t hiddenBytes = 0; + uint64_t doneBytes = 0; + uint64_t statusBytes = 0; + uint64_t maskBytes = 0; + uint64_t cursor = 0; + if (!Layout::CheckedMul(nvS, hiddenStride, &hiddenBytes) || + !Layout::CheckedMul(nvS, sizeof(int32_t), &maskBytes) || + !Layout::CheckedMul(world, kMoonEpCombineV2TokenStrideBytes, &doneBytes) || + !Layout::CheckedMul(activeCores, kMoonEpCombineV2TokenStrideBytes, &statusBytes) || + !Layout::AppendRegion(hiddenBytes, kMoonEpStageAlignment, &cursor, + &next.sourceHiddenOffset) || + !Layout::AppendRegion(hiddenBytes, kMoonEpStageAlignment, &cursor, + &next.receiveHiddenOffset) || + !Layout::AppendRegion(weightsBytes, kMoonEpStageAlignment, &cursor, + &next.sourceWeightsOffset) || + !Layout::AppendRegion(weightsBytes, kMoonEpStageAlignment, &cursor, + &next.receiveWeightsOffset) || + !Layout::AppendRegion(maskBytes, kMoonEpStageAlignment, &cursor, + &next.duplicateMaskOffset) || + !Layout::AppendRegion(doneBytes, kMoonEpCombineV2TokenStrideBytes, &cursor, + &next.doneOffset) || + !Layout::AppendRegion(statusBytes, kMoonEpCombineV2TokenStrideBytes, &cursor, + &next.coreStatusOffset) || + cursor > static_cast(TileXR::IPC_BUFF_MAX_SIZE)) { + return false; + } + next.hiddenPayloadBytes = hiddenBytes; + next.routeWeightsBytes = weightsBytes; + next.duplicateMaskBytes = maskBytes; + next.doneBytes = doneBytes; + next.coreStatusBytes = statusBytes; + next.windowBytes = cursor; + *layout = next; + return true; +} + +uint64_t MaxHiddenStride(uint64_t nvS, uint64_t weightsBytes, + uint64_t world, uint64_t activeCores) +{ + const uint64_t capacity = static_cast(TileXR::IPC_BUFF_MAX_SIZE); + const uint64_t maxUnits = capacity / kMoonEpStageAlignment / nvS / 2U; + uint64_t low = 0; + uint64_t high = maxUnits; + while (low < high) { + const uint64_t mid = low + (high - low + 1U) / 2U; + CombineLayout candidate {}; + if (BuildRegions(nvS, mid * kMoonEpStageAlignment, weightsBytes, + world, activeCores, &candidate)) { + low = mid; + } else { + high = mid - 1U; + } + } + return low * kMoonEpStageAlignment; +} + } // namespace int TileXRMoonEpBuildCombineLayout(int64_t commRank, int64_t commWorld, @@ -42,62 +108,46 @@ int TileXRMoonEpBuildCombineLayout(int64_t commRank, int64_t commWorld, return TILEXR_MOONEP_ERROR_INVALID_ARGUMENT; } *layout = CombineLayout {}; - const uint64_t splitFlags = TILEXR_MOONEP_FLAG_COMBINE_PUBLISH_ONLY | - TILEXR_MOONEP_FLAG_COMBINE_CONSUME_ONLY; - const uint64_t allowedFlags = TILEXR_MOONEP_FLAG_SKIP_INTER_RANK_SYNC | splitFlags; - if ((flags & ~allowedFlags) != 0) { - return TILEXR_MOONEP_ERROR_INVALID_ARGUMENT; - } - if ((flags & splitFlags) == splitFlags) { - return TILEXR_MOONEP_ERROR_INVALID_ARGUMENT; - } int64_t s = 0; uint64_t hiddenRowBytes = 0; - if (!Layout::PlanValid(commRank, commWorld, plan, &s) || + if (flags != TILEXR_MOONEP_FLAG_NONE || + !MoonEpCombineV2RankSizeSupported(static_cast(commWorld)) || + !Layout::PlanValid(commRank, commWorld, plan, &s) || !HiddenValid(hiddenNvsh, hiddenSh, s, plan->nvS, &hiddenRowBytes) || + hiddenRowBytes > static_cast(std::numeric_limits::max()) || !WeightsValid(routeWeightsNvs, routeWeightsSk, s, plan->k, plan->nvS)) { return TILEXR_MOONEP_ERROR_INVALID_ARGUMENT; } + + const uint64_t world = static_cast(commWorld); const uint64_t nvS = static_cast(plan->nvS); + const uint64_t activeCores = MoonEpCombineV2ActiveCoreCount( + static_cast(commWorld)); uint64_t routeWeightsBytes = 0; if (routeWeightsNvs != nullptr && !Layout::CheckedMul(nvS, sizeof(float), &routeWeightsBytes)) { return TILEXR_MOONEP_ERROR_INVALID_ARGUMENT; } - if (routeWeightsBytes > static_cast(TileXR::IPC_BUFF_MAX_SIZE) - - (kMoonEpStageAlignment - 1)) { - return TILEXR_MOONEP_ERROR_INVALID_ARGUMENT; - } - const uint64_t available = static_cast(TileXR::IPC_BUFF_MAX_SIZE) - - routeWeightsBytes - (kMoonEpStageAlignment - 1); - const uint64_t maxStride = (available / nvS / kMoonEpStageAlignment) * - kMoonEpStageAlignment; + const uint64_t maxStride = MaxHiddenStride( + nvS, routeWeightsBytes, world, activeCores); if (maxStride == 0) { return TILEXR_MOONEP_ERROR_INVALID_ARGUMENT; } const uint64_t hiddenChunkBytes = hiddenRowBytes < maxStride ? hiddenRowBytes : maxStride; const uint64_t hiddenChunkStride = Layout::AlignUp(hiddenChunkBytes, kMoonEpStageAlignment); - uint64_t hiddenPayloadBytes = 0; - uint64_t cursor = 0; - CombineLayout next {}; - if (hiddenChunkStride == std::numeric_limits::max() || - !Layout::CheckedMul(nvS, hiddenChunkStride, &hiddenPayloadBytes) || - !Layout::AppendRegion(hiddenPayloadBytes, kMoonEpStageAlignment, &cursor, - &next.hiddenPayloadBytes) || next.hiddenPayloadBytes != 0 || - !Layout::AppendRegion(routeWeightsBytes, kMoonEpStageAlignment, &cursor, - &next.routeWeightsOffset) || - cursor > static_cast(TileXR::IPC_BUFF_MAX_SIZE)) { + if (hiddenChunkStride == std::numeric_limits::max()) { return TILEXR_MOONEP_ERROR_INVALID_ARGUMENT; } - const uint64_t chunkCount = (hiddenRowBytes - 1) / hiddenChunkBytes + 1; - if ((flags & splitFlags) != 0 && chunkCount != 1) { - return TILEXR_MOONEP_ERROR_NOT_SUPPORTED; - } - if (chunkCount > static_cast(std::numeric_limits::max() - - kMoonEpCombineWindowDrainedStep) / 4U) { + const uint64_t chunkCount = (hiddenRowBytes - 1U) / hiddenChunkBytes + 1U; + if (chunkCount > 0xFFFFFFU) { return TILEXR_MOONEP_ERROR_INVALID_ARGUMENT; } + CombineLayout next {}; + if (!BuildRegions(nvS, hiddenChunkStride, routeWeightsBytes, + world, activeCores, &next)) { + return TILEXR_MOONEP_ERROR_INVALID_ARGUMENT; + } next.rank = commRank; next.world = commWorld; next.s = s; @@ -105,15 +155,15 @@ int TileXRMoonEpBuildCombineLayout(int64_t commRank, int64_t commWorld, next.n = plan->n; next.nvS = plan->nvS; next.hiddenSize = hiddenNvsh->shape[1]; - next.blockDim = kMoonEpStageAivBlockCount; + next.blockDim = static_cast(activeCores); + next.stepCount = static_cast(MoonEpCombineV2StepCount( + static_cast(commWorld))); + next.sourcesPerCore = commWorld / next.blockDim; next.chunkCount = static_cast(chunkCount); next.flags = flags; next.hiddenRowBytes = hiddenRowBytes; next.hiddenChunkBytes = hiddenChunkBytes; next.hiddenChunkStride = hiddenChunkStride; - next.hiddenPayloadBytes = hiddenPayloadBytes; - next.routeWeightsBytes = routeWeightsBytes; - next.windowBytes = cursor; *layout = next; return TILEXR_MOONEP_SUCCESS; } diff --git a/src/moonep/combine/host/combine_layout.h b/src/moonep/combine/host/combine_layout.h index 7297ac6..2e2def8 100644 --- a/src/moonep/combine/host/combine_layout.h +++ b/src/moonep/combine/host/combine_layout.h @@ -16,14 +16,25 @@ struct CombineLayout { int64_t nvS = 0; int64_t hiddenSize = 0; int64_t blockDim = 0; + int64_t stepCount = 0; + int64_t sourcesPerCore = 0; int64_t chunkCount = 0; uint64_t flags = 0; uint64_t hiddenRowBytes = 0; uint64_t hiddenChunkBytes = 0; uint64_t hiddenChunkStride = 0; + uint64_t sourceHiddenOffset = 0; uint64_t hiddenPayloadBytes = 0; - uint64_t routeWeightsOffset = 0; + uint64_t receiveHiddenOffset = 0; + uint64_t sourceWeightsOffset = 0; + uint64_t receiveWeightsOffset = 0; uint64_t routeWeightsBytes = 0; + uint64_t duplicateMaskOffset = 0; + uint64_t duplicateMaskBytes = 0; + uint64_t doneOffset = 0; + uint64_t doneBytes = 0; + uint64_t coreStatusOffset = 0; + uint64_t coreStatusBytes = 0; uint64_t windowBytes = 0; }; diff --git a/src/moonep/combine/kernels/tilexr_moonep_combine_kernel.cpp b/src/moonep/combine/kernels/tilexr_moonep_combine_kernel.cpp index 23b1644..9a640d1 100644 --- a/src/moonep/combine/kernels/tilexr_moonep_combine_kernel.cpp +++ b/src/moonep/combine/kernels/tilexr_moonep_combine_kernel.cpp @@ -4,12 +4,16 @@ #include "comm_args.h" #include "combine_common.h" +#include "moonep_combine_schedule.h" #include "tilexr_sync.h" -#include "tilexr_udma.h" namespace TileXRMoonEp { namespace Kernel { +constexpr uint32_t kChunkReadyStepBase = 1U << 10U; +constexpr uint32_t kChunkDrainedStepBase = 1U << 25U; +constexpr uint32_t kMaxChunkCount = 0xFFFFFFU; + __aicore__ inline int64_t MinInt64(int64_t lhs, int64_t rhs) { return lhs < rhs ? lhs : rhs; @@ -26,12 +30,12 @@ __aicore__ inline void CopyBytesGmToGm(GM_ADDR dstAddr, GM_ADDR srcAddr, AscendC::GlobalTensor dst; src.SetGlobalBuffer(reinterpret_cast<__gm__ uint8_t *>(srcAddr), bytes); dst.SetGlobalBuffer(reinterpret_cast<__gm__ uint8_t *>(dstAddr), bytes); - - for (int64_t copied = 0; copied < bytes; copied += kMoonEpCombineFloatScratchBytes) { - const int64_t tileBytes = MinInt64(bytes - copied, kMoonEpCombineFloatScratchBytes); + for (int64_t copied = 0; copied < bytes; + copied += kMoonEpCombineFloatScratchBytes) { + const int64_t tileBytes = MinInt64( + bytes - copied, kMoonEpCombineFloatScratchBytes); AscendC::DataCopyExtParams params { - 1, static_cast(tileBytes), 0, 0, 0 - }; + 1, static_cast(tileBytes), 0, 0, 0}; AscendC::DataCopyPadExtParams pad {false, 0, 0, 0}; AscendC::DataCopyPad(local, src[copied], params, pad); AscendC::SetFlag(EVENT_ID0); @@ -45,16 +49,22 @@ __aicore__ inline void CopyBytesGmToGm(GM_ADDR dstAddr, GM_ADDR srcAddr, class CombineKernel { public: - __aicore__ inline void Init(GM_ADDR commArgs, GM_ADDR dst, GM_ADDR dupGroups, - GM_ADDR dupLoffs, GM_ADDR dupCounts, GM_ADDR hiddenNvsh, - GM_ADDR routeWeightsNvs, GM_ADDR hiddenSh, GM_ADDR routeWeightsSk, - GM_ADDR status, int64_t s, int64_t k, int64_t n, int64_t nvS, - int64_t hiddenRowBytes, int64_t hiddenChunkBytes, int64_t hiddenChunkStride, - int64_t chunkCount, uint64_t hiddenPayloadBytes, uint64_t routeWeightsOffset, - uint64_t routeWeightsBytes, uint64_t waitIterations, uint64_t flags, + __aicore__ inline void Init(GM_ADDR commArgs, GM_ADDR dstLocal, GM_ADDR dst, + GM_ADDR dupGroups, GM_ADDR dupLoffs, GM_ADDR dupCounts, + GM_ADDR hiddenNvsh, GM_ADDR routeWeightsNvs, GM_ADDR hiddenSh, + GM_ADDR routeWeightsSk, GM_ADDR status, int64_t s, int64_t k, + int64_t n, int64_t nvS, int64_t hiddenRowBytes, + int64_t hiddenChunkBytes, int64_t hiddenChunkStride, + int64_t chunkCount, uint64_t sourceHiddenOffset, + uint64_t receiveHiddenOffset, uint64_t hiddenPayloadBytes, + uint64_t sourceWeightsOffset, uint64_t receiveWeightsOffset, + uint64_t routeWeightsBytes, uint64_t duplicateMaskOffset, + uint64_t doneOffset, uint64_t coreStatusOffset, + uint64_t windowBytes, uint64_t waitIterations, uint64_t flags, int64_t magic) { args_ = reinterpret_cast<__gm__ TileXR::CommArgs *>(commArgs); + dstLocalAddr_ = dstLocal; dstAddr_ = dst; dupGroupsAddr_ = dupGroups; dupLoffsAddr_ = dupLoffs; @@ -72,214 +82,146 @@ class CombineKernel { hiddenChunkBytes_ = hiddenChunkBytes; hiddenChunkStride_ = hiddenChunkStride; chunkCount_ = chunkCount; + sourceHiddenOffset_ = sourceHiddenOffset; + receiveHiddenOffset_ = receiveHiddenOffset; hiddenPayloadBytes_ = hiddenPayloadBytes; - routeWeightsOffset_ = routeWeightsOffset; + sourceWeightsOffset_ = sourceWeightsOffset; + receiveWeightsOffset_ = receiveWeightsOffset; routeWeightsBytes_ = routeWeightsBytes; + duplicateMaskOffset_ = duplicateMaskOffset; + doneOffset_ = doneOffset; + coreStatusOffset_ = coreStatusOffset; + windowBytes_ = windowBytes; waitIterations_ = waitIterations; flags_ = flags; magic_ = magic; - if (args_ == nullptr) { return; } rank_ = args_->rank; rankSize_ = args_->rankSize; + core_ = static_cast(AscendC::GetBlockIdx()); + if (!MoonEpCombineV2RankSizeSupported(static_cast(rankSize_))) { + return; + } + activeCoreCount_ = MoonEpCombineV2ActiveCoreCount( + static_cast(rankSize_)); + stepCount_ = MoonEpCombineV2StepCount(static_cast(rankSize_)); + sourcesPerCore_ = static_cast(rankSize_) / activeCoreCount_; for (int32_t peer = 0; peer < rankSize_; ++peer) { shareAddrs_[peer] = args_->peerMems[peer]; } - pipe_.InitBuffer(syncBuf_, kMoonEpSyncUbBytes); - pipe_.InitBuffer(bfloatBuf_, kMoonEpCombineBfloatScratchBytes); - pipe_.InitBuffer(routeBuf_, kMoonEpCombineFloatScratchBytes); - pipe_.InitBuffer(accumulatorBuf_, kMoonEpCombineFloatScratchBytes); - sync_.Init(rank_, rankSize_, shareAddrs_, syncBuf_); - initialized_ = true; + localWindow_ = shareAddrs_[rank_] + TileXR::IPC_DATA_OFFSET; + valid_ = ValidateConfiguration(); } __aicore__ inline void Process() { - if (!Valid() || AscendC::GetBlockIdx() != 0) { + if (!valid_) { return; } - if (IsPublishOnly() || IsConsumeOnly()) { - if (!InitRegisteredWindow()) { - return; - } - } else { - localWindow_ = shareAddrs_[rank_] + TileXR::IPC_DATA_OFFSET; - } - if (IsPublishOnly()) { - RunPublishOnly(); - return; + InitBuffers(); + sync_.Init(rank_, rankSize_, shareAddrs_, syncBuf_); + StoreCoreStatus(0); + if (core_ == 0U) { + StoreInt(statusAddr_, 0); } - if (IsConsumeOnly()) { - RunConsumeOnly(); + AscendC::SyncAll(); + + BuildDuplicateMask(); + ValidateReverseRoutes(); + AscendC::SyncAll(); + if (!CrossRankBarrier(kMoonEpCombineDataReadyStep)) { + Finish(); return; } - for (int64_t chunk = 0; chunk < chunkCount_; ++chunk) { - const int64_t chunkOffset = chunk * hiddenChunkBytes_; + + for (uint32_t chunk = 0U; + chunk < static_cast(chunkCount_); ++chunk) { + const int64_t chunkOffset = + static_cast(chunk) * hiddenChunkBytes_; const int64_t bytesThisChunk = MinInt64( hiddenRowBytes_ - chunkOffset, hiddenChunkBytes_); - PublishLocalInput(chunkOffset, bytesThisChunk, chunk == 0); - if (!PreReduceDuplicates(bytesThisChunk)) { - return; - } - PublishStep(ChunkStep(kMoonEpCombineDataReadyStep, chunk)); - if (!WaitAllPeers(ChunkStep(kMoonEpCombineDataReadyStep, chunk))) { + PrepareChunk(chunkOffset, bytesThisChunk, chunk == 0U); + AscendC::SyncAll(); + PreReduceDuplicates(bytesThisChunk); + AscendC::SyncAll(); + if (!CrossRankBarrier(kChunkReadyStepBase + chunk)) { + Finish(); return; } - if (!ReduceHiddenChunk(chunkOffset, bytesThisChunk)) { - return; + for (uint32_t step = 0U; step < stepCount_; ++step) { + const uint32_t peer = MoonEpCombineV2Peer( + static_cast(rank_), step, core_, + static_cast(rankSize_)); + PushPeerRows(peer, step, chunk, bytesThisChunk, chunk == 0U); } - if (chunk == 0 && routeWeightsBytes_ > 0 && !GatherWeights()) { + WaitInboundDone(chunk); + AscendC::SyncAll(); + if (FirstFailure() != 0) { + CrossRankBarrier(kChunkDrainedStepBase + chunk); + Finish(); return; } - PublishStep(ChunkStep(kMoonEpCombineWindowDrainedStep, chunk)); - if (!WaitAllPeers(ChunkStep(kMoonEpCombineWindowDrainedStep, chunk))) { + ReduceHiddenChunk(chunkOffset, bytesThisChunk); + if (chunk == 0U && routeWeightsBytes_ != 0U) { + CopyReceivedWeights(); + } + AscendC::SyncAll(); + if (!CrossRankBarrier(kChunkDrainedStepBase + chunk)) { + Finish(); return; } } - StoreStatus(kMoonEpCombineStatusSuccess); + Finish(); } private: - __aicore__ inline bool IsPublishOnly() const - { - return (flags_ & kMoonEpFlagCombinePublishOnly) != 0; - } - - __aicore__ inline bool IsConsumeOnly() const - { - return (flags_ & kMoonEpFlagCombineConsumeOnly) != 0; - } - - __aicore__ inline bool InitRegisteredWindow() - { - if (!TileXR::UDMARegistryEnabled(args_)) { - Fail(kMoonEpCombineStatusUdmaInvalid); - return false; - } - const uint64_t weightsEnd = routeWeightsOffset_ + routeWeightsBytes_; - const uint64_t payloadEnd = hiddenPayloadBytes_ > weightsEnd ? - hiddenPayloadBytes_ : weightsEnd; - if (payloadEnd > UINT64_MAX - (kMoonEpStageAlignment - 1U)) { - Fail(kMoonEpCombineStatusUdmaInvalid); - return false; - } - const uint64_t scratchOffset = - ((payloadEnd + kMoonEpStageAlignment - 1U) / kMoonEpStageAlignment) * - kMoonEpStageAlignment; - const uint64_t scratchBytes = hiddenChunkBytes_ > static_cast(sizeof(float)) ? - static_cast(hiddenChunkBytes_) : sizeof(float); - if (scratchOffset > UINT64_MAX - scratchBytes) { - Fail(kMoonEpCombineStatusUdmaInvalid); - return false; - } - const uint64_t requiredBytes = scratchOffset + scratchBytes; - __gm__ TileXR::TileXRUDMARegistry *registry = TileXR::GetUDMARegistry(args_); - if (!TileXR::UDMARegisteredRangeValid(registry, rank_, 0, requiredBytes)) { - Fail(kMoonEpCombineStatusUdmaInvalid); - return false; - } - localWindow_ = TileXR::UDMARegisteredRemoteAddr(registry, rank_, 0); - registeredScratch_ = localWindow_ + scratchOffset; - return true; - } - - __aicore__ inline bool FetchRegisteredPayload( - int64_t peer, uint64_t byteOffset, uint32_t byteCount, GM_ADDR &source) - { - if (peer == rank_) { - source = localWindow_ + static_cast(byteOffset); - return true; - } - __gm__ TileXR::TileXRUDMARegistry *registry = TileXR::GetUDMARegistry(args_); - if (registeredScratch_ == nullptr || - !TileXR::UDMARegisteredRangeValid(registry, static_cast(peer), - byteOffset, byteCount)) { - Fail(kMoonEpCombineStatusUdmaInvalid); - return false; - } - AscendC::LocalTensor wqeScratch = syncBuf_.Get(); - const uint32_t qpCount = TileXR::UDMAQpCount(args_); - if (qpCount == 0U) { - Fail(kMoonEpCombineStatusUdmaInvalid); - return false; - } - const uint32_t qpIdx = qpCount > 1U ? 1U : 0U; - const uint32_t submitStatus = TileXR::UDMAGetNbiOnQp( - args_, wqeScratch, static_cast(peer), qpIdx, - reinterpret_cast<__gm__ uint8_t *>(registeredScratch_), byteOffset, - byteCount); - if (submitStatus != TileXR::TILEXR_UDMA_STATUS_SUCCESS) { - Fail(kMoonEpCombineStatusUdmaSubmitBase + static_cast(peer)); - return false; - } - const uint32_t completionStatus = TileXR::UDMAQuietStatusOnQp( - args_, static_cast(peer), qpIdx); - if (completionStatus != TileXR::TILEXR_UDMA_STATUS_SUCCESS) { - Fail(kMoonEpCombineStatusUdmaCompletionBase + static_cast(peer)); - return false; - } - AscendC::GlobalTensor cache; - cache.SetGlobalBuffer(reinterpret_cast<__gm__ uint64_t *>(0), 1); - AscendC::DataCacheCleanAndInvalid(cache); - source = registeredScratch_; - return true; - } - - __aicore__ inline void RunPublishOnly() - { - PublishLocalInput(0, hiddenChunkBytes_, true); - if (PreReduceDuplicates(hiddenChunkBytes_)) { - StoreStatus(kMoonEpCombineStatusSuccess); - } - } - - __aicore__ inline void RunConsumeOnly() - { - if (!ReduceHiddenChunk(0, hiddenChunkBytes_)) { - return; - } - if (routeWeightsBytes_ > 0 && !GatherWeights()) { - return; - } - StoreStatus(kMoonEpCombineStatusSuccess); - } - - __aicore__ inline bool Valid() const + __aicore__ inline bool ValidateConfiguration() const { - const uint64_t splitFlags = kMoonEpFlagCombinePublishOnly | - kMoonEpFlagCombineConsumeOnly; - const uint64_t allowedFlags = kMoonEpFlagSkipInterRankSync | splitFlags; - uint64_t windowBytes = hiddenPayloadBytes_; - const uint64_t weightsEnd = routeWeightsOffset_ + routeWeightsBytes_; - windowBytes = weightsEnd > windowBytes ? weightsEnd : windowBytes; - return initialized_ && args_ != nullptr && dstAddr_ != nullptr && + const uint64_t hiddenEnd = receiveHiddenOffset_ + hiddenPayloadBytes_; + const uint64_t sourceWeightsEnd = sourceWeightsOffset_ + routeWeightsBytes_; + const uint64_t receiveWeightsEnd = receiveWeightsOffset_ + routeWeightsBytes_; + const uint64_t maskEnd = duplicateMaskOffset_ + + static_cast(nvS_) * sizeof(int32_t); + const uint64_t doneEnd = doneOffset_ + + static_cast(rankSize_) * kMoonEpCombineV2TokenStrideBytes; + const uint64_t statusEnd = coreStatusOffset_ + + static_cast(activeCoreCount_) * + kMoonEpCombineV2TokenStrideBytes; + return dstLocalAddr_ != nullptr && dstAddr_ != nullptr && dupGroupsAddr_ != nullptr && dupLoffsAddr_ != nullptr && dupCountsAddr_ != nullptr && hiddenNvshAddr_ != nullptr && - hiddenShAddr_ != nullptr && statusAddr_ != nullptr && rank_ >= 0 && - rank_ < rankSize_ && s_ > 0 && k_ > 0 && k_ <= 32 && n_ == s_ * k_ && - nvS_ >= n_ && hiddenRowBytes_ > 0 && hiddenChunkBytes_ > 0 && + hiddenShAddr_ != nullptr && statusAddr_ != nullptr && + rank_ >= 0 && rank_ < rankSize_ && core_ < activeCoreCount_ && + AscendC::GetBlockNum() == activeCoreCount_ && s_ > 0 && + k_ > 0 && k_ <= 32 && n_ == s_ * k_ && nvS_ >= n_ && + hiddenRowBytes_ > 0 && hiddenChunkBytes_ > 0 && hiddenChunkStride_ >= hiddenChunkBytes_ && hiddenChunkStride_ % static_cast(kMoonEpStageAlignment) == 0 && - chunkCount_ > 0 && hiddenPayloadBytes_ == - static_cast(nvS_) * static_cast(hiddenChunkStride_) && - windowBytes <= static_cast(TileXR::IPC_BUFF_MAX_SIZE) && - waitIterations_ > 0 && magic_ > 0 && (flags_ & ~allowedFlags) == 0 && - (flags_ & splitFlags) != splitFlags && - ((flags_ & splitFlags) == 0 || chunkCount_ == 1) && - ((routeWeightsBytes_ == 0 && routeWeightsNvsAddr_ == nullptr && + chunkCount_ > 0 && chunkCount_ <= static_cast(kMaxChunkCount) && + hiddenPayloadBytes_ == static_cast(nvS_) * + static_cast(hiddenChunkStride_) && + sourceHiddenOffset_ == 0U && receiveHiddenOffset_ >= hiddenPayloadBytes_ && + hiddenEnd <= windowBytes_ && sourceWeightsEnd <= windowBytes_ && + receiveWeightsEnd <= windowBytes_ && maskEnd <= windowBytes_ && + doneEnd <= windowBytes_ && statusEnd <= windowBytes_ && + windowBytes_ <= static_cast(TileXR::IPC_BUFF_MAX_SIZE) && + waitIterations_ > 0 && magic_ > 0 && magic_ <= INT32_MAX && + flags_ == 0U && + ((routeWeightsBytes_ == 0U && routeWeightsNvsAddr_ == nullptr && routeWeightsSkAddr_ == nullptr) || (routeWeightsBytes_ == static_cast(nvS_) * sizeof(float) && routeWeightsNvsAddr_ != nullptr && routeWeightsSkAddr_ != nullptr)); } - __aicore__ inline int32_t ChunkStep(int32_t base, int64_t chunk) const + __aicore__ inline void InitBuffers() { - return base + static_cast(chunk * 4); + pipe_.InitBuffer(syncBuf_, kMoonEpSyncUbBytes); + pipe_.InitBuffer(bfloatBuf_, kMoonEpCombineBfloatScratchBytes); + pipe_.InitBuffer(routeBuf_, kMoonEpCombineFloatScratchBytes); + pipe_.InitBuffer(accumulatorBuf_, kMoonEpCombineFloatScratchBytes); } __aicore__ inline int32_t LoadInt(GM_ADDR address) @@ -309,107 +251,290 @@ class CombineKernel { AscendC::WaitFlag(EVENT_ID0); } - __aicore__ inline int32_t LoadRoute(int64_t route) + __aicore__ inline void ZeroBytes(GM_ADDR address, int64_t bytes) { - return LoadInt(dstAddr_ + route * static_cast(sizeof(int32_t))); + AscendC::LocalTensor local = routeBuf_.Get(); + for (int64_t offset = 0; offset < bytes; + offset += kMoonEpCombineFloatScratchBytes) { + const int64_t tileBytes = MinInt64( + bytes - offset, kMoonEpCombineFloatScratchBytes); + const int32_t words = static_cast((tileBytes + 3) / 4); + AscendC::Duplicate(local, static_cast(0), words); + AscendC::SetFlag(EVENT_ID0); + AscendC::WaitFlag(EVENT_ID0); + AscendC::GlobalTensor dst; + dst.SetGlobalBuffer(reinterpret_cast<__gm__ uint8_t *>(address + offset), + tileBytes); + AscendC::DataCopyExtParams params { + 1, static_cast(tileBytes), 0, 0, 0}; + AscendC::DataCopyPad(dst, local.ReinterpretCast(), params); + AscendC::SetFlag(EVENT_ID0); + AscendC::WaitFlag(EVENT_ID0); + } + AscendC::PipeBarrier(); } - __aicore__ inline bool DecodeRoute(int32_t encoded, int64_t &peer, int64_t &offset) + __aicore__ inline void StoreCoreStatus(int32_t status) { - const int64_t encoded64 = static_cast(encoded); - const int64_t raw = encoded64 >= 0 ? encoded64 : -encoded64 - 1; - peer = raw / nvS_; - offset = raw % nvS_; - if (raw < 0 || peer < 0 || peer >= rankSize_ || offset < 0 || offset >= nvS_) { - Fail(kMoonEpCombineStatusInvalidRoute); + failureStatus_ = status; + StoreInt(localWindow_ + coreStatusOffset_ + + static_cast(core_) * kMoonEpCombineV2TokenStrideBytes, + status); + } + + __aicore__ inline void Fail(int32_t status) + { + if (failureStatus_ == 0) { + StoreCoreStatus(status); + } + } + + __aicore__ inline int32_t FirstFailure() + { + for (uint32_t core = 0U; core < activeCoreCount_; ++core) { + const int32_t status = LoadInt(localWindow_ + coreStatusOffset_ + + static_cast(core) * kMoonEpCombineV2TokenStrideBytes); + if (status != 0) { + return status; + } + } + return 0; + } + + __aicore__ inline bool DecodeReverseRoute( + int32_t encoded, int32_t *peer, int64_t *target) const + { + if (encoded == -1) { return false; } - return true; + if (encoded < 0) { + return false; + } + *peer = static_cast(static_cast(encoded) / nvS_); + *target = static_cast(encoded) % nvS_; + return *peer >= 0 && *peer < rankSize_ && *target >= 0 && *target < n_; + } + + __aicore__ inline void BuildDuplicateMask() + { + const int64_t rowBegin = nvS_ * core_ / activeCoreCount_; + const int64_t rowEnd = nvS_ * (core_ + 1U) / activeCoreCount_; + ZeroBytes(localWindow_ + duplicateMaskOffset_ + + rowBegin * static_cast(sizeof(int32_t)), + (rowEnd - rowBegin) * static_cast(sizeof(int32_t))); + AscendC::SyncAll(); + + const int32_t groupCount = LoadInt(dupCountsAddr_); + const int32_t duplicateCount = LoadInt( + dupCountsAddr_ + sizeof(int32_t)); + if (groupCount < 0 || groupCount > nvS_ || duplicateCount < 0 || + duplicateCount > nvS_) { + Fail(kMoonEpCombineStatusInvalidRoute); + return; + } + const int32_t groupBegin = static_cast( + static_cast(groupCount) * core_ / activeCoreCount_); + const int32_t groupEnd = static_cast( + static_cast(groupCount) * (core_ + 1U) / activeCoreCount_); + for (int32_t group = groupBegin; group < groupEnd; ++group) { + const int64_t base = static_cast(group) * 3; + const int32_t primary = LoadInt( + dupGroupsAddr_ + base * sizeof(int32_t)); + const int32_t start = LoadInt( + dupGroupsAddr_ + (base + 1) * sizeof(int32_t)); + const int32_t count = LoadInt( + dupGroupsAddr_ + (base + 2) * sizeof(int32_t)); + if (primary < 0 || primary >= nvS_ || start < 0 || count <= 0 || + static_cast(start) + count > duplicateCount) { + Fail(kMoonEpCombineStatusInvalidRoute); + return; + } + for (int32_t index = 0; index < count; ++index) { + const int32_t duplicate = LoadInt(dupLoffsAddr_ + + static_cast(start + index) * sizeof(int32_t)); + if (duplicate < 0 || duplicate >= nvS_ || duplicate == primary) { + Fail(kMoonEpCombineStatusInvalidRoute); + return; + } + StoreInt(localWindow_ + duplicateMaskOffset_ + + static_cast(duplicate) * sizeof(int32_t), 1); + } + } + AscendC::PipeBarrier(); + } + + __aicore__ inline void ValidateReverseRoutes() + { + const int64_t begin = nvS_ * core_ / activeCoreCount_; + const int64_t end = nvS_ * (core_ + 1U) / activeCoreCount_; + for (int64_t row = begin; row < end; ++row) { + const int32_t encoded = LoadInt( + dstLocalAddr_ + row * static_cast(sizeof(int32_t))); + if (encoded == -1) { + continue; + } + int32_t peer = 0; + int64_t target = 0; + if (!DecodeReverseRoute(encoded, &peer, &target)) { + Fail(kMoonEpCombineStatusInvalidRoute); + return; + } + } } - __aicore__ inline void PublishLocalInput( + __aicore__ inline int32_t WaitPeerStep(int32_t peer, uint32_t expectedStep) + { + AscendC::GlobalTensor flag; + flag.SetGlobalBuffer(reinterpret_cast<__gm__ int64_t *>(shareAddrs_[peer]), + FLAG_UNIT_INT_NUM); + const int64_t expectedMagic = static_cast( + static_cast(magic_)) << MAGIC_OFFSET; + for (uint64_t iteration = 0; iteration < waitIterations_; ++iteration) { + AscendC::DataCacheCleanAndInvalid(flag); + const int64_t value = flag.GetValue(0); + if ((value & MAGIC_MASK) != (expectedMagic & MAGIC_MASK)) { + continue; + } + const int32_t step = static_cast(value & ~MAGIC_MASK); + if (step == static_cast(expectedStep)) { + return 0; + } + if (step == kMoonEpCombineFailedStep) { + return 1; + } + } + return 2; + } + + __aicore__ inline bool CrossRankBarrier(uint32_t step) + { + AscendC::SyncAll(); + if (core_ == 0U) { + const int32_t localFailure = FirstFailure(); + sync_.SetInnerFlag(static_cast(magic_), + localFailure == 0 ? static_cast(step) : + kMoonEpCombineFailedStep); + if (localFailure == 0) { + for (int32_t peer = 0; peer < rankSize_; ++peer) { + const int32_t result = WaitPeerStep(peer, step); + if (result == 1) { + Fail(kMoonEpCombineStatusRemoteFailureBase + peer); + break; + } + if (result == 2) { + Fail(kMoonEpCombineStatusTimeoutBase + peer); + break; + } + } + } + } + AscendC::SyncAll(); + return FirstFailure() == 0; + } + + __aicore__ inline void PrepareChunk( int64_t chunkOffset, int64_t bytesThisChunk, bool firstChunk) { - for (int64_t row = 0; row < nvS_; ++row) { - CopyBytesGmToGm(localWindow_ + row * hiddenChunkStride_, + const int64_t sourceBegin = nvS_ * core_ / activeCoreCount_; + const int64_t sourceEnd = nvS_ * (core_ + 1U) / activeCoreCount_; + for (int64_t row = sourceBegin; row < sourceEnd; ++row) { + CopyBytesGmToGm(localWindow_ + sourceHiddenOffset_ + + row * hiddenChunkStride_, hiddenNvshAddr_ + row * hiddenRowBytes_ + chunkOffset, routeBuf_, bytesThisChunk); } - if (firstChunk && routeWeightsBytes_ > 0) { - CopyBytesGmToGm(localWindow_ + static_cast(routeWeightsOffset_), - routeWeightsNvsAddr_, routeBuf_, static_cast(routeWeightsBytes_)); + const int64_t receiveBegin = n_ * core_ / activeCoreCount_; + const int64_t receiveEnd = n_ * (core_ + 1U) / activeCoreCount_; + ZeroBytes(localWindow_ + receiveHiddenOffset_ + + receiveBegin * hiddenChunkStride_, + (receiveEnd - receiveBegin) * hiddenChunkStride_); + if (firstChunk && routeWeightsBytes_ != 0U) { + CopyBytesGmToGm(localWindow_ + sourceWeightsOffset_ + + sourceBegin * static_cast(sizeof(float)), + routeWeightsNvsAddr_ + + sourceBegin * static_cast(sizeof(float)), + routeBuf_, (sourceEnd - sourceBegin) * sizeof(float)); + ZeroBytes(localWindow_ + receiveWeightsOffset_ + + receiveBegin * static_cast(sizeof(float)), + (receiveEnd - receiveBegin) * sizeof(float)); } - AscendC::PipeBarrier(); } - __aicore__ inline bool PreReduceDuplicates(int64_t bytesThisChunk) + __aicore__ inline void PreReduceDuplicates(int64_t bytesThisChunk) { const int32_t groupCount = LoadInt(dupCountsAddr_); - const int32_t duplicateCount = LoadInt(dupCountsAddr_ + sizeof(int32_t)); - if (groupCount < 0 || groupCount > nvS_ || duplicateCount < 0 || - duplicateCount > nvS_ || bytesThisChunk % sizeof(bfloat16_t) != 0) { - Fail(kMoonEpCombineStatusInvalidRoute); - return false; + const int32_t duplicateCount = LoadInt( + dupCountsAddr_ + sizeof(int32_t)); + if (failureStatus_ != 0 || groupCount <= 0) { + return; } - - AscendC::LocalTensor bfloatScratch = bfloatBuf_.Get(); + AscendC::LocalTensor bfloatScratch = + bfloatBuf_.Get(); AscendC::LocalTensor routeScratch = routeBuf_.Get(); AscendC::LocalTensor accumulator = accumulatorBuf_.Get(); const int64_t elements = bytesThisChunk / sizeof(bfloat16_t); - for (int32_t group = 0; group < groupCount; ++group) { + const int32_t groupBegin = static_cast( + static_cast(groupCount) * core_ / activeCoreCount_); + const int32_t groupEnd = static_cast( + static_cast(groupCount) * (core_ + 1U) / activeCoreCount_); + for (int32_t group = groupBegin; group < groupEnd; ++group) { const int64_t base = static_cast(group) * 3; - const int32_t primary = LoadInt(dupGroupsAddr_ + base * sizeof(int32_t)); - const int32_t start = LoadInt(dupGroupsAddr_ + (base + 1) * sizeof(int32_t)); - const int32_t count = LoadInt(dupGroupsAddr_ + (base + 2) * sizeof(int32_t)); - if (primary < 0 || primary >= nvS_ || start < 0 || count <= 0 || + const int32_t primary = LoadInt( + dupGroupsAddr_ + base * sizeof(int32_t)); + const int32_t start = LoadInt( + dupGroupsAddr_ + (base + 1) * sizeof(int32_t)); + const int32_t count = LoadInt( + dupGroupsAddr_ + (base + 2) * sizeof(int32_t)); + if (start < 0 || count <= 0 || static_cast(start) + count > duplicateCount) { Fail(kMoonEpCombineStatusInvalidRoute); - return false; + return; } - for (int64_t tileOffset = 0; tileOffset < elements; - tileOffset += kMoonEpCombineHiddenTileElements) { + tileOffset += kMoonEpCombineHiddenTileElements) { const int64_t tileElements = MinInt64( elements - tileOffset, kMoonEpCombineHiddenTileElements); AscendC::GlobalTensor primaryGm; primaryGm.SetGlobalBuffer(reinterpret_cast<__gm__ bfloat16_t *>( - localWindow_ + static_cast(primary) * hiddenChunkStride_) + - tileOffset, tileElements); - AscendC::DataCopyExtParams params { - 1, static_cast(tileElements * sizeof(bfloat16_t)), 0, 0, 0 - }; + localWindow_ + sourceHiddenOffset_ + + static_cast(primary) * hiddenChunkStride_) + + tileOffset, tileElements); + AscendC::DataCopyExtParams params {1, + static_cast(tileElements * sizeof(bfloat16_t)), + 0, 0, 0}; AscendC::DataCopyPadExtParams pad {false, 0, 0, 0}; AscendC::DataCopyPad(bfloatScratch, primaryGm, params, pad); AscendC::SetFlag(EVENT_ID0); AscendC::WaitFlag(EVENT_ID0); AscendC::Cast(accumulator, bfloatScratch, - AscendC::RoundMode::CAST_NONE, static_cast(tileElements)); + AscendC::RoundMode::CAST_NONE, + static_cast(tileElements)); AscendC::PipeBarrier(); - for (int32_t index = 0; index < count; ++index) { const int32_t duplicate = LoadInt(dupLoffsAddr_ + static_cast(start + index) * sizeof(int32_t)); - if (duplicate < 0 || duplicate >= nvS_) { - Fail(kMoonEpCombineStatusInvalidRoute); - return false; - } AscendC::GlobalTensor duplicateGm; - duplicateGm.SetGlobalBuffer(reinterpret_cast<__gm__ bfloat16_t *>( - localWindow_ + static_cast(duplicate) * hiddenChunkStride_) + - tileOffset, tileElements); + duplicateGm.SetGlobalBuffer( + reinterpret_cast<__gm__ bfloat16_t *>(localWindow_ + + sourceHiddenOffset_ + + static_cast(duplicate) * hiddenChunkStride_) + + tileOffset, tileElements); AscendC::DataCopyPad(bfloatScratch, duplicateGm, params, pad); AscendC::SetFlag(EVENT_ID0); AscendC::WaitFlag(EVENT_ID0); AscendC::Cast(routeScratch, bfloatScratch, - AscendC::RoundMode::CAST_NONE, static_cast(tileElements)); + AscendC::RoundMode::CAST_NONE, + static_cast(tileElements)); AscendC::PipeBarrier(); AscendC::Add(accumulator, accumulator, routeScratch, static_cast(tileElements)); AscendC::PipeBarrier(); } - AscendC::Cast(bfloatScratch, accumulator, - AscendC::RoundMode::CAST_RINT, static_cast(tileElements)); + AscendC::RoundMode::CAST_RINT, + static_cast(tileElements)); AscendC::SetFlag(EVENT_ID0); AscendC::WaitFlag(EVENT_ID0); AscendC::DataCopyPad(primaryGm, bfloatScratch, params); @@ -418,183 +543,199 @@ class CombineKernel { } } AscendC::PipeBarrier(); + } + + __aicore__ inline uint64_t DoneToken(uint32_t chunk, uint32_t step) const + { + return (static_cast(static_cast(magic_)) << 32U) | + (static_cast(chunk) << 8U) | step; + } + + __aicore__ inline void StoreToken(GM_ADDR address, uint64_t value) + { + AscendC::LocalTensor local = bfloatBuf_.Get(); + AscendC::GlobalTensor dst; + dst.SetGlobalBuffer(reinterpret_cast<__gm__ uint64_t *>(address), 1); + local.SetValue(0, value); + AscendC::SetFlag(EVENT_ID0); + AscendC::WaitFlag(EVENT_ID0); + AscendC::DataCopyExtParams params {1, sizeof(uint64_t), 0, 0, 0}; + AscendC::DataCopyPad(dst, local, params); + AscendC::SetFlag(EVENT_ID0); + AscendC::WaitFlag(EVENT_ID0); + } + + __aicore__ inline void PushPeerRows(uint32_t peer, uint32_t step, + uint32_t chunk, int64_t bytesThisChunk, bool firstChunk) + { + GM_ADDR remoteWindow = shareAddrs_[peer] + TileXR::IPC_DATA_OFFSET; + for (int64_t source = 0; source < nvS_; ++source) { + const int32_t encoded = LoadInt(dstLocalAddr_ + + source * static_cast(sizeof(int32_t))); + if (encoded == -1) { + continue; + } + int32_t targetPeer = 0; + int64_t target = 0; + if (!DecodeReverseRoute(encoded, &targetPeer, &target) || + targetPeer != static_cast(peer)) { + continue; + } + const int32_t duplicate = LoadInt(localWindow_ + duplicateMaskOffset_ + + source * static_cast(sizeof(int32_t))); + if (duplicate == 0) { + CopyBytesGmToGm(remoteWindow + receiveHiddenOffset_ + + target * hiddenChunkStride_, + localWindow_ + sourceHiddenOffset_ + + source * hiddenChunkStride_, + routeBuf_, bytesThisChunk); + } + if (firstChunk && routeWeightsBytes_ != 0U) { + CopyBytesGmToGm(remoteWindow + receiveWeightsOffset_ + + target * static_cast(sizeof(float)), + localWindow_ + sourceWeightsOffset_ + + source * static_cast(sizeof(float)), + routeBuf_, sizeof(float)); + } + } + AscendC::PipeBarrier(); + StoreToken(remoteWindow + doneOffset_ + + static_cast(rank_) * kMoonEpCombineV2TokenStrideBytes, + DoneToken(chunk, step)); + } + + __aicore__ inline uint64_t LoadToken(GM_ADDR address) + { + AscendC::GlobalTensor token; + token.SetGlobalBuffer(reinterpret_cast<__gm__ uint64_t *>(address), 1); + AscendC::DataCacheCleanAndInvalid(token); + return token.GetValue(0); + } + + __aicore__ inline bool WaitInboundDone(uint32_t chunk) + { + for (uint32_t sourceIndex = 0U; + sourceIndex < sourcesPerCore_; ++sourceIndex) { + const uint32_t source = MoonEpCombineV2SourceForCore( + core_, sourceIndex, static_cast(rankSize_)); + const uint32_t step = MoonEpCombineV2ReceiveStep( + static_cast(rank_), source, + static_cast(rankSize_)); + const uint64_t expected = DoneToken(chunk, step); + GM_ADDR address = localWindow_ + doneOffset_ + + static_cast(source) * + kMoonEpCombineV2TokenStrideBytes; + bool ready = false; + for (uint64_t iteration = 0; iteration < waitIterations_; ++iteration) { + if (LoadToken(address) == expected) { + ready = true; + break; + } + } + if (!ready) { + Fail(kMoonEpCombineStatusTimeoutBase + + static_cast(source)); + return false; + } + } return true; } - __aicore__ inline bool ReduceHiddenChunk( + __aicore__ inline void ReduceHiddenChunk( int64_t chunkOffset, int64_t bytesThisChunk) { - if (bytesThisChunk % sizeof(bfloat16_t) != 0) { - Fail(kMoonEpCombineStatusInvalidRoute); - return false; + if (failureStatus_ != 0) { + return; } const int64_t elements = bytesThisChunk / sizeof(bfloat16_t); - AscendC::LocalTensor bfloatScratch = bfloatBuf_.Get(); + AscendC::LocalTensor bfloatScratch = + bfloatBuf_.Get(); AscendC::LocalTensor routeScratch = routeBuf_.Get(); AscendC::LocalTensor accumulator = accumulatorBuf_.Get(); - - for (int64_t token = 0; token < s_; ++token) { + const int64_t tokenBegin = s_ * core_ / activeCoreCount_; + const int64_t tokenEnd = s_ * (core_ + 1U) / activeCoreCount_; + for (int64_t token = tokenBegin; token < tokenEnd; ++token) { for (int64_t tileOffset = 0; tileOffset < elements; - tileOffset += kMoonEpCombineHiddenTileElements) { + tileOffset += kMoonEpCombineHiddenTileElements) { const int64_t tileElements = MinInt64( elements - tileOffset, kMoonEpCombineHiddenTileElements); - AscendC::Duplicate(accumulator, 0.0f, static_cast(tileElements)); + AscendC::Duplicate(accumulator, 0.0f, + static_cast(tileElements)); AscendC::PipeBarrier(); - for (int64_t topk = 0; topk < k_; ++topk) { const int64_t route = token * k_ + topk; - const int32_t encoded = LoadRoute(route); - if (encoded < 0) { - continue; - } - int64_t peer = 0; - int64_t offset = 0; - if (!DecodeRoute(encoded, peer, offset)) { - return false; - } - const uint64_t remoteByteOffset = - static_cast(offset) * hiddenChunkStride_ + - static_cast(tileOffset) * sizeof(bfloat16_t); - GM_ADDR source = nullptr; - if (IsConsumeOnly() && !FetchRegisteredPayload( - peer, remoteByteOffset, - static_cast(tileElements * sizeof(bfloat16_t)), - source)) { - return false; - } - if (!IsConsumeOnly()) { - source = shareAddrs_[peer] + TileXR::IPC_DATA_OFFSET + - offset * hiddenChunkStride_ + - tileOffset * static_cast(sizeof(bfloat16_t)); - } - AscendC::GlobalTensor remote; - remote.SetGlobalBuffer(reinterpret_cast<__gm__ bfloat16_t *>( - source), tileElements); - AscendC::DataCopyExtParams params { - 1, static_cast(tileElements * sizeof(bfloat16_t)), 0, 0, 0 - }; - AscendC::DataCopyPadExtParams pad {false, 0, 0, 0}; - AscendC::DataCopyPad(bfloatScratch, remote, params, pad); + AscendC::GlobalTensor input; + input.SetGlobalBuffer(reinterpret_cast<__gm__ bfloat16_t *>( + localWindow_ + receiveHiddenOffset_ + + route * hiddenChunkStride_) + tileOffset, + tileElements); + AscendC::DataCopyExtParams params {1, + static_cast( + tileElements * sizeof(bfloat16_t)), 0, 0, 0}; + AscendC::DataCopyPadExtParams pad { + false, 0, 0, 0}; + AscendC::DataCopyPad(bfloatScratch, input, params, pad); AscendC::SetFlag(EVENT_ID0); AscendC::WaitFlag(EVENT_ID0); AscendC::Cast(routeScratch, bfloatScratch, - AscendC::RoundMode::CAST_NONE, static_cast(tileElements)); + AscendC::RoundMode::CAST_NONE, + static_cast(tileElements)); AscendC::PipeBarrier(); AscendC::Add(accumulator, accumulator, routeScratch, static_cast(tileElements)); AscendC::PipeBarrier(); } - AscendC::Cast(bfloatScratch, accumulator, - AscendC::RoundMode::CAST_RINT, static_cast(tileElements)); + AscendC::RoundMode::CAST_RINT, + static_cast(tileElements)); AscendC::SetFlag(EVENT_ID0); AscendC::WaitFlag(EVENT_ID0); AscendC::GlobalTensor output; output.SetGlobalBuffer(reinterpret_cast<__gm__ bfloat16_t *>( - hiddenShAddr_ + token * hiddenRowBytes_ + chunkOffset) + tileOffset, - tileElements); - AscendC::DataCopyExtParams params { - 1, static_cast(tileElements * sizeof(bfloat16_t)), 0, 0, 0 - }; + hiddenShAddr_ + token * hiddenRowBytes_ + chunkOffset) + + tileOffset, tileElements); + AscendC::DataCopyExtParams params {1, + static_cast( + tileElements * sizeof(bfloat16_t)), 0, 0, 0}; AscendC::DataCopyPad(output, bfloatScratch, params); AscendC::SetFlag(EVENT_ID0); AscendC::WaitFlag(EVENT_ID0); } } AscendC::PipeBarrier(); - return true; - } - - __aicore__ inline bool GatherWeights() - { - for (int64_t route = 0; route < n_; ++route) { - const int32_t encoded = LoadRoute(route); - int64_t peer = 0; - int64_t offset = 0; - if (!DecodeRoute(encoded, peer, offset)) { - return false; - } - GM_ADDR source = nullptr; - if (IsConsumeOnly() && !FetchRegisteredPayload(peer, - routeWeightsOffset_ + static_cast(offset) * sizeof(float), - sizeof(float), source)) { - return false; - } - if (!IsConsumeOnly()) { - source = shareAddrs_[peer] + TileXR::IPC_DATA_OFFSET + - static_cast(routeWeightsOffset_) + offset * sizeof(float); - } - CopyBytesGmToGm(routeWeightsSkAddr_ + route * sizeof(float), source, - routeBuf_, sizeof(float)); - } - return true; - } - - __aicore__ inline void StoreStatus(int32_t value) - { - StoreInt(statusAddr_, value); } - __aicore__ inline void PublishStep(int32_t step) + __aicore__ inline void CopyReceivedWeights() { - AscendC::PipeBarrier(); - sync_.SetInnerFlag(static_cast(magic_), step); + const int64_t begin = n_ * core_ / activeCoreCount_; + const int64_t end = n_ * (core_ + 1U) / activeCoreCount_; + CopyBytesGmToGm(routeWeightsSkAddr_ + + begin * static_cast(sizeof(float)), + localWindow_ + receiveWeightsOffset_ + + begin * static_cast(sizeof(float)), + routeBuf_, (end - begin) * sizeof(float)); } - __aicore__ inline int32_t WaitPeerStep(int32_t peer, int32_t expectedStep) - { - AscendC::GlobalTensor flag; - flag.SetGlobalBuffer(reinterpret_cast<__gm__ int64_t *>(shareAddrs_[peer]), - FLAG_UNIT_INT_NUM); - const int64_t expectedMagic = - static_cast(static_cast(magic_)) << MAGIC_OFFSET; - for (uint64_t iteration = 0; iteration < waitIterations_; ++iteration) { - AscendC::DataCacheCleanAndInvalid(flag); - const int64_t value = flag.GetValue(0); - if ((value & MAGIC_MASK) != (expectedMagic & MAGIC_MASK)) { - continue; - } - const int32_t step = static_cast(value & ~MAGIC_MASK); - if (step == expectedStep) { - return 0; - } - if (step == kMoonEpCombineFailedStep) { - return 1; - } - } - return 2; - } - - __aicore__ inline bool WaitAllPeers(int32_t expectedStep) - { - for (int32_t offset = 0; offset < rankSize_; ++offset) { - const int32_t peer = (rank_ + offset) % rankSize_; - const int32_t result = WaitPeerStep(peer, expectedStep); - if (result == 1) { - Fail(kMoonEpCombineStatusRemoteFailureBase + peer); - return false; - } - if (result == 2) { - Fail(kMoonEpCombineStatusTimeoutBase + peer); - return false; - } - } - return true; - } - - __aicore__ inline void Fail(int32_t status) + __aicore__ inline void Finish() { - StoreStatus(status); - if (!IsPublishOnly() && !IsConsumeOnly()) { - PublishStep(kMoonEpCombineFailedStep); + AscendC::SyncAll(); + if (core_ == 0U) { + StoreInt(statusAddr_, FirstFailure()); } + AscendC::SyncAll(); } __gm__ TileXR::CommArgs *args_ = nullptr; + bool valid_ = false; int32_t rank_ = 0; int32_t rankSize_ = 0; + uint32_t core_ = 0U; + uint32_t activeCoreCount_ = 0U; + uint32_t stepCount_ = 0U; + uint32_t sourcesPerCore_ = 0U; + int32_t failureStatus_ = 0; int64_t s_ = 0; int64_t k_ = 0; int64_t n_ = 0; @@ -603,13 +744,20 @@ class CombineKernel { int64_t hiddenChunkBytes_ = 0; int64_t hiddenChunkStride_ = 0; int64_t chunkCount_ = 0; - uint64_t hiddenPayloadBytes_ = 0; - uint64_t routeWeightsOffset_ = 0; - uint64_t routeWeightsBytes_ = 0; - uint64_t waitIterations_ = 0; - uint64_t flags_ = 0; + uint64_t sourceHiddenOffset_ = 0U; + uint64_t receiveHiddenOffset_ = 0U; + uint64_t hiddenPayloadBytes_ = 0U; + uint64_t sourceWeightsOffset_ = 0U; + uint64_t receiveWeightsOffset_ = 0U; + uint64_t routeWeightsBytes_ = 0U; + uint64_t duplicateMaskOffset_ = 0U; + uint64_t doneOffset_ = 0U; + uint64_t coreStatusOffset_ = 0U; + uint64_t windowBytes_ = 0U; + uint64_t waitIterations_ = 0U; + uint64_t flags_ = 0U; int64_t magic_ = 0; - bool initialized_ = false; + GM_ADDR dstLocalAddr_ = nullptr; GM_ADDR dstAddr_ = nullptr; GM_ADDR dupGroupsAddr_ = nullptr; GM_ADDR dupLoffsAddr_ = nullptr; @@ -620,7 +768,6 @@ class CombineKernel { GM_ADDR routeWeightsSkAddr_ = nullptr; GM_ADDR statusAddr_ = nullptr; GM_ADDR localWindow_ = nullptr; - GM_ADDR registeredScratch_ = nullptr; GM_ADDR shareAddrs_[TileXR::TILEXR_MAX_RANK_SIZE] = {}; AscendC::TPipe pipe_; AscendC::TBuf syncBuf_; @@ -633,20 +780,28 @@ class CombineKernel { } // namespace Kernel } // namespace TileXRMoonEp -extern "C" __global__ __aicore__ void tilexr_moonep_combine_kernel(GM_ADDR commArgs, - GM_ADDR dst, GM_ADDR dupGroups, GM_ADDR dupLoffs, GM_ADDR dupCounts, - GM_ADDR hiddenNvsh, GM_ADDR routeWeightsNvs, GM_ADDR hiddenSh, - GM_ADDR routeWeightsSk, GM_ADDR status, int64_t s, int64_t k, int64_t n, - int64_t nvS, int64_t hiddenRowBytes, int64_t hiddenChunkBytes, - int64_t hiddenChunkStride, int64_t chunkCount, uint64_t hiddenPayloadBytes, - uint64_t routeWeightsOffset, uint64_t routeWeightsBytes, +extern "C" __global__ __aicore__ void tilexr_moonep_combine_kernel( + GM_ADDR commArgs, GM_ADDR dstLocal, GM_ADDR dst, GM_ADDR dupGroups, + GM_ADDR dupLoffs, GM_ADDR dupCounts, GM_ADDR hiddenNvsh, + GM_ADDR routeWeightsNvs, GM_ADDR hiddenSh, GM_ADDR routeWeightsSk, + GM_ADDR status, int64_t s, int64_t k, int64_t n, int64_t nvS, + int64_t hiddenRowBytes, int64_t hiddenChunkBytes, + int64_t hiddenChunkStride, int64_t chunkCount, + uint64_t sourceHiddenOffset, uint64_t receiveHiddenOffset, + uint64_t hiddenPayloadBytes, uint64_t sourceWeightsOffset, + uint64_t receiveWeightsOffset, uint64_t routeWeightsBytes, + uint64_t duplicateMaskOffset, uint64_t doneOffset, + uint64_t coreStatusOffset, uint64_t windowBytes, uint64_t waitIterations, uint64_t flags, int64_t magic) { TileXRMoonEp::Kernel::CombineKernel op; - op.Init(commArgs, dst, dupGroups, dupLoffs, dupCounts, hiddenNvsh, - routeWeightsNvs, hiddenSh, routeWeightsSk, status, s, k, n, nvS, - hiddenRowBytes, hiddenChunkBytes, hiddenChunkStride, chunkCount, - hiddenPayloadBytes, routeWeightsOffset, routeWeightsBytes, - waitIterations, flags, magic); + op.Init(commArgs, dstLocal, dst, dupGroups, dupLoffs, dupCounts, + hiddenNvsh, routeWeightsNvs, hiddenSh, routeWeightsSk, status, + s, k, n, nvS, hiddenRowBytes, hiddenChunkBytes, + hiddenChunkStride, chunkCount, sourceHiddenOffset, + receiveHiddenOffset, hiddenPayloadBytes, sourceWeightsOffset, + receiveWeightsOffset, routeWeightsBytes, duplicateMaskOffset, + doneOffset, coreStatusOffset, windowBytes, waitIterations, flags, + magic); op.Process(); } diff --git a/src/moonep/combine_v2/CMakeLists.txt b/src/moonep/combine_v2/CMakeLists.txt index bbecd9c..caa04d1 100644 --- a/src/moonep/combine_v2/CMakeLists.txt +++ b/src/moonep/combine_v2/CMakeLists.txt @@ -75,6 +75,7 @@ tilexr_add_moonep_kernel(tilexr_moonep_combine_v2_kernel "${CMAKE_CURRENT_SOURCE_DIR}/kernels/tilexr_moonep_combine_v2_kernel.h" "${CMAKE_CURRENT_SOURCE_DIR}/common/combine_v2_profile.h" "${CMAKE_CURRENT_SOURCE_DIR}/common/combine_v2_schedule.h" + "${CMAKE_SOURCE_DIR}/src/moonep/common/moonep_combine_schedule.h" "${CMAKE_CURRENT_SOURCE_DIR}/common/combine_v2_wqe_batch.h" "${CMAKE_SOURCE_DIR}/src/include/comm_args.h" "${CMAKE_SOURCE_DIR}/src/include/tilexr_udma.h" diff --git a/src/moonep/combine_v2/common/combine_v2_schedule.h b/src/moonep/combine_v2/common/combine_v2_schedule.h index f6789e3..9985c36 100644 --- a/src/moonep/combine_v2/common/combine_v2_schedule.h +++ b/src/moonep/combine_v2/common/combine_v2_schedule.h @@ -1,366 +1,6 @@ -#ifndef TILEXR_MOONEP_COMBINE_V2_SCHEDULE_H -#define TILEXR_MOONEP_COMBINE_V2_SCHEDULE_H +#ifndef TILEXR_MOONEP_COMBINE_V2_SCHEDULE_WRAPPER_H +#define TILEXR_MOONEP_COMBINE_V2_SCHEDULE_WRAPPER_H -#include -#include +#include "moonep_combine_schedule.h" -#if defined(__CCE__) && defined(__CCE_IS_AICORE__) -#define TILEXR_MOONEP_COMBINE_V2_INLINE \ - __attribute__((always_inline)) inline __aicore__ -#else -#define TILEXR_MOONEP_COMBINE_V2_INLINE inline constexpr -#endif - -namespace TileXRMoonEp { - -constexpr uint32_t kMoonEpCombineV2RankCount = 128U; -constexpr uint32_t kMoonEpCombineV2GroupSize = 8U; -constexpr uint32_t kMoonEpCombineV2GroupCount = 16U; -constexpr uint32_t kMoonEpCombineV2GroupsPerHalf = 8U; -constexpr uint32_t kMoonEpCombineV2StepCount = 8U; -constexpr uint32_t kMoonEpCombineV2GrantStepCount = - kMoonEpCombineV2StepCount - 1U; -constexpr uint32_t kMoonEpCombineV2CoreCount = 16U; -constexpr uint32_t kMoonEpCombineV2LaneCount = 2U; -constexpr uint32_t kMoonEpCombineV2QpCount = 32U; -constexpr uint32_t kMoonEpCombineV2MaxSourcesPerCore = - kMoonEpCombineV2RankCount / kMoonEpCombineV2CoreCount; -constexpr uint32_t kMoonEpCombineV2EpochCount = 2U; -constexpr uint32_t kMoonEpCombineV2LogicalBatchRows = 128U; -constexpr uint32_t kMoonEpCombineV2SixPortRows = 96U; -constexpr uint32_t kMoonEpCombineV2TwoPortRows = 32U; -constexpr uint32_t kMoonEpCombineV2SelectionChunkRows = 16384U; -constexpr uint32_t kMoonEpCombineV2MaxOutstanding = 16384U; -constexpr int64_t kMoonEpCombineV2SmallBs = 8; -constexpr int64_t kMoonEpCombineV2SmallSlots = 128; -constexpr int64_t kMoonEpCombineV2TargetBs = 8192; -constexpr int64_t kMoonEpCombineV2TargetH = 3584; -constexpr int64_t kMoonEpCombineV2TargetTopK = 16; -constexpr int64_t kMoonEpCombineV2TargetSlots = 131072; -constexpr uint64_t kMoonEpCombineV2TokenStrideBytes = 64U; -constexpr uint64_t kMoonEpCombineV2GrantSlotBytes = 512U; -constexpr uint64_t kMoonEpCombineV2GrantReceiveOffsetBytes = 0U; -constexpr uint64_t kMoonEpCombineV2GrantSourceOffsetBytes = 64U; -constexpr uint32_t kMoonEpCombineV2FailureMarker = 0x47505632U; -constexpr uint64_t kMoonEpCombineV2MaxMagic = - std::numeric_limits::max() >> 3U; - -enum MoonEpCombineV2Lane : uint32_t { - MOONEP_COMBINE_V2_SIX_PORT = 0U, - MOONEP_COMBINE_V2_TWO_PORT = 1U, -}; - -enum MoonEpCombineV2FailureStatus : uint32_t { - MOONEP_COMBINE_V2_SUCCESS = 0U, - MOONEP_COMBINE_V2_INVALID_CONFIG = 1U, - MOONEP_COMBINE_V2_POISONED = 2U, - MOONEP_COMBINE_V2_OUTSTANDING_LIMIT = 3U, - MOONEP_COMBINE_V2_CQ_TIMEOUT = 4U, - MOONEP_COMBINE_V2_CQ_ERROR = 5U, - MOONEP_COMBINE_V2_GRANT_TIMEOUT = 6U, - MOONEP_COMBINE_V2_DONE_TIMEOUT = 7U, - MOONEP_COMBINE_V2_BAD_DESTINATION = 8U, -}; - -struct alignas(64) MoonEpCombineV2FailureRecord { - uint64_t magic; - uint32_t status; - uint32_t rank; - uint32_t core; - uint32_t step; - uint32_t peer; - uint32_t lane; - uint32_t qp; - uint32_t cqStatus; - uint64_t expected; - uint64_t observed; - uint32_t poison; - uint32_t marker; -}; - -struct MoonEpCombineV2LaneCounts { - uint32_t sixPort; - uint32_t twoPort; -}; - -struct MoonEpCombineV2RingSegments { - uint32_t first; - uint32_t second; -}; - -static_assert(sizeof(MoonEpCombineV2FailureRecord) == 64U, - "MoonEP Combine V2 failure record ABI changed"); -static_assert(kMoonEpCombineV2GrantSourceOffsetBytes + sizeof(uint64_t) <= - kMoonEpCombineV2GrantSlotBytes, - "MoonEP Combine V2 Grant source exceeds its slot"); - -TILEXR_MOONEP_COMBINE_V2_INLINE bool -MoonEpCombineV2RankSizeSupported(uint32_t rankSize) -{ - return (rankSize >= 2U && rankSize <= kMoonEpCombineV2GroupSize) || - rankSize == 16U || rankSize == 32U || rankSize == 64U || - rankSize == kMoonEpCombineV2RankCount; -} - -TILEXR_MOONEP_COMBINE_V2_INLINE bool -MoonEpCombineV2RankValid(uint32_t rank, uint32_t rankSize) -{ - return MoonEpCombineV2RankSizeSupported(rankSize) && rank < rankSize; -} - -TILEXR_MOONEP_COMBINE_V2_INLINE uint32_t -MoonEpCombineV2ActiveCoreCount(uint32_t rankSize) -{ - return rankSize <= kMoonEpCombineV2GroupSize ? rankSize : - kMoonEpCombineV2CoreCount; -} - -TILEXR_MOONEP_COMBINE_V2_INLINE uint32_t -MoonEpCombineV2StepCount(uint32_t rankSize) -{ - return rankSize <= kMoonEpCombineV2GroupSize ? 1U : - rankSize / kMoonEpCombineV2CoreCount; -} - -TILEXR_MOONEP_COMBINE_V2_INLINE uint32_t -MoonEpCombineV2GroupsPerHalf(uint32_t rankSize) -{ - return rankSize / kMoonEpCombineV2CoreCount; -} - -TILEXR_MOONEP_COMBINE_V2_INLINE uint32_t -MoonEpCombineV2LocalRankSize(uint32_t rankSize) -{ - return rankSize <= kMoonEpCombineV2GroupSize ? rankSize : - kMoonEpCombineV2GroupSize; -} - -TILEXR_MOONEP_COMBINE_V2_INLINE bool -MoonEpCombineV2CoreValid(uint32_t core, uint32_t rankSize) -{ - return MoonEpCombineV2RankSizeSupported(rankSize) && - core < MoonEpCombineV2ActiveCoreCount(rankSize); -} - -TILEXR_MOONEP_COMBINE_V2_INLINE bool -MoonEpCombineV2StepValid(uint32_t step, uint32_t rankSize) -{ - return MoonEpCombineV2RankSizeSupported(rankSize) && - step < MoonEpCombineV2StepCount(rankSize); -} - -TILEXR_MOONEP_COMBINE_V2_INLINE bool -MoonEpCombineV2MagicValid(uint64_t magic) -{ - return magic != 0U && magic <= kMoonEpCombineV2MaxMagic; -} - -TILEXR_MOONEP_COMBINE_V2_INLINE uint32_t -MoonEpCombineV2Epoch(uint64_t magic) -{ - return static_cast(magic & 1U); -} - -TILEXR_MOONEP_COMBINE_V2_INLINE uint32_t -MoonEpCombineV2Qp(uint32_t core, uint32_t lane) -{ - return lane == MOONEP_COMBINE_V2_SIX_PORT ? core : - kMoonEpCombineV2CoreCount + core; -} - -TILEXR_MOONEP_COMBINE_V2_INLINE uint32_t -MoonEpCombineV2LaneForPosition(uint32_t position) -{ - return (position & 3U) == 3U ? MOONEP_COMBINE_V2_TWO_PORT : - MOONEP_COMBINE_V2_SIX_PORT; -} - -TILEXR_MOONEP_COMBINE_V2_INLINE uint32_t -MoonEpCombineV2ControlWqesPerLane( - uint32_t step, uint32_t stepCount, bool finalBatch) -{ - return !finalBatch ? 0U : - (step + 1U < stepCount ? 2U : 1U); -} - -TILEXR_MOONEP_COMBINE_V2_INLINE MoonEpCombineV2LaneCounts -MoonEpCombineV2BatchLaneCounts(uint32_t logicalRows, - uint32_t sequencePhase, uint32_t step, uint32_t stepCount, - bool finalBatch) -{ - const uint32_t firstTwoPort = - (3U - (sequencePhase & 3U)) & 3U; - const uint32_t twoPort = logicalRows <= firstTwoPort ? 0U : - 1U + (logicalRows - 1U - firstTwoPort) / 4U; - const uint32_t control = - MoonEpCombineV2ControlWqesPerLane(step, stepCount, finalBatch); - return MoonEpCombineV2LaneCounts { - logicalRows - twoPort + control, twoPort + control}; -} - -TILEXR_MOONEP_COMBINE_V2_INLINE MoonEpCombineV2RingSegments -MoonEpCombineV2SplitRingCopy(uint32_t absoluteHead, - uint32_t count, uint32_t ringEntries) -{ - const uint32_t ringIndex = ringEntries == 0U ? 0U : - absoluteHead % ringEntries; - const uint32_t untilEnd = ringEntries - ringIndex; - const uint32_t first = count < untilEnd ? count : untilEnd; - return ringEntries == 0U ? MoonEpCombineV2RingSegments {0U, 0U} : - MoonEpCombineV2RingSegments {first, count - first}; -} - -TILEXR_MOONEP_COMBINE_V2_INLINE uint32_t -MoonEpCombineV2NextCqTarget( - uint32_t cqTail, bool finalBatch) -{ - return finalBatch ? cqTail + 1U : cqTail; -} - -TILEXR_MOONEP_COMBINE_V2_INLINE uint32_t -MoonEpCombineV2CompletionCount(bool finalBatch) -{ - return finalBatch ? 1U : 0U; -} - -TILEXR_MOONEP_COMBINE_V2_INLINE bool -MoonEpCombineV2CqTargetReached( - uint32_t cqTail, uint32_t cqTarget) -{ - return cqTail == cqTarget; -} - -TILEXR_MOONEP_COMBINE_V2_INLINE bool -MoonEpCombineV2ShapeValid( - int64_t bs, int64_t h, int64_t topK, int64_t nvS) -{ - return bs > 0 && h > 0 && topK > 0 && topK <= 32 && nvS > 0 && - bs <= nvS / topK && nvS <= std::numeric_limits::max(); -} - -TILEXR_MOONEP_COMBINE_V2_INLINE bool -MoonEpCombineV2DestinationValid( - int32_t encodedDestination, uint64_t slots, uint32_t rankSize) -{ - return MoonEpCombineV2RankSizeSupported(rankSize) && - slots != 0U && - (encodedDestination == -1 || - (encodedDestination >= 0 && - static_cast(encodedDestination) < - static_cast(rankSize) * slots)); -} - -TILEXR_MOONEP_COMBINE_V2_INLINE uint32_t -MoonEpCombineV2Peer( - uint32_t sourceRank, uint32_t step, uint32_t core, uint32_t rankSize) -{ - if (rankSize <= kMoonEpCombineV2GroupSize) { - return core; - } - const uint32_t groupsPerHalf = MoonEpCombineV2GroupsPerHalf(rankSize); - const uint32_t sourceGroup = - (sourceRank / kMoonEpCombineV2GroupSize) % - groupsPerHalf; - const uint32_t targetHalf = core / kMoonEpCombineV2GroupSize; - const uint32_t targetOffset = core % kMoonEpCombineV2GroupSize; - const uint32_t distance = (step + 1U) % groupsPerHalf; - const uint32_t targetIndex = core < kMoonEpCombineV2GroupSize ? - (sourceGroup + distance) % groupsPerHalf : - (sourceGroup + groupsPerHalf - distance) % groupsPerHalf; - return (targetIndex + targetHalf * groupsPerHalf) * - kMoonEpCombineV2GroupSize + targetOffset; -} - -TILEXR_MOONEP_COMBINE_V2_INLINE uint32_t -MoonEpCombineV2Successor( - uint32_t sourceRank, uint32_t core, uint32_t rankSize) -{ - if (rankSize <= kMoonEpCombineV2GroupSize) { - return sourceRank; - } - const uint32_t halfRankCount = rankSize / 2U; - const uint32_t groupsPerHalf = MoonEpCombineV2GroupsPerHalf(rankSize); - const uint32_t halfBase = sourceRank / halfRankCount * halfRankCount; - const uint32_t groupInHalf = (sourceRank % halfRankCount) / - kMoonEpCombineV2GroupSize; - const uint32_t successorGroup = core < kMoonEpCombineV2GroupSize ? - (groupInHalf + groupsPerHalf - 1U) % groupsPerHalf : - (groupInHalf + 1U) % groupsPerHalf; - return halfBase + successorGroup * kMoonEpCombineV2GroupSize + - sourceRank % kMoonEpCombineV2GroupSize; -} - -TILEXR_MOONEP_COMBINE_V2_INLINE uint32_t -MoonEpCombineV2ReceiveStep( - uint32_t destinationRank, uint32_t sourceRank, uint32_t rankSize) -{ - if (rankSize <= kMoonEpCombineV2GroupSize) { - return 0U; - } - const uint32_t groupsPerHalf = MoonEpCombineV2GroupsPerHalf(rankSize); - const uint32_t halfRankCount = rankSize / 2U; - const uint32_t destinationIndex = - (destinationRank / kMoonEpCombineV2GroupSize) % - groupsPerHalf; - const uint32_t sourceIndex = - (sourceRank / kMoonEpCombineV2GroupSize) % - groupsPerHalf; - const uint32_t delta = destinationRank < halfRankCount ? - (destinationIndex + groupsPerHalf - sourceIndex) % groupsPerHalf : - (sourceIndex + groupsPerHalf - destinationIndex) % groupsPerHalf; - const uint32_t distance = delta == 0U ? groupsPerHalf : delta; - return distance - 1U; -} - -TILEXR_MOONEP_COMBINE_V2_INLINE uint32_t -MoonEpCombineV2SourceForCore( - uint32_t core, uint32_t sourceIndex, uint32_t rankSize) -{ - return core + sourceIndex * MoonEpCombineV2ActiveCoreCount(rankSize); -} - -TILEXR_MOONEP_COMBINE_V2_INLINE uint64_t -MoonEpCombineV2Token( - uint64_t magic, uint32_t step) -{ - return (magic << 3U) | static_cast(step); -} - -TILEXR_MOONEP_COMBINE_V2_INLINE bool -MoonEpCombineV2TokenMatches( - uint64_t token, uint64_t magic, uint32_t step, uint32_t rankSize) -{ - return MoonEpCombineV2MagicValid(magic) && - MoonEpCombineV2StepValid(step, rankSize) && - token == MoonEpCombineV2Token(magic, step); -} - -TILEXR_MOONEP_COMBINE_V2_INLINE uint64_t -MoonEpCombineV2DoneIndex( - uint32_t epoch, uint32_t sourceRank, uint32_t lane) -{ - return (static_cast(epoch) * kMoonEpCombineV2RankCount + - sourceRank) * kMoonEpCombineV2LaneCount + lane; -} - -TILEXR_MOONEP_COMBINE_V2_INLINE uint64_t -MoonEpCombineV2GrantIndex( - uint32_t epoch, uint32_t core, uint32_t lane, uint32_t step) -{ - return (((static_cast(epoch) * kMoonEpCombineV2CoreCount + - core) * kMoonEpCombineV2LaneCount + lane) * - kMoonEpCombineV2GrantStepCount) + (step - 1U); -} - -TILEXR_MOONEP_COMBINE_V2_INLINE uint64_t -MoonEpCombineV2FailureIndex( - uint32_t epoch, uint32_t core) -{ - return static_cast(epoch) * kMoonEpCombineV2CoreCount + core; -} - -} // namespace TileXRMoonEp - -#undef TILEXR_MOONEP_COMBINE_V2_INLINE - -#endif // TILEXR_MOONEP_COMBINE_V2_SCHEDULE_H +#endif // TILEXR_MOONEP_COMBINE_V2_SCHEDULE_WRAPPER_H diff --git a/src/moonep/combine_v2/kernels/tilexr_moonep_combine_v2_kernel.h b/src/moonep/combine_v2/kernels/tilexr_moonep_combine_v2_kernel.h index 923507b..6137c39 100644 --- a/src/moonep/combine_v2/kernels/tilexr_moonep_combine_v2_kernel.h +++ b/src/moonep/combine_v2/kernels/tilexr_moonep_combine_v2_kernel.h @@ -1276,6 +1276,9 @@ __aicore__ inline bool MoonEpCombineV2::ReduceHidden() const int64_t tileElements = h_ - hiddenOffset < static_cast(kReduceTileElements) ? h_ - hiddenOffset : static_cast(kReduceTileElements); + const uint32_t inputStrideElements = + TileXRMoonEp::MoonEpCombineV2ReduceInputStrideElements( + static_cast(tileElements)); Duplicate(accumulator, 0.0f, static_cast(tileElements)); PipeBarrier(); @@ -1298,7 +1301,7 @@ __aicore__ inline bool MoonEpCombineV2::ReduceHidden() tileElements * sizeof(bfloat16_t)), 0U, 0U, 0U}; const DataCopyPadExtParams pad { false, 0U, 0U, 0U}; - DataCopyPad(inputRows[batchRoute * tileElements], + DataCopyPad(inputRows[batchRoute * inputStrideElements], input, copyIn, pad); } reduceInputQueue_.EnQue(inputRows); @@ -1306,7 +1309,7 @@ __aicore__ inline bool MoonEpCombineV2::ReduceHidden() for (int64_t batchRoute = 0; batchRoute < batchRoutes; ++batchRoute) { - Cast(row, inputRows[batchRoute * tileElements], + Cast(row, inputRows[batchRoute * inputStrideElements], RoundMode::CAST_NONE, static_cast(tileElements)); PipeBarrier(); diff --git a/src/moonep/common/moonep_combine_schedule.h b/src/moonep/common/moonep_combine_schedule.h new file mode 100644 index 0000000..9f3b002 --- /dev/null +++ b/src/moonep/common/moonep_combine_schedule.h @@ -0,0 +1,376 @@ +#ifndef TILEXR_MOONEP_COMBINE_V2_SCHEDULE_H +#define TILEXR_MOONEP_COMBINE_V2_SCHEDULE_H + +#include +#include + +#if defined(__CCE__) && defined(__CCE_IS_AICORE__) +#define TILEXR_MOONEP_COMBINE_V2_INLINE \ + __attribute__((always_inline)) inline __aicore__ +#else +#define TILEXR_MOONEP_COMBINE_V2_INLINE inline constexpr +#endif + +namespace TileXRMoonEp { + +constexpr uint32_t kMoonEpCombineV2RankCount = 128U; +constexpr uint32_t kMoonEpCombineV2GroupSize = 8U; +constexpr uint32_t kMoonEpCombineV2GroupCount = 16U; +constexpr uint32_t kMoonEpCombineV2GroupsPerHalf = 8U; +constexpr uint32_t kMoonEpCombineV2StepCount = 8U; +constexpr uint32_t kMoonEpCombineV2GrantStepCount = + kMoonEpCombineV2StepCount - 1U; +constexpr uint32_t kMoonEpCombineV2CoreCount = 16U; +constexpr uint32_t kMoonEpCombineV2LaneCount = 2U; +constexpr uint32_t kMoonEpCombineV2QpCount = 32U; +constexpr uint32_t kMoonEpCombineV2MaxSourcesPerCore = + kMoonEpCombineV2RankCount / kMoonEpCombineV2CoreCount; +constexpr uint32_t kMoonEpCombineV2EpochCount = 2U; +constexpr uint32_t kMoonEpCombineV2LogicalBatchRows = 128U; +constexpr uint32_t kMoonEpCombineV2SixPortRows = 96U; +constexpr uint32_t kMoonEpCombineV2TwoPortRows = 32U; +constexpr uint32_t kMoonEpCombineV2SelectionChunkRows = 16384U; +constexpr uint32_t kMoonEpCombineV2MaxOutstanding = 16384U; +constexpr uint32_t kMoonEpCombineV2MteBlockBytes = 32U; +constexpr int64_t kMoonEpCombineV2SmallBs = 8; +constexpr int64_t kMoonEpCombineV2SmallSlots = 128; +constexpr int64_t kMoonEpCombineV2TargetBs = 8192; +constexpr int64_t kMoonEpCombineV2TargetH = 3584; +constexpr int64_t kMoonEpCombineV2TargetTopK = 16; +constexpr int64_t kMoonEpCombineV2TargetSlots = 131072; +constexpr uint64_t kMoonEpCombineV2TokenStrideBytes = 64U; +constexpr uint64_t kMoonEpCombineV2GrantSlotBytes = 512U; +constexpr uint64_t kMoonEpCombineV2GrantReceiveOffsetBytes = 0U; +constexpr uint64_t kMoonEpCombineV2GrantSourceOffsetBytes = 64U; +constexpr uint32_t kMoonEpCombineV2FailureMarker = 0x47505632U; +constexpr uint64_t kMoonEpCombineV2MaxMagic = + std::numeric_limits::max() >> 3U; + +enum MoonEpCombineV2Lane : uint32_t { + MOONEP_COMBINE_V2_SIX_PORT = 0U, + MOONEP_COMBINE_V2_TWO_PORT = 1U, +}; + +enum MoonEpCombineV2FailureStatus : uint32_t { + MOONEP_COMBINE_V2_SUCCESS = 0U, + MOONEP_COMBINE_V2_INVALID_CONFIG = 1U, + MOONEP_COMBINE_V2_POISONED = 2U, + MOONEP_COMBINE_V2_OUTSTANDING_LIMIT = 3U, + MOONEP_COMBINE_V2_CQ_TIMEOUT = 4U, + MOONEP_COMBINE_V2_CQ_ERROR = 5U, + MOONEP_COMBINE_V2_GRANT_TIMEOUT = 6U, + MOONEP_COMBINE_V2_DONE_TIMEOUT = 7U, + MOONEP_COMBINE_V2_BAD_DESTINATION = 8U, +}; + +struct alignas(64) MoonEpCombineV2FailureRecord { + uint64_t magic; + uint32_t status; + uint32_t rank; + uint32_t core; + uint32_t step; + uint32_t peer; + uint32_t lane; + uint32_t qp; + uint32_t cqStatus; + uint64_t expected; + uint64_t observed; + uint32_t poison; + uint32_t marker; +}; + +struct MoonEpCombineV2LaneCounts { + uint32_t sixPort; + uint32_t twoPort; +}; + +struct MoonEpCombineV2RingSegments { + uint32_t first; + uint32_t second; +}; + +static_assert(sizeof(MoonEpCombineV2FailureRecord) == 64U, + "MoonEP Combine V2 failure record ABI changed"); +static_assert(kMoonEpCombineV2GrantSourceOffsetBytes + sizeof(uint64_t) <= + kMoonEpCombineV2GrantSlotBytes, + "MoonEP Combine V2 Grant source exceeds its slot"); + +TILEXR_MOONEP_COMBINE_V2_INLINE bool +MoonEpCombineV2RankSizeSupported(uint32_t rankSize) +{ + return (rankSize >= 2U && rankSize <= kMoonEpCombineV2GroupSize) || + rankSize == 16U || rankSize == 32U || rankSize == 64U || + rankSize == kMoonEpCombineV2RankCount; +} + +TILEXR_MOONEP_COMBINE_V2_INLINE bool +MoonEpCombineV2RankValid(uint32_t rank, uint32_t rankSize) +{ + return MoonEpCombineV2RankSizeSupported(rankSize) && rank < rankSize; +} + +TILEXR_MOONEP_COMBINE_V2_INLINE uint32_t +MoonEpCombineV2ActiveCoreCount(uint32_t rankSize) +{ + return rankSize <= kMoonEpCombineV2GroupSize ? rankSize : + kMoonEpCombineV2CoreCount; +} + +TILEXR_MOONEP_COMBINE_V2_INLINE uint32_t +MoonEpCombineV2StepCount(uint32_t rankSize) +{ + return rankSize <= kMoonEpCombineV2GroupSize ? 1U : + rankSize / kMoonEpCombineV2CoreCount; +} + +TILEXR_MOONEP_COMBINE_V2_INLINE uint32_t +MoonEpCombineV2GroupsPerHalf(uint32_t rankSize) +{ + return rankSize / kMoonEpCombineV2CoreCount; +} + +TILEXR_MOONEP_COMBINE_V2_INLINE uint32_t +MoonEpCombineV2LocalRankSize(uint32_t rankSize) +{ + return rankSize <= kMoonEpCombineV2GroupSize ? rankSize : + kMoonEpCombineV2GroupSize; +} + +TILEXR_MOONEP_COMBINE_V2_INLINE bool +MoonEpCombineV2CoreValid(uint32_t core, uint32_t rankSize) +{ + return MoonEpCombineV2RankSizeSupported(rankSize) && + core < MoonEpCombineV2ActiveCoreCount(rankSize); +} + +TILEXR_MOONEP_COMBINE_V2_INLINE bool +MoonEpCombineV2StepValid(uint32_t step, uint32_t rankSize) +{ + return MoonEpCombineV2RankSizeSupported(rankSize) && + step < MoonEpCombineV2StepCount(rankSize); +} + +TILEXR_MOONEP_COMBINE_V2_INLINE bool +MoonEpCombineV2MagicValid(uint64_t magic) +{ + return magic != 0U && magic <= kMoonEpCombineV2MaxMagic; +} + +TILEXR_MOONEP_COMBINE_V2_INLINE uint32_t +MoonEpCombineV2Epoch(uint64_t magic) +{ + return static_cast(magic & 1U); +} + +TILEXR_MOONEP_COMBINE_V2_INLINE uint32_t +MoonEpCombineV2Qp(uint32_t core, uint32_t lane) +{ + return lane == MOONEP_COMBINE_V2_SIX_PORT ? core : + kMoonEpCombineV2CoreCount + core; +} + +TILEXR_MOONEP_COMBINE_V2_INLINE uint32_t +MoonEpCombineV2LaneForPosition(uint32_t position) +{ + return (position & 3U) == 3U ? MOONEP_COMBINE_V2_TWO_PORT : + MOONEP_COMBINE_V2_SIX_PORT; +} + +TILEXR_MOONEP_COMBINE_V2_INLINE uint32_t +MoonEpCombineV2ControlWqesPerLane( + uint32_t step, uint32_t stepCount, bool finalBatch) +{ + return !finalBatch ? 0U : + (step + 1U < stepCount ? 2U : 1U); +} + +TILEXR_MOONEP_COMBINE_V2_INLINE MoonEpCombineV2LaneCounts +MoonEpCombineV2BatchLaneCounts(uint32_t logicalRows, + uint32_t sequencePhase, uint32_t step, uint32_t stepCount, + bool finalBatch) +{ + const uint32_t firstTwoPort = + (3U - (sequencePhase & 3U)) & 3U; + const uint32_t twoPort = logicalRows <= firstTwoPort ? 0U : + 1U + (logicalRows - 1U - firstTwoPort) / 4U; + const uint32_t control = + MoonEpCombineV2ControlWqesPerLane(step, stepCount, finalBatch); + return MoonEpCombineV2LaneCounts { + logicalRows - twoPort + control, twoPort + control}; +} + +TILEXR_MOONEP_COMBINE_V2_INLINE MoonEpCombineV2RingSegments +MoonEpCombineV2SplitRingCopy(uint32_t absoluteHead, + uint32_t count, uint32_t ringEntries) +{ + const uint32_t ringIndex = ringEntries == 0U ? 0U : + absoluteHead % ringEntries; + const uint32_t untilEnd = ringEntries - ringIndex; + const uint32_t first = count < untilEnd ? count : untilEnd; + return ringEntries == 0U ? MoonEpCombineV2RingSegments {0U, 0U} : + MoonEpCombineV2RingSegments {first, count - first}; +} + +TILEXR_MOONEP_COMBINE_V2_INLINE uint32_t +MoonEpCombineV2NextCqTarget( + uint32_t cqTail, bool finalBatch) +{ + return finalBatch ? cqTail + 1U : cqTail; +} + +TILEXR_MOONEP_COMBINE_V2_INLINE uint32_t +MoonEpCombineV2CompletionCount(bool finalBatch) +{ + return finalBatch ? 1U : 0U; +} + +TILEXR_MOONEP_COMBINE_V2_INLINE bool +MoonEpCombineV2CqTargetReached( + uint32_t cqTail, uint32_t cqTarget) +{ + return cqTail == cqTarget; +} + +TILEXR_MOONEP_COMBINE_V2_INLINE bool +MoonEpCombineV2ShapeValid( + int64_t bs, int64_t h, int64_t topK, int64_t nvS) +{ + return bs > 0 && h > 0 && topK > 0 && topK <= 32 && nvS > 0 && + bs <= nvS / topK && nvS <= std::numeric_limits::max(); +} + +TILEXR_MOONEP_COMBINE_V2_INLINE uint32_t +MoonEpCombineV2ReduceInputStrideElements(uint32_t tileElements) +{ + const uint32_t bytes = tileElements * sizeof(uint16_t); + return ((bytes + kMoonEpCombineV2MteBlockBytes - 1U) / + kMoonEpCombineV2MteBlockBytes * kMoonEpCombineV2MteBlockBytes) / + sizeof(uint16_t); +} + +TILEXR_MOONEP_COMBINE_V2_INLINE bool +MoonEpCombineV2DestinationValid( + int32_t encodedDestination, uint64_t slots, uint32_t rankSize) +{ + return MoonEpCombineV2RankSizeSupported(rankSize) && + slots != 0U && + (encodedDestination == -1 || + (encodedDestination >= 0 && + static_cast(encodedDestination) < + static_cast(rankSize) * slots)); +} + +TILEXR_MOONEP_COMBINE_V2_INLINE uint32_t +MoonEpCombineV2Peer( + uint32_t sourceRank, uint32_t step, uint32_t core, uint32_t rankSize) +{ + if (rankSize <= kMoonEpCombineV2GroupSize) { + return core; + } + const uint32_t groupsPerHalf = MoonEpCombineV2GroupsPerHalf(rankSize); + const uint32_t sourceGroup = + (sourceRank / kMoonEpCombineV2GroupSize) % + groupsPerHalf; + const uint32_t targetHalf = core / kMoonEpCombineV2GroupSize; + const uint32_t targetOffset = core % kMoonEpCombineV2GroupSize; + const uint32_t distance = (step + 1U) % groupsPerHalf; + const uint32_t targetIndex = core < kMoonEpCombineV2GroupSize ? + (sourceGroup + distance) % groupsPerHalf : + (sourceGroup + groupsPerHalf - distance) % groupsPerHalf; + return (targetIndex + targetHalf * groupsPerHalf) * + kMoonEpCombineV2GroupSize + targetOffset; +} + +TILEXR_MOONEP_COMBINE_V2_INLINE uint32_t +MoonEpCombineV2Successor( + uint32_t sourceRank, uint32_t core, uint32_t rankSize) +{ + if (rankSize <= kMoonEpCombineV2GroupSize) { + return sourceRank; + } + const uint32_t halfRankCount = rankSize / 2U; + const uint32_t groupsPerHalf = MoonEpCombineV2GroupsPerHalf(rankSize); + const uint32_t halfBase = sourceRank / halfRankCount * halfRankCount; + const uint32_t groupInHalf = (sourceRank % halfRankCount) / + kMoonEpCombineV2GroupSize; + const uint32_t successorGroup = core < kMoonEpCombineV2GroupSize ? + (groupInHalf + groupsPerHalf - 1U) % groupsPerHalf : + (groupInHalf + 1U) % groupsPerHalf; + return halfBase + successorGroup * kMoonEpCombineV2GroupSize + + sourceRank % kMoonEpCombineV2GroupSize; +} + +TILEXR_MOONEP_COMBINE_V2_INLINE uint32_t +MoonEpCombineV2ReceiveStep( + uint32_t destinationRank, uint32_t sourceRank, uint32_t rankSize) +{ + if (rankSize <= kMoonEpCombineV2GroupSize) { + return 0U; + } + const uint32_t groupsPerHalf = MoonEpCombineV2GroupsPerHalf(rankSize); + const uint32_t halfRankCount = rankSize / 2U; + const uint32_t destinationIndex = + (destinationRank / kMoonEpCombineV2GroupSize) % + groupsPerHalf; + const uint32_t sourceIndex = + (sourceRank / kMoonEpCombineV2GroupSize) % + groupsPerHalf; + const uint32_t delta = destinationRank < halfRankCount ? + (destinationIndex + groupsPerHalf - sourceIndex) % groupsPerHalf : + (sourceIndex + groupsPerHalf - destinationIndex) % groupsPerHalf; + const uint32_t distance = delta == 0U ? groupsPerHalf : delta; + return distance - 1U; +} + +TILEXR_MOONEP_COMBINE_V2_INLINE uint32_t +MoonEpCombineV2SourceForCore( + uint32_t core, uint32_t sourceIndex, uint32_t rankSize) +{ + return core + sourceIndex * MoonEpCombineV2ActiveCoreCount(rankSize); +} + +TILEXR_MOONEP_COMBINE_V2_INLINE uint64_t +MoonEpCombineV2Token( + uint64_t magic, uint32_t step) +{ + return (magic << 3U) | static_cast(step); +} + +TILEXR_MOONEP_COMBINE_V2_INLINE bool +MoonEpCombineV2TokenMatches( + uint64_t token, uint64_t magic, uint32_t step, uint32_t rankSize) +{ + return MoonEpCombineV2MagicValid(magic) && + MoonEpCombineV2StepValid(step, rankSize) && + token == MoonEpCombineV2Token(magic, step); +} + +TILEXR_MOONEP_COMBINE_V2_INLINE uint64_t +MoonEpCombineV2DoneIndex( + uint32_t epoch, uint32_t sourceRank, uint32_t lane) +{ + return (static_cast(epoch) * kMoonEpCombineV2RankCount + + sourceRank) * kMoonEpCombineV2LaneCount + lane; +} + +TILEXR_MOONEP_COMBINE_V2_INLINE uint64_t +MoonEpCombineV2GrantIndex( + uint32_t epoch, uint32_t core, uint32_t lane, uint32_t step) +{ + return (((static_cast(epoch) * kMoonEpCombineV2CoreCount + + core) * kMoonEpCombineV2LaneCount + lane) * + kMoonEpCombineV2GrantStepCount) + (step - 1U); +} + +TILEXR_MOONEP_COMBINE_V2_INLINE uint64_t +MoonEpCombineV2FailureIndex( + uint32_t epoch, uint32_t core) +{ + return static_cast(epoch) * kMoonEpCombineV2CoreCount + core; +} + +} // namespace TileXRMoonEp + +#undef TILEXR_MOONEP_COMBINE_V2_INLINE + +#endif // TILEXR_MOONEP_COMBINE_V2_SCHEDULE_H diff --git a/src/moonep/dispatch/urma/host/dispatch_host.cpp b/src/moonep/dispatch/urma/host/dispatch_host.cpp index 92800e9..ad1169c 100644 --- a/src/moonep/dispatch/urma/host/dispatch_host.cpp +++ b/src/moonep/dispatch/urma/host/dispatch_host.cpp @@ -306,11 +306,13 @@ int TileXRMoonEpQueryDispatchUrmaWorkspace(TileXRCommPtr comm, int64_t s, return TILEXR_MOONEP_SUCCESS; } -int TileXRMoonEpRunDispatchUrmaV1(const TileXRMoonEpDispatchArgsV1 *args, - aclrtStream stream) +static int RunDispatchUrma(const TileXRMoonEpDispatchArgsV1 *args, + aclrtStream stream, bool resetStatus) { if (args == nullptr || args->structSize < sizeof(*args) || - args->abiVersion != TILEXR_MOONEP_ABI_VERSION_V1 || stream == nullptr || + (args->abiVersion != TILEXR_MOONEP_ABI_VERSION_V1 && + args->abiVersion != TILEXR_MOONEP_ABI_VERSION_V2) || + stream == nullptr || args->comm == nullptr || args->plan == nullptr || args->hiddenSh == nullptr || args->hiddenNvsh == nullptr || args->flags != TILEXR_MOONEP_FLAG_NONE || !PlanValid(args->plan) || !TensorDescriptorValid(args->hiddenSh) || @@ -438,6 +440,11 @@ int TileXRMoonEpRunDispatchUrmaV1(const TileXRMoonEpDispatchArgsV1 *args, } } + if (resetStatus && aclrtMemsetAsync(args->plan->status, + sizeof(int32_t), 0, sizeof(int32_t), stream) != ACL_SUCCESS) { + return TILEXR_MOONEP_ERROR_INTERNAL; + } + DispatchUrmaLaunchParams params {}; params.commArgs = devArgs; params.dst = static_cast(args->plan->dst); @@ -465,4 +472,16 @@ int TileXRMoonEpRunDispatchUrmaV1(const TileXRMoonEpDispatchArgsV1 *args, return MapLaunchStatus(TileXRMoonEpLaunchDispatchUrmaKernel(params)); } +int TileXRMoonEpRunDispatchUrmaV1(const TileXRMoonEpDispatchArgsV1 *args, + aclrtStream stream) +{ + return RunDispatchUrma(args, stream, false); +} + +int TileXRMoonEpRunDispatchUrmaV2(const TileXRMoonEpDispatchArgsV2 *args, + aclrtStream stream) +{ + return RunDispatchUrma(args, stream, true); +} + } // namespace TileXRMoonEp diff --git a/src/moonep/host/tilexr_moonep.cpp b/src/moonep/host/tilexr_moonep.cpp index b574f9c..3ca5e54 100644 --- a/src/moonep/host/tilexr_moonep.cpp +++ b/src/moonep/host/tilexr_moonep.cpp @@ -17,6 +17,8 @@ int TileXRMoonEpQueryDispatchUrmaWorkspace(TileXRCommPtr comm, int64_t s, uint64_t *workspaceAlignment); int TileXRMoonEpRunDispatchUrmaV1( const TileXRMoonEpDispatchArgsV1 *args, aclrtStream stream); +int TileXRMoonEpRunDispatchUrmaV2( + const TileXRMoonEpDispatchArgsV2 *args, aclrtStream stream); int TileXRMoonEpRunCombineV1( const TileXRMoonEpCombineArgsV1 *args, aclrtStream stream); int TileXRMoonEpRunPrefetchWeightV1( @@ -328,6 +330,26 @@ extern "C" int TileXRMoonEpDispatchGetWorkspaceSizeV1(TileXRCommPtr comm, comm, s, k, h, hiddenDtype, workspaceBytes, workspaceAlignment); } +extern "C" int TileXRMoonEpDispatchGetWorkspaceSizeV2(TileXRCommPtr comm, + int64_t s, int64_t k, int64_t h, uint32_t hiddenDtype, + uint64_t *workspaceBytes, uint64_t *workspaceAlignment) +{ + return TileXRMoonEp::TileXRMoonEpQueryDispatchUrmaWorkspace( + comm, s, k, h, hiddenDtype, workspaceBytes, workspaceAlignment); +} + +extern "C" int TileXRMoonEpDispatchV2(const TileXRMoonEpDispatchArgsV2 *args, + aclrtStream stream) +{ + if (args == nullptr || args->structSize < sizeof(*args) || + args->abiVersion != TILEXR_MOONEP_ABI_VERSION_V2 || + args->registeredWorkspace == nullptr || + args->registeredWorkspaceBytes == 0) { + return TILEXR_MOONEP_ERROR_INVALID_ARGUMENT; + } + return TileXRMoonEp::TileXRMoonEpRunDispatchUrmaV2(args, stream); +} + extern "C" int TileXRMoonEpPrefetchWeightV1( const TileXRMoonEpPrefetchWeightArgsV1 *args, aclrtStream stream) { diff --git a/src/moonep/prefetch_weight/host/prefetch_weight_layout.cpp b/src/moonep/prefetch_weight/host/prefetch_weight_layout.cpp index b563749..b7ec8cf 100644 --- a/src/moonep/prefetch_weight/host/prefetch_weight_layout.cpp +++ b/src/moonep/prefetch_weight/host/prefetch_weight_layout.cpp @@ -120,7 +120,7 @@ int TileXRMoonEpBuildPrefetchWeightLayout( commArgs.udmaInfoPtr == nullptr || commArgs.udmaRegistryPtr == nullptr || !TileXR::UDMARegistryValid(®istry, commArgs.rankSize) || commArgs.rank < 0 || commArgs.rank >= commArgs.rankSize || - !SupportedWorkerCount(qpNum) || qpNum > kPrefetchWeightMaxWorkers) { + qpNum == 0) { return TILEXR_MOONEP_ERROR_INVALID_ARGUMENT; } @@ -147,7 +147,8 @@ int TileXRMoonEpBuildPrefetchWeightLayout( } const bool hasOverride = blockDimOverride != nullptr && blockDimOverride[0] != '\0'; - uint32_t workers = qpNum; + uint32_t workers = qpNum < kPrefetchWeightMaxWorkers ? + qpNum : kPrefetchWeightMaxWorkers; if (!ParseWorkerOverride(blockDimOverride, &workers) || workers > qpNum || (hasOverride && workers > static_cast(args.plan->b))) { return TILEXR_MOONEP_ERROR_INVALID_ARGUMENT; diff --git a/tests/moonep/python/test_ffi_unittest.py b/tests/moonep/python/test_ffi_unittest.py index dde5a13..cd93687 100644 --- a/tests/moonep/python/test_ffi_unittest.py +++ b/tests/moonep/python/test_ffi_unittest.py @@ -1,9 +1,12 @@ from __future__ import annotations import ctypes +import os import sys import unittest from pathlib import Path +from types import SimpleNamespace +from unittest.mock import patch ROOT = Path(__file__).resolve().parents[3] @@ -17,6 +20,7 @@ from tilexr_moonep.abi import ( TileXRMoonEPCombineArgsV1, TileXRMoonEPDispatchArgsV1, + TileXRMoonEPDispatchArgsV2, TileXRMoonEPPlanV1, TileXRMoonEPPlanningArgsV1, TileXRMoonEPPrefetchWeightArgsV1, @@ -253,20 +257,46 @@ def dispatch_workspace(comm, s, k, h, dtype, workspace_bytes, alignment): library.TileXRMoonEpGetCapabilitiesV1 = FakeFunction(capabilities) library.TileXRMoonEpGetCapabilitiesV2 = FakeFunction(capabilities) library.TileXRMoonEpPlanningGetWorkspaceSizeV1 = FakeFunction(workspace) - library.TileXRMoonEpDispatchGetWorkspaceSizeV1 = FakeFunction(dispatch_workspace) + library.TileXRMoonEpDispatchGetWorkspaceSizeV2 = FakeFunction(dispatch_workspace) library.TileXRMoonEpPlanningV1 = FakeFunction(planning) for name, args_type in ( - ("dispatch", TileXRMoonEPDispatchArgsV1), + ("dispatch", TileXRMoonEPDispatchArgsV2), ("prefetch_weight", TileXRMoonEPPrefetchWeightArgsV1), ): setattr( library, { - "dispatch": "TileXRMoonEpDispatchV1", + "dispatch": "TileXRMoonEpDispatchV2", "prefetch_weight": "TileXRMoonEpPrefetchWeightV1", }[name], FakeFunction(self._stage_callback(name, args_type)), ) + + def combine_v1(args_ptr, stream): + args = ctypes.cast( + args_ptr, ctypes.POINTER(TileXRMoonEPCombineArgsV1) + ).contents + plan = args.plan.contents + self.stage_records.append({ + "name": "combine_v1", + "stream": stream.value, + "flags": args.flags, + "plan_capacity": plan.nvS, + "dst_local": int(args.dstLocal), + "shapes": { + name: None if not getattr(args, name) else tuple( + getattr(args, name).contents.shape[ + : getattr(args, name).contents.rank + ] + ) + for name in ( + "hiddenNvsh", "routeWeightsNvs", "hiddenSh", "routeWeightsSk" + ) + }, + }) + return 0 + + library.TileXRMoonEpCombineV1 = FakeFunction(combine_v1) def reduce_grad_query(query_ptr, info_ptr): query = ctypes.cast( query_ptr, ctypes.POINTER(TileXRMoonEPReduceGradWorkspaceQueryV2) @@ -330,6 +360,7 @@ def callback(args_ptr, stream): "name": name, "stream": stream.value, "flags": args.flags, + "abi_version": args.abiVersion, "plan_capacity": plan.nvS, "shapes": shapes, }) @@ -339,9 +370,16 @@ def callback(args_ptr, stream): def __call__(self, path, mode): self.loads.append((str(path), mode)) - return (self.comm, self.planner, self.combine_v2, self.moonep)[ - len(self.loads) - 1 - ] + name = Path(path).name + if name == "libtile-comm.so": + return self.comm + if "planner" in name: + return self.planner + if "combine-v2" in name: + return self.combine_v2 + if "moonep" in name: + return self.moonep + raise AssertionError(f"unexpected library path {path}") def tensor(shape, dtype): @@ -354,8 +392,10 @@ def test_ctypes_layout_matches_tilexr_moonep_header(self): self.assertEqual(ctypes.sizeof(TileXRMoonEPPlanV1), 120) self.assertEqual(ctypes.sizeof(TileXRMoonEPPlanningArgsV1), 80) self.assertEqual(ctypes.sizeof(TileXRMoonEPDispatchArgsV1), 80) + self.assertEqual(ctypes.sizeof(TileXRMoonEPDispatchArgsV2), 80) self.assertEqual(ctypes.sizeof(TileXRMoonEPPrefetchWeightArgsV1), 56) - self.assertEqual(ctypes.sizeof(TileXRMoonEPCombineArgsV1), 64) + self.assertEqual(ctypes.sizeof(TileXRMoonEPCombineArgsV1), 72) + self.assertEqual(TileXRMoonEPCombineArgsV1.dstLocal.offset, 24) self.assertEqual(ctypes.sizeof(TileXRMoonEPReduceGradArgsV1), 48) self.assertEqual(ctypes.sizeof(TileXRMoonEPReduceGradWorkspaceQueryV2), 64) self.assertEqual(ctypes.sizeof(TileXRMoonEPReduceGradWorkspaceInfoV2), 96) @@ -368,6 +408,88 @@ def test_ctypes_layout_matches_tilexr_moonep_header(self): self.assertEqual(TileXRMoonEPDispatchArgsV1.flags.offset, 56) self.assertEqual(TileXRMoonEPDispatchArgsV1.registeredWorkspace.offset, 64) self.assertEqual(TileXRMoonEPDispatchArgsV1.registeredWorkspaceBytes.offset, 72) + self.assertEqual(TileXRMoonEPDispatchArgsV2.registeredWorkspace.offset, 64) + + def test_invalid_combine_version_fails_before_library_load(self): + loader = FakeCDLLLoader() + with patch.dict(os.environ, {"TILEXR_MOONEP_COMBINE_VERSION": "3"}): + with self.assertRaisesRegex(ValueError, "TILEXR_MOONEP_COMBINE_VERSION"): + TileXRMoonEPRuntime( + rank=0, + world_size=2, + library_paths={ + "comm": "libtile-comm.so", + "planner": "libtilexr-moonep-planner.so", + "combine_v2": "libtilexr-moonep-combine-v2.so.2", + "moonep": "libtilexr-moonep.so.1", + }, + cdll_loader=loader, + ) + self.assertEqual(loader.loads, []) + + def test_explicit_combine_v1_uses_one_descriptor_call(self): + loader = FakeCDLLLoader() + with patch.dict(os.environ, {"TILEXR_MOONEP_COMBINE_VERSION": "1"}): + runtime = TileXRMoonEPRuntime( + rank=0, + world_size=2, + library_paths={ + "comm": "libtile-comm.so", + "planner": "libtilexr-moonep-planner.so", + "combine_v2": "libtilexr-moonep-combine-v2.so.2", + "moonep": "libtilexr-moonep.so.1", + }, + cdll_loader=loader, + ) + torch = FakeTorch() + context = SimpleNamespace( + tokens_per_rank=4, + hidden_size=8, + topk=2, + expert_count=4, + prefetch_slots=2, + nv_s=12, + dtype=torch.bfloat16, + ) + plan = SimpleNamespace( + n=8, + rank_size=2, + expert_count=4, + prefetch_slots=2, + nv_s=12, + topk=2, + dst=tensor((8,), torch.int32), + experts_to_copy=tensor((4,), torch.int32), + zero_fill_ranges=tensor((6, 2), torch.int32), + remote_stats=tensor((2,), torch.int32), + dup_groups=tensor((12, 3), torch.int32), + dup_loffs=tensor((12,), torch.int32), + dup_counts=tensor((2,), torch.int32), + status=tensor((1,), torch.int32), + workspace=tensor((256,), torch.uint8), + dst_local_offset=128, + ) + runtime.combine( + context, + plan, + tensor((12, 8), torch.bfloat16), + tensor((4, 8), torch.bfloat16), + 0xCAFE, + tensor((12,), torch.float32), + tensor((4, 2), torch.float32), + inter_rank_sync=True, + ) + self.assertEqual(runtime.combine_version, 1) + self.assertNotIn( + "libtilexr-moonep-combine-v2.so.2", + [Path(path).name for path, _ in loader.loads], + ) + records = [record for record in loader.stage_records if record["name"] == "combine_v1"] + self.assertEqual(len(records), 1) + self.assertEqual(records[0]["dst_local"], plan.workspace.data_ptr() + 128) + self.assertEqual(records[0]["shapes"]["hiddenNvsh"], (12, 8)) + self.assertEqual(records[0]["shapes"]["routeWeightsSk"], (4, 2)) + runtime.close() def test_fake_cdll_receives_v1_descriptors_and_combine_v2_pointers(self): loader = FakeCDLLLoader() @@ -495,6 +617,8 @@ def test_fake_cdll_receives_v1_descriptors_and_combine_v2_pointers(self): ) self.assertTrue(all(record["stream"] == 0xCAFE for record in loader.stage_records)) self.assertEqual(loader.stage_records[0]["flags"], 0) + self.assertEqual(loader.stage_records[0]["abi_version"], 2) + self.assertEqual(loader.stage_records[1]["abi_version"], 2) self.assertTrue(all(record["flags"] == 0 for record in loader.stage_records[1:])) self.assertEqual(loader.stage_records[0]["shapes"]["hiddenNvsh"], (12, 8)) self.assertEqual(loader.stage_records[0]["shapes"]["routeWeightsNvs"], (12,)) diff --git a/tests/moonep/python/test_moonep_modes.py b/tests/moonep/python/test_moonep_modes.py index ac5a0c8..02646b0 100644 --- a/tests/moonep/python/test_moonep_modes.py +++ b/tests/moonep/python/test_moonep_modes.py @@ -169,6 +169,10 @@ def combine(*args, **kwargs): assert torch.count_nonzero(buffer.combine_inputs[0][2:]).item() == 0 +def test_v2_full_flow_uses_dispatch_status_as_final_shared_status() -> None: + assert benchmark.FINAL_SHARED_STATUS_SUCCESS == 0 + + def test_case_aliases_add_reference_dimensions_without_changing_old_defaults() -> None: old = BenchmarkCase("old", 4, 2, 4, 8) assert old.intermediate_size is None diff --git a/tests/moonep/python/test_moonep_performance_report.py b/tests/moonep/python/test_moonep_performance_report.py index 7c49c0c..aeb702f 100644 --- a/tests/moonep/python/test_moonep_performance_report.py +++ b/tests/moonep/python/test_moonep_performance_report.py @@ -66,6 +66,24 @@ def test_algorithm_bytes_follow_routes_and_runtime_tensor_dtypes() -> None: } +def test_stage_metadata_reports_selected_combine_v1_memory_kernel() -> None: + capabilities = { + "implementations": { + "planning": "native", + "dispatch": "native", + "prefetch_weight": "native", + "combine": "native", + "reduce_grad": "native", + } + } + metadata = stage_execution_metadata( + capabilities, torch_npu_version="test", combine_version=1 + ) + assert metadata["combine"]["kernel_version"] == ( + "tilexr_moonep_combine_kernel (CombineV1Memory)" + ) + + def _result(rank: int, world_size: int, node_count: int) -> dict[str, object]: capabilities = { "stage_mask": 31, @@ -261,7 +279,8 @@ def test_flow_report_uses_stage_critical_rank_and_its_bytes(tmp_path: Path) -> N assert "Native" in table assert "Kernel/API version" in table assert "tilexr_ep_plan_kernel (PlannerV3)" in table - assert "tilexr_moonep_dispatch_urma_kernel (DispatchV1)" in table + assert "tilexr_moonep_dispatch_urma_kernel (DispatchV2)" in table + assert "tilexr_moonep_combine_v2_kernel (CombineV2)" in table assert "torch_npu 2.10.0.post2 (GMM+SwiGLU)" in table stage_csv_header = (case_dir / "stage_summary.csv").read_text( encoding="utf-8" diff --git a/tests/moonep/unit/test_tilexr_moonep_abi_layout.cpp b/tests/moonep/unit/test_tilexr_moonep_abi_layout.cpp index 1f27c2e..0a9ba63 100644 --- a/tests/moonep/unit/test_tilexr_moonep_abi_layout.cpp +++ b/tests/moonep/unit/test_tilexr_moonep_abi_layout.cpp @@ -54,13 +54,21 @@ int main() "Unexpected DispatchArgsV1 registered workspace offset"); static_assert(offsetof(TileXRMoonEpDispatchArgsV1, registeredWorkspaceBytes) == 72, "Unexpected DispatchArgsV1 registered workspace size offset"); + static_assert(std::is_standard_layout::value, + "Dispatch V2 args must be standard layout"); + static_assert(sizeof(TileXRMoonEpDispatchArgsV2) == 80, + "Unexpected DispatchArgsV2 size"); + static_assert(offsetof(TileXRMoonEpDispatchArgsV2, registeredWorkspace) == 64, + "Unexpected DispatchArgsV2 registered workspace offset"); static_assert(sizeof(TileXRMoonEpPrefetchWeightArgsV1) == 56, "Unexpected PrefetchWeightArgsV1 size"); static_assert(offsetof(TileXRMoonEpPrefetchWeightArgsV1, gate) == 24, "Unexpected PrefetchWeightArgsV1 gate offset"); - static_assert(sizeof(TileXRMoonEpCombineArgsV1) == 64, + static_assert(sizeof(TileXRMoonEpCombineArgsV1) == 72, "Unexpected CombineArgsV1 size"); - static_assert(offsetof(TileXRMoonEpCombineArgsV1, hiddenNvsh) == 24, + static_assert(offsetof(TileXRMoonEpCombineArgsV1, dstLocal) == 24, + "Unexpected CombineArgsV1 reverse-route offset"); + static_assert(offsetof(TileXRMoonEpCombineArgsV1, hiddenNvsh) == 32, "Unexpected CombineArgsV1 hidden input offset"); static_assert(sizeof(TileXRMoonEpReduceGradArgsV1) == 48, "Unexpected ReduceGradArgsV1 size"); diff --git a/tests/moonep/unit/test_tilexr_moonep_c_header.c b/tests/moonep/unit/test_tilexr_moonep_c_header.c index f3d9f4f..c8b880b 100644 --- a/tests/moonep/unit/test_tilexr_moonep_c_header.c +++ b/tests/moonep/unit/test_tilexr_moonep_c_header.c @@ -7,6 +7,7 @@ int main(void) TileXRMoonEpTensorV1 tensor = {0}; TileXRMoonEpPlanV1 plan = {0}; TileXRMoonEpDispatchArgsV1 dispatch = {0}; + TileXRMoonEpDispatchArgsV2 dispatchV2 = {0}; TileXRMoonEpPrefetchWeightArgsV1 prefetch = {0}; TileXRMoonEpCombineArgsV1 combine = {0}; TileXRMoonEpReduceGradArgsV1 reduce = {0}; @@ -24,7 +25,10 @@ int main(void) dispatch.structSize = (uint32_t)sizeof(dispatch); dispatch.abiVersion = TILEXR_MOONEP_ABI_VERSION_V1; dispatch.hiddenSh = &tensor; + dispatchV2.abiVersion = TILEXR_MOONEP_ABI_VERSION_V2; + dispatchV2.hiddenSh = &tensor; prefetch.gate = &tensor; + combine.dstLocal = (const int32_t *)(uintptr_t)0x1000; combine.hiddenNvsh = &tensor; reduce.input = &tensor; reduce.output = &tensor; @@ -34,7 +38,9 @@ int main(void) return tensor.dtype == TILEXR_MOONEP_DTYPE_FLOAT16 && plan.nvS == 1 && dispatch.hiddenSh == &tensor && - prefetch.gate == &tensor && combine.hiddenNvsh == &tensor && + dispatchV2.hiddenSh == &tensor && + prefetch.gate == &tensor && combine.dstLocal != 0 && + combine.hiddenNvsh == &tensor && reduce.input == &tensor && reduce.output == &tensor && query.abiVersion == TILEXR_MOONEP_ABI_VERSION_V2 && info.abiVersion == TILEXR_MOONEP_ABI_VERSION_V2 && diff --git a/tests/moonep/unit/test_tilexr_moonep_combine_host.cpp b/tests/moonep/unit/test_tilexr_moonep_combine_host.cpp index a302e6f..6f16f88 100644 --- a/tests/moonep/unit/test_tilexr_moonep_combine_host.cpp +++ b/tests/moonep/unit/test_tilexr_moonep_combine_host.cpp @@ -116,6 +116,7 @@ TileXRMoonEpCombineArgsV1 Args(const TileXRMoonEpPlanV1 *plan, args.abiVersion = TILEXR_MOONEP_ABI_VERSION_V1; args.comm = reinterpret_cast(uintptr_t {0x1000}); args.plan = plan; + args.dstLocal = reinterpret_cast(uintptr_t {0x3800}); args.hiddenNvsh = hiddenNvsh; args.routeWeightsNvs = weightsNvs; args.hiddenSh = hiddenSh; @@ -137,14 +138,15 @@ void TestPairedLaunch() TileXRMoonEpTensorV1 weightsSk = Tensor(reinterpret_cast(uintptr_t {0x7000}), TILEXR_MOONEP_DTYPE_FLOAT32, 2, 2, 2, 4); TileXRMoonEpCombineArgsV1 args = Args(&plan, &hiddenNvsh, &weightsNvs, - &hiddenSh, &weightsSk, TILEXR_MOONEP_FLAG_SKIP_INTER_RANK_SYNC); + &hiddenSh, &weightsSk, TILEXR_MOONEP_FLAG_NONE); aclrtStream stream = reinterpret_cast(uintptr_t {0x8000}); CheckStatus("paired combine", TileXRMoonEp::TileXRMoonEpRunCombineV1(&args, stream), TILEXR_MOONEP_SUCCESS); Check(hostCalls == 1 && devCalls == 1 && magicCalls == 1 && launchCalls == 1, "paired combine call counts mismatch"); - Check(launchedParams.dst == plan.dst && launchedParams.dupGroups == plan.dupGroups && + Check(launchedParams.dstLocal == args.dstLocal && + launchedParams.dst == plan.dst && launchedParams.dupGroups == plan.dupGroups && launchedParams.dupLoffs == plan.dupLoffs && launchedParams.dupCounts == plan.dupCounts && launchedParams.hiddenNvsh == hiddenNvsh.data && launchedParams.hiddenSh == hiddenSh.data && launchedParams.routeWeightsNvs == weightsNvs.data && @@ -182,6 +184,10 @@ void TestValidationAndFailures() CheckStatus("unpaired weights", TileXRMoonEp::TileXRMoonEpRunCombineV1(&args, stream), TILEXR_MOONEP_ERROR_INVALID_ARGUMENT); args.routeWeightsNvs = nullptr; + args.dstLocal = nullptr; + CheckStatus("missing reverse route", TileXRMoonEp::TileXRMoonEpRunCombineV1( + &args, stream), TILEXR_MOONEP_ERROR_INVALID_ARGUMENT); + args.dstLocal = reinterpret_cast(uintptr_t {0x3800}); args.flags = TILEXR_MOONEP_FLAG_BUILD_DEDUP; CheckStatus("build dedup", TileXRMoonEp::TileXRMoonEpRunCombineV1(&args, stream), TILEXR_MOONEP_ERROR_INVALID_ARGUMENT); @@ -234,7 +240,7 @@ void TestValidationAndFailures() CheckStatus("launch error", TileXRMoonEp::TileXRMoonEpRunCombineV1(&args, stream), -54); } -void TestSplitPhaseLaunch() +void TestSplitPhaseRejected() { Reset(); TileXRMoonEpPlanV1 plan = ValidPlan(); @@ -247,17 +253,13 @@ void TestSplitPhaseLaunch() aclrtStream stream = reinterpret_cast(uintptr_t {0x8000}); CheckStatus("publish-only launch", TileXRMoonEp::TileXRMoonEpRunCombineV1( - &args, stream), TILEXR_MOONEP_SUCCESS); - Check(launchCalls == 1 && launchedParams.flags == - TILEXR_MOONEP_FLAG_COMBINE_PUBLISH_ONLY, - "publish-only flags were not forwarded"); + &args, stream), TILEXR_MOONEP_ERROR_INVALID_ARGUMENT); + Check(launchCalls == 0, "publish-only unexpectedly launched a kernel"); args.flags = TILEXR_MOONEP_FLAG_COMBINE_CONSUME_ONLY; CheckStatus("consume-only launch", TileXRMoonEp::TileXRMoonEpRunCombineV1( - &args, stream), TILEXR_MOONEP_SUCCESS); - Check(launchCalls == 2 && launchedParams.flags == - TILEXR_MOONEP_FLAG_COMBINE_CONSUME_ONLY, - "consume-only flags were not forwarded"); + &args, stream), TILEXR_MOONEP_ERROR_INVALID_ARGUMENT); + Check(launchCalls == 0, "consume-only unexpectedly launched a kernel"); } } // namespace @@ -300,6 +302,6 @@ int main() { TestPairedLaunch(); TestValidationAndFailures(); - TestSplitPhaseLaunch(); + TestSplitPhaseRejected(); return failures == 0 ? EXIT_SUCCESS : EXIT_FAILURE; } diff --git a/tests/moonep/unit/test_tilexr_moonep_host.cpp b/tests/moonep/unit/test_tilexr_moonep_host.cpp index e71bff0..717da50 100644 --- a/tests/moonep/unit/test_tilexr_moonep_host.cpp +++ b/tests/moonep/unit/test_tilexr_moonep_host.cpp @@ -22,6 +22,7 @@ int queryCalls = 0; int plannerCalls = 0; int dispatchCalls = 0; int dispatchUrmaCalls = 0; +int dispatchUrmaV2Calls = 0; int dispatchWorkspaceQueryCalls = 0; int combineCalls = 0; int prefetchCalls = 0; @@ -71,6 +72,7 @@ void Reset() TILEXR_MOONEP_SUCCESS; prefetchReturn = TILEXR_MOONEP_SUCCESS; queryCalls = plannerCalls = dispatchCalls = dispatchUrmaCalls = + dispatchUrmaV2Calls = dispatchWorkspaceQueryCalls = combineCalls = prefetchCalls = 0; queryWorkspaceBytes = 512; queryNvS = 12; @@ -205,6 +207,11 @@ void TestWorkspaceQuery() &dispatchBytes, &dispatchAlignment), TILEXR_MOONEP_SUCCESS); Check(dispatchWorkspaceQueryCalls == 1 && dispatchBytes == 4096 && dispatchAlignment == 2097152, "dispatch workspace query delegation mismatch"); + CheckStatus("dispatch V2 workspace query", TileXRMoonEpDispatchGetWorkspaceSizeV2( + comm, 2, 2, 64, TILEXR_MOONEP_DTYPE_BFLOAT16, + &dispatchBytes, &dispatchAlignment), TILEXR_MOONEP_SUCCESS); + Check(dispatchWorkspaceQueryCalls == 2, + "dispatch V2 workspace query delegation mismatch"); } void TestPlanningDelegation() @@ -278,6 +285,23 @@ void TestStageDelegation() Check(dispatchCalls == 1 && dispatchUrmaCalls == 1 && seenDispatchUrma == &dispatch, "URMA dispatch selection mismatch"); + TileXRMoonEpDispatchArgsV2 dispatchV2 {}; + dispatchV2.structSize = sizeof(dispatchV2); + dispatchV2.abiVersion = TILEXR_MOONEP_ABI_VERSION_V2; + CheckStatus("V2 dispatch requires workspace", + TileXRMoonEpDispatchV2(&dispatchV2, stream), + TILEXR_MOONEP_ERROR_INVALID_ARGUMENT); + Check(dispatchUrmaCalls == 1, + "invalid V2 dispatch must not reach the URMA implementation"); + dispatchV2.registeredWorkspace = + reinterpret_cast(uintptr_t {0x800000}); + dispatchV2.registeredWorkspaceBytes = 2097152; + CheckStatus("V2 dispatch", TileXRMoonEpDispatchV2(&dispatchV2, stream), + dispatchUrmaReturn); + Check(dispatchUrmaCalls == 1 && dispatchUrmaV2Calls == 1 && + seenDispatchUrma == &dispatchV2, + "V2 dispatch must select the URMA implementation"); + TileXRMoonEpCombineArgsV1 combine {}; combine.structSize = sizeof(combine); combine.abiVersion = TILEXR_MOONEP_ABI_VERSION_V1; @@ -387,6 +411,14 @@ int TileXRMoonEpRunDispatchUrmaV1( return dispatchUrmaReturn; } +int TileXRMoonEpRunDispatchUrmaV2( + const TileXRMoonEpDispatchArgsV2 *args, aclrtStream) +{ + ++dispatchUrmaV2Calls; + seenDispatchUrma = args; + return dispatchUrmaReturn; +} + int TileXRMoonEpRunCombineV1( const TileXRMoonEpCombineArgsV1 *args, aclrtStream stream) { diff --git a/tests/moonep/unit/test_tilexr_moonep_kernel_sources.cpp b/tests/moonep/unit/test_tilexr_moonep_kernel_sources.cpp index a66bfcf..e487456 100644 --- a/tests/moonep/unit/test_tilexr_moonep_kernel_sources.cpp +++ b/tests/moonep/unit/test_tilexr_moonep_kernel_sources.cpp @@ -171,37 +171,32 @@ int main() Contains("combine kernel", combineKernel, "TileXR::CommArgs"); Contains("combine kernel", combineKernel, "TileXR::IPC_DATA_OFFSET"); Contains("combine kernel", combineKernel, "SyncCollectives"); - Contains("combine kernel", combineKernel, "PublishLocalInput"); + Contains("combine kernel", combineKernel, + "sync_.Init(rank_, rankSize_, shareAddrs_, syncBuf_)"); + Contains("combine kernel", combineKernel, "dstLocal"); + Contains("combine kernel", combineKernel, "MoonEpCombineV2Peer"); + Contains("combine kernel", combineKernel, "PushPeerRows"); Contains("combine kernel", combineKernel, "CopyBytesGmToGm"); Contains("combine kernel", combineKernel, "hiddenChunkStride"); Contains("combine kernel", combineKernel, "routeWeightsNvs"); Contains("combine kernel", combineKernel, "PreReduceDuplicates"); - Contains("combine kernel", combineKernel, "ChunkStep"); + Contains("combine kernel", combineKernel, "BuildDuplicateMask"); Contains("combine kernel", combineKernel, "DataCacheCleanAndInvalid"); - Contains("combine kernel", combineKernel, "static_cast(encoded)"); - Contains("combine kernel", combineKernel, "-encoded64 - 1"); + Contains("combine kernel", combineKernel, "DecodeReverseRoute"); Contains("combine kernel", combineKernel, "ReduceHiddenChunk"); Contains("combine kernel", combineKernel, "AscendC::LocalTensor"); Contains("combine kernel", combineKernel, "AscendC::LocalTensor"); Contains("combine kernel", combineKernel, "AscendC::RoundMode::CAST_NONE"); Contains("combine kernel", combineKernel, "AscendC::Add"); Contains("combine kernel", combineKernel, "AscendC::RoundMode::CAST_RINT"); - Contains("combine kernel", combineKernel, "GatherWeights"); - Contains("combine kernel", combineKernel, "IsPublishOnly"); - Contains("combine kernel", combineKernel, "IsConsumeOnly"); - Contains("combine kernel", combineKernel, "RunPublishOnly"); - Contains("combine kernel", combineKernel, "RunConsumeOnly"); - Contains("combine kernel", combineKernel, "#include \"tilexr_udma.h\""); - Contains("combine kernel", combineKernel, "UDMAGetNbiOnQp"); - Contains("combine kernel", combineKernel, "UDMAQuietStatusOnQp"); - Contains("combine kernel", combineKernel, "UDMARegisteredRangeValid"); - Contains("combine kernel", combineKernel, - "const uint32_t qpCount = TileXR::UDMAQpCount(args_)"); - Contains("combine kernel", combineKernel, - "const uint32_t qpIdx = qpCount > 1U ? 1U : 0U"); - Excludes("combine kernel", combineKernel, "const uint32_t qpIdx = 0U"); + Contains("combine kernel", combineKernel, "CopyReceivedWeights"); + Excludes("combine kernel", combineKernel, "IsPublishOnly"); + Excludes("combine kernel", combineKernel, "IsConsumeOnly"); + Excludes("combine kernel", combineKernel, "RunPublishOnly"); + Excludes("combine kernel", combineKernel, "RunConsumeOnly"); + Excludes("combine kernel", Lower(combineKernel), "udma"); Contains("combine kernel", combineKernel, "kMoonEpCombineDataReadyStep"); - Contains("combine kernel", combineKernel, "kMoonEpCombineWindowDrainedStep"); + Contains("combine kernel", combineKernel, "kChunkDrainedStepBase"); Contains("combine kernel", combineKernel, "kMoonEpCombineFailedStep"); Contains("combine kernel", combineKernel, "kMoonEpCombineStatusInvalidRoute"); Contains("combine kernel", combineKernel, "kMoonEpCombineStatusTimeoutBase"); diff --git a/tests/moonep/unit/test_tilexr_moonep_prefetch_weight_host.cpp b/tests/moonep/unit/test_tilexr_moonep_prefetch_weight_host.cpp index 31fa7ae..9b84859 100644 --- a/tests/moonep/unit/test_tilexr_moonep_prefetch_weight_host.cpp +++ b/tests/moonep/unit/test_tilexr_moonep_prefetch_weight_host.cpp @@ -99,6 +99,20 @@ void TestLaunch() seenContext.layout.expertsPerRank == 4, "prefetch UDMA layout mismatch"); + Reset(); + qpNum = 32; + plan = Plan(); + gate = Weight(0x100000, 4, 8); + up = Weight(0x101000, 4, 16); + down = Weight(0x102000, 8, 8); + args = Args(&plan, &gate, &up, &down); + Status("prefetch shared-domain QPs", + TileXRMoonEp::TileXRMoonEpRunPrefetchWeightV1(&args, stream), + TILEXR_MOONEP_SUCCESS); + Check(launchCalls == 1 && seenContext.layout.qpNum == 32 && + seenContext.layout.blockDim == 4, + "prefetch must cap workers without rejecting the shared-domain QP count"); + gate.dtype = TILEXR_MOONEP_DTYPE_FLOAT32; Status("prefetch dtype", TileXRMoonEp::TileXRMoonEpRunPrefetchWeightV1(&args, stream), TILEXR_MOONEP_ERROR_INVALID_ARGUMENT); gate = Weight(0x100000, 4, 8); args.flags = 1; diff --git a/tests/moonep/unit/test_tilexr_moonep_sources.cpp b/tests/moonep/unit/test_tilexr_moonep_sources.cpp index dc85eae..fcb875c 100644 --- a/tests/moonep/unit/test_tilexr_moonep_sources.cpp +++ b/tests/moonep/unit/test_tilexr_moonep_sources.cpp @@ -106,6 +106,8 @@ int main() Contains("public header", header, "TileXRMoonEpPlanV1"); Contains("public header", header, "registeredWorkspace"); Contains("public header", header, "TileXRMoonEpDispatchGetWorkspaceSizeV1"); + Contains("public header", header, "TileXRMoonEpDispatchGetWorkspaceSizeV2"); + Contains("public header", header, "TileXRMoonEpDispatchV2"); Excludes("public header", header, "tilexr_api.h"); Excludes("public header", header, "std::"); @@ -116,6 +118,8 @@ int main() Contains("host", host, "dupCounts"); Contains("host", host, "TileXRMoonEpRunDispatchV1"); Contains("host", host, "TileXRMoonEpRunDispatchUrmaV1"); + Contains("host", host, "TileXRMoonEpRunDispatchUrmaV2"); + Contains("host", host, "TileXRMoonEpDispatchV2"); Contains("host", host, "TileXRMoonEpRunCombineV1"); Excludes("host dispatch stub", host, "RunLocalStub(args, stream, StubStage::Dispatch)"); @@ -156,6 +160,7 @@ int main() Contains("dispatch CMake", dispatchCmake, "INSTALL_RPATH \"$ORIGIN\""); Contains("URMA dispatch Host", dispatchUrmaHost, "TileXRMoonEpBuildDispatchUrmaLayout"); + Contains("URMA dispatch Host", dispatchUrmaHost, "aclrtMemsetAsync"); Contains("URMA dispatch Host", dispatchUrmaHost, "registeredWorkspaceBytes"); Contains("URMA dispatch Host", dispatchUrmaHost, "params.zeroFillRanges = static_cast(args->plan->zeroFillRanges)"); diff --git a/tests/moonep/unit/test_tilexr_moonep_stage_layout.cpp b/tests/moonep/unit/test_tilexr_moonep_stage_layout.cpp index 12c8b3f..d3f506f 100644 --- a/tests/moonep/unit/test_tilexr_moonep_stage_layout.cpp +++ b/tests/moonep/unit/test_tilexr_moonep_stage_layout.cpp @@ -131,12 +131,13 @@ void TestChunkingAndPlanAgreement() CheckStatus("chunked combine", TileXRMoonEp::TileXRMoonEpBuildCombineLayout( 0, 2, &plan, &output, nullptr, &input, nullptr, 0, &combine), TILEXR_MOONEP_SUCCESS); - CHECK_TRUE(combine.chunkCount == 2 && - combine.hiddenPayloadBytes <= static_cast(TileXR::IPC_BUFF_MAX_SIZE)); + CHECK_TRUE(combine.chunkCount >= 3 && combine.hiddenPayloadBytes > 0 && + combine.receiveHiddenOffset >= combine.hiddenPayloadBytes && + combine.windowBytes <= static_cast(TileXR::IPC_BUFF_MAX_SIZE)); CheckStatus("chunked split combine", TileXRMoonEp::TileXRMoonEpBuildCombineLayout( 0, 2, &plan, &output, nullptr, &input, nullptr, TILEXR_MOONEP_FLAG_COMBINE_PUBLISH_ONLY, &combine), - TILEXR_MOONEP_ERROR_NOT_SUPPORTED); + TILEXR_MOONEP_ERROR_INVALID_ARGUMENT); plan.r = 4; CheckStatus("world mismatch", TileXRMoonEp::TileXRMoonEpBuildCombineLayout( @@ -169,15 +170,21 @@ void TestCombinePairedLayout() TILEXR_MOONEP_SUCCESS); CHECK_TRUE(layout.s == 4 && layout.n == 8 && layout.nvS == 12); CHECK_TRUE(layout.hiddenRowBytes == 32 && layout.hiddenChunkStride == 32); - CHECK_TRUE(layout.hiddenPayloadBytes == 384 && layout.routeWeightsOffset == 384); - CHECK_TRUE(layout.routeWeightsBytes == 48 && layout.windowBytes == 432); + CHECK_TRUE(layout.blockDim == 2 && layout.stepCount == 1); + CHECK_TRUE(layout.hiddenPayloadBytes == 384 && layout.receiveHiddenOffset == 384); + CHECK_TRUE(layout.sourceWeightsOffset == 768 && layout.receiveWeightsOffset == 832); + CHECK_TRUE(layout.routeWeightsBytes == 48 && layout.duplicateMaskOffset == 896); + CHECK_TRUE(layout.doneOffset == 960 && layout.doneBytes == 128); + CHECK_TRUE(layout.coreStatusOffset == 1088 && layout.windowBytes == 1216); CheckStatus("publish-only combine", TileXRMoonEp::TileXRMoonEpBuildCombineLayout( 0, 2, &plan, &hiddenNvsh, &weightsNvs, &hiddenSh, &weightsSk, - TILEXR_MOONEP_FLAG_COMBINE_PUBLISH_ONLY, &layout), TILEXR_MOONEP_SUCCESS); + TILEXR_MOONEP_FLAG_COMBINE_PUBLISH_ONLY, &layout), + TILEXR_MOONEP_ERROR_INVALID_ARGUMENT); CheckStatus("consume-only combine", TileXRMoonEp::TileXRMoonEpBuildCombineLayout( 0, 2, &plan, &hiddenNvsh, &weightsNvs, &hiddenSh, &weightsSk, - TILEXR_MOONEP_FLAG_COMBINE_CONSUME_ONLY, &layout), TILEXR_MOONEP_SUCCESS); + TILEXR_MOONEP_FLAG_COMBINE_CONSUME_ONLY, &layout), + TILEXR_MOONEP_ERROR_INVALID_ARGUMENT); CheckStatus("ambiguous split combine", TileXRMoonEp::TileXRMoonEpBuildCombineLayout( 0, 2, &plan, &hiddenNvsh, &weightsNvs, &hiddenSh, &weightsSk, TILEXR_MOONEP_FLAG_COMBINE_PUBLISH_ONLY | diff --git a/tests/moonep_combine_v2/unit/test_combine_v2_schedule.cpp b/tests/moonep_combine_v2/unit/test_combine_v2_schedule.cpp index e291ac6..3bc7248 100644 --- a/tests/moonep_combine_v2/unit/test_combine_v2_schedule.cpp +++ b/tests/moonep_combine_v2/unit/test_combine_v2_schedule.cpp @@ -164,6 +164,10 @@ void TestTokensAndShapes() kMoonEpCombineV2SmallSlots), "small shape rejected"); Check(MoonEpCombineV2ShapeValid(256, 1024, 4, 2040), "PR113 shape rejected"); + Check(MoonEpCombineV2ReduceInputStrideElements(8U) == 16U, + "short BF16 reduction rows must use a 32-byte UB stride"); + Check(MoonEpCombineV2ReduceInputStrideElements(4096U) == 4096U, + "aligned BF16 reduction rows changed stride"); Check(!MoonEpCombineV2ShapeValid(256, 1024, 4, 1023), "undersized NvS accepted"); diff --git a/tools/moonep/benchmark.py b/tools/moonep/benchmark.py index ba30e17..861a59e 100644 --- a/tools/moonep/benchmark.py +++ b/tools/moonep/benchmark.py @@ -78,24 +78,34 @@ def _reference_process_group_backend( _NATIVE_STAGE_KERNEL_VERSIONS = { "planning": "tilexr_ep_plan_kernel (PlannerV3)", - "dispatch": "tilexr_moonep_dispatch_urma_kernel (DispatchV1)", + "dispatch": "tilexr_moonep_dispatch_urma_kernel (DispatchV2)", "prefetch_weight": "tilexr_moonep_prefetch_weight_kernel (V1)", - "combine": "tilexr_moonep_combine_kernel (V1)", + "combine": "tilexr_moonep_combine_v2_kernel (CombineV2)", "reduce_grad": "tilexr_moonep_reduce_grad_kernel (V2)", } -FINAL_SHARED_STATUS_SUCCESS = 3000 +_COMBINE_KERNEL_VERSIONS = { + 1: "tilexr_moonep_combine_kernel (CombineV1Memory)", + 2: "tilexr_moonep_combine_v2_kernel (CombineV2)", +} + +FINAL_SHARED_STATUS_SUCCESS = 0 REDUCE_GRAD_STATUS_SUCCESS = 0 def stage_execution_metadata( - capabilities: Mapping[str, object], *, torch_npu_version: str + capabilities: Mapping[str, object], *, torch_npu_version: str, + combine_version: int = 2, ) -> dict[str, dict[str, object]]: + if combine_version not in _COMBINE_KERNEL_VERSIONS: + raise ValueError(f"unsupported MoonEP Combine version {combine_version}") implementations = capabilities.get("implementations") if not isinstance(implementations, Mapping): raise ValueError("MoonEP capabilities do not contain stage implementations") result = {} - for stage, kernel_version in _NATIVE_STAGE_KERNEL_VERSIONS.items(): + versions = dict(_NATIVE_STAGE_KERNEL_VERSIONS) + versions["combine"] = _COMBINE_KERNEL_VERSIONS[combine_version] + for stage, kernel_version in versions.items(): implementation = str(implementations.get(stage, "unavailable")) native = implementation == "native" result[stage] = { @@ -851,6 +861,7 @@ def record_failure(step: str, exc: Exception) -> None: result["stage_execution"] = stage_execution_metadata( capabilities, torch_npu_version=str(environment["torch_npu"]), + combine_version=int(getattr(context.runtime, "combine_version", 2)), ) result["topology"] = topology_metadata(context) if os.environ.get("TILEXR_MOONEP_TRACE_STAGES", "0") == "1":