diff --git a/.gitattributes b/.gitattributes index dfdb8b77..b192d7cb 100644 --- a/.gitattributes +++ b/.gitattributes @@ -1 +1,4 @@ *.sh text eol=lf + +# Unified diff context lines intentionally contain the required leading space. +tools/moonep/mindspeed/*.patch whitespace=-blank-at-eol diff --git a/.gitignore b/.gitignore index ac3c32bd..2e797e7a 100644 --- a/.gitignore +++ b/.gitignore @@ -58,9 +58,9 @@ MANIFEST.in # Agent planning scratch files /.planning/ -/task_plan.md -/findings.md -/progress.md +task_plan.md +findings.md +progress.md # custom .worktrees/ diff --git a/AGENTS.md b/AGENTS.md index ced562b0..0ab41785 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -42,6 +42,11 @@ Optional CMake switches are `TILEXR_BUILD_COLLECTIVES`, `TILEXR_BUILD_EP`, `TILE - Keep these notes concise, actionable, and evidence-based. Capture the triggering context, failure mode or impact, root cause, correct approach, and validation boundary; update an existing module, architecture, validation, or troubleshooting document instead of creating scattered task notes when possible. - Add a lesson to `AGENTS.md` only when it is project-wide, high-impact, and easy to get wrong repeatedly. Keep task-specific details, transient environment observations, and long investigations in the relevant documentation instead. +## Debugging + +- Before debugging TileXR, read [docs/moonep/MINDSPEED_DEBUGGING_EXPERIENCE.md](docs/moonep/MINDSPEED_DEBUGGING_EXPERIENCE.md) and use its evidence-driven workflow, test ladder, and state, queue, and ownership checklists to guide the investigation. +- Treat the historical root causes in that document as hypotheses rather than conclusions. Reproduce the current failure, identify the first failing boundary, and verify the actual source and binary provenance before changing code. + ## Architecture - `src/comm` builds `libtile-comm.so`, owns communicator setup, peer mappings, capability flags, and `CommArgs`, and exposes the core API through `tilexr_api.h`. diff --git a/docs/moonep/MINDSPEED_DEBUGGING_EXPERIENCE.md b/docs/moonep/MINDSPEED_DEBUGGING_EXPERIENCE.md new file mode 100644 index 00000000..39327780 --- /dev/null +++ b/docs/moonep/MINDSPEED_DEBUGGING_EXPERIENCE.md @@ -0,0 +1,196 @@ +# TileXR MoonEP 接入 MindSpeed 调试经验 + +本文记录 2026-08-10 至 2026-08-12 在单机 8 卡、4K/8P/EP8 模型中调试 +TileXR MoonEP 接入 MindSpeed 的问题、证据链和可复用方法。后续调试应把本文作为 +排查入口,但不能把历史根因直接套到新故障上;必须先用当前源码、二进制和运行日志 +重新确认第一个失败边界。 + +## 最终结论 + +这次适配不是一个单点故障,而是模型规模和完整前反向流程依次暴露了五类契约问题: + +1. 通信 runtime 和 RA 的所有权不唯一。 +2. MindSpeed 依赖的 Buffer、Tensor view 和 zero-copy 契约超出了 MoonEP 公共 API。 +3. 逻辑 worker 数、物理 QP 数和 AIV block 数被错误地绑定到同一固定上限。 +4. CQE、SQ tail、ring index 和 cycle bit 的语义没有明确区分。 +5. 同一个 plan 跨 Planner、Dispatch、Prefetch、Combine、反向 Dispatch 和 ReduceGrad + 复用,但 status 的输入、输出和清理责任没有形成完整状态机。 + +最后阻塞 grouped-URMA Dispatch + Combine V2 完整模型的直接原因,是前向 Combine +成功后在复用 plan 中留下 `status=3000`。反向 URMA Dispatch 要求输入状态为 `0`, +并通过 `CAS(0, error)` 发布首错误。旧状态同时导致正常发送路径被跳过,并阻止真实 +错误码写入,最终表现为反向 Dispatch 等待 flag 超时。 + +Python 侧的 `plan.status.zero_()` 不是可靠修复:NPU task queue 和 Kernel launch 可能 +位于不同 stream,无法保证 reset 先于 consumer Kernel。最终方案是通过 +`TILEXR_MOONEP_FLAG_RESET_STATUS`,由 Host 在 Dispatch 的同一 stream 上排队 +`aclrtMemsetAsync`,然后紧邻 Kernel launch。 + +## 已确认问题 + +| 现象 | 根因 | 修复原则 | 发现方式 | +|---|---|---|---| +| `RaInit=328002`,UDMA 注册继发 `-7` | SHMEM/HCCL 和 TileXR 重复拥有 RA;rank 间 owned/attached 状态可能不一致 | 通信 runtime 只能有一个 owner;允许显式 attach,并跨 rank 校验所有权 | TileXR-only、SHMEM -> TileXR、TileXR -> SHMEM 三组初始化顺序 A/B | +| 替换 `moonep.Buffer` 后模型仍失败 | MindSpeed 依赖 `_ctx`、token/route buffers、packed Prefetch/Reduce、caller buffer 等私有契约 | 用适配层显式桥接,不能只替换公共类名 | 从 traceback 沿 caller/adapter/native 三层逐项核对对象契约 | +| 权重被 `storage_offset != 0` 拒绝 | 权重是连续 arena 的合法 subview,随后还会复制到独立注册 backing | 区分 source view 和 registered backing;只在立即复制边界允许 offset | 打印 storage、data pointer、offset 和注册对象身份 | +| backward 要求 SHMEM view | zero-copy alias 缺少 owner、generation 和 plan 元数据 | 地址别名必须携带所有权和生命周期标签 | 验证 alias storage 相同,并校验 owner/plan/generation | +| Combine V1 `DecodeRoute=3002` | 4 字节 `DataCopyPad` 标量读取返回陈旧 UB 数据 | GM 小标量使用当前 CANN 版本实机验证过的读取模式 | 逐分支状态码加设备侧原值捕获;立即在 Combine 后同步状态 | +| PrefetchWeight 在 32 QP 返回 `-3` | 把物理 `qpNum` 当作 worker 数,只接受 1/2/4/8 | `workerCount=min(qpNum, maxWorkers)`,Kernel 仍接收完整 `qpNum` | 查询真实 QP 数并对照 Host 参数校验分支 | +| ReduceGrad workspace query 返回 `-1` | Host/Kernel QP 上限固定为 8,实机为 32 | Host、Kernel 和测试统一支持目标硬件 QP 上限 | 1/2/4/8/32 QP 参数矩阵 | +| ReduceGrad 第二轮 CQ 失败,`entryIdx=0x4000` | `entryIdx` 携带 SQ cycle,原实现错误地要求它小于 ring depth | 先按 depth 归一化,再依据绝对 SQ tail 计算完成 BB 数 | 两轮同 QP 精确复现;dump raw CQE、SQ head/tail、outstanding | +| Combine 后复用 plan 的反向 Dispatch 超时 | 旧 `status=3000` 违反 URMA 输入协议,异步 Python reset 又存在跨 stream 竞态 | consumer Host 在同一 stream reset;先检查旧首错误,不能掩盖真实失败 | 生产规模 oracle 对比不清理、异步清理、同步清理和 Host same-stream reset | +| 4K grouped Dispatch 超时或 SQ 满 | 全量 route 无法一次装入 UB,WQE 也不能一次塞入 SQ | route tiling、WQE 分批发布、每批 CQ 回收,以 `head-tail` 计算 outstanding | case 15 生产规模单算子和 H=7168 grouped oracle | + +“重复 MR 注册泄漏”“完全没有 poll CQ”“peer 调度不对称”都曾是合理假设,但被后续 +A/B、原始队列状态和成功对照推翻,不应继续作为既定根因传播。 + +## 关键测试如何定位根因 + +### 1. 初始化顺序 A/B + +分别运行 TileXR-only、SHMEM -> TileXR 和 TileXR -> SHMEM。结果显示失败随初始化顺序 +变化,把第一个失败边界从模型逻辑缩小到 RA 所有权,而不是 UDMA 数据面。 + +### 2. 两轮 ReduceGrad 精确复现 + +第一轮全部通过,第二轮仅部分 rank 失败。raw CQE 显示 `entryIdx=16384`、ring-local +tail 为 0、outstanding 为 1。16384 恰好是一个 SQ cycle,证明错误来自 cycle 归一化, +而不是超时长度、WQE 地址或 MR 注册。加入 modulo-depth 修复后,两轮 16/16 rank-round +通过;扩大 WQE 数量后仍通过。 + +### 3. 生产规模 grouped-URMA oracle + +oracle 使用以下真实模型特征,而不是缩小到会绕开问题的 toy shape: + +- 单机 8 卡、8 rank; +- `S=4096`、`K=8`、`H=7168`、EP8、`NvS=32768`; +- grouped-URMA,group width 16; +- model-skew 路由和 route weights; +- PrefetchWeight、Combine V2; +- 连续创建 5 个额外 plan 后复用最后一个 plan; +- 所有 rank 对 Dispatch 和 Combine 输出做 exact comparison。 + +第一次 Dispatch 和 Combine 正确,而 Combine 后复用 plan 的 Dispatch 失败。A/B 显示: + +- 保留 `status=3000`:失败; +- Python `zero_()` 不同步:不稳定; +- Python `zero_()` 后显式同步:通过; +- Host 在 Dispatch stream 上 reset:稳定通过。 + +该测试把根因从宽泛的“反向 flag 超时”缩小为 plan status 的跨阶段和跨 stream 生命周期 +错误。它同时验证了 4K route tiling、WQE 分批和 CQ 回收,因此是完整模型前最重要的 +root-cause oracle。 + +### 4. 完整模型与性能复测 + +修复后按以下顺序扩大验证范围: + +1. case 15:4K/EP8 grouped-URMA Dispatch 单算子复测。 +2. grouped oracle:H=7168、model-skew、Prefetch、Combine V2、5 个额外 plan 和 plan reuse, + 8 rank exact comparison。 +3. 完整正确性模型 `tilexr_urma_correctness_4k_8p_ep8_0812_015220`:8/8 迭代通过。 +4. 关闭 DFX、trace、dump 和 profiler 后运行 + `tilexr_urma_perf_v2_4k_8p_ep8_0812_095543`:8/8 迭代、退出码 0、无 skip/NaN。 + +完整模型通过只能证明该组合路径闭环,不能替代单算子对边界语义的证明;单算子通过也 +不能证明 plan reuse、前反向和多轮资源复用。 + +## 推荐调试流程 + +### 1. 固定运行身份 + +每次运行先保存以下信息: + +- commit、工作树 diff 和源码哈希; +- 实际加载的 `.so`、AICore binary 路径和哈希; +- CANN、驱动、固件、SoC 和 conda 环境; +- launcher、环境变量、rank/device 映射; +- 运行前后的 NPU PID、stdout、stderr、plog 和输出目录。 + +构建成功不等于运行时加载了新库。没有二进制 provenance 的通过或失败结果都不能作为 +最终证据。 + +### 2. 找第一个失败边界 + +按 Planner -> Dispatch -> Prefetch -> Combine -> reused Dispatch -> ReduceGrad 顺序,在 +每个异步 stage 后只增加一个必要同步点。记录该点的 Host 返回码、`plan.status`、DFX +首错误和相关队列状态。后续超时往往只是更早的异步错误在同步点被暴露。 + +### 3. 一次验证一个可证伪假设 + +写清楚: + +> X 导致症状 Y;如果成立,只改变 Z 后结果应从 A 变成 B。 + +优先改变初始化顺序、transport、是否 Prefetch、是否 Combine、是否复用 plan、是否同步 +reset 等单一变量。一次同时修改 Kernel、Host、timeout 和路由,会使任何成功都无法归因。 + +### 4. 从协议状态入手,而不是先加 timeout + +超时时至少采集: + +- plan status 的阶段来源和预期值; +- flag 矩阵在 producer 前后、consumer 前后和超时后的差异; +- SQ/CQ head、tail、depth、owner、`entryIdx`、outstanding; +- 第一个失败 peer、QP、phase 和原始错误码。 + +延长 timeout 只有在状态持续前进但速度不足时才有意义。状态完全不变、已进入错误分支或 +队列字段非法时,延长 timeout 只会降低调试效率。 + +### 5. 修复协议拥有者 + +- 跨 stream reset:由 consumer Host 在目标 stream 上排队。 +- 跨 rank runtime:明确 owned/attached,并做一致性校验。 +- 错误发布:使用 first-error/CAS,禁止后续 core 覆盖首错误。 +- ring 字段:在类型、命名和测试中区分 absolute counter、ring index 和 cycle bit。 +- Tensor:区分 source view、copy destination 和 registered backing。 + +### 6. 逐级扩大验证 + +推荐门禁顺序: + +1. Host/unit/mock:参数、状态转换、边界算术。 +2. 单 rank 或最小多 rank:验证 ABI 和第一条真实数据路径。 +3. 两轮以上真实 UDMA:验证 QP、MR、SQ/CQ 和 magic 复用。 +4. case 15:生产 route 数和 grouped-URMA 流控。 +5. grouped oracle:model-skew、Prefetch、Combine、额外 plan、plan reuse、exact comparison。 +6. 完整 4K/8P/EP8 前反向模型。 +7. 关闭所有 DFX 后的独占性能运行。 + +## 必须长期保留的回归维度 + +- forward 和 backward 都执行; +- 同一个 plan 在 Combine 后被 Dispatch 复用; +- 至少两轮 MR、QP、SQ/CQ 和 magic 复用; +- route 数覆盖 UB tile 和 SQ depth 边界; +- QP 数覆盖 1/2/4/8/32; +- SQ 位置覆盖 `depth-1`、`depth`、`depth+1`、`2*depth`; +- 同时覆盖 1-BB 和多 BB WQE; +- balanced、model-skew、sparse 和 unique routing; +- `S=4096`、`K=8`、`H=7168`、EP8 的生产规模; +- 正确性运行开启失败时 DFX,性能运行关闭 trace、dump、DFX 和 profiler。 + +## 避免重复踩坑 + +- 不要把 error code 的字面含义直接当根因;查调用边界和相邻日志。 +- 不要在第一次成功后停止,队列 cycle 和资源复用问题通常第二轮才出现。 +- 不要用 Python 异步 tensor 操作实现 Host/Kernel 协议同步。 +- 不要把 worker、QP、AIV block 和 rank 数视为同一种并行度。 +- 不要用 toy shape 证明生产规模的 UB/SQ 容量安全。 +- 不要依赖已消费 SQE 的内容推导 CQ completion。 +- 不要在调试运行和性能运行之间复用未审计的环境变量。 +- 不要因某次补丁通过完整模型就跳过最小 reproducer;最小 reproducer 才能证明因果。 +- 不要删除被推翻的假设记录。保留否定证据可以防止后续重复猜测。 + +## 当前实现状态说明 + +本文记录的是已验证经验,不代表所有修复都已经进入 `main`。截至 2026-08-12: + +- 已进入 `main`:MoonEP 上游 API 兼容、Combine V2 路由与 reduction、阶段性能报告。 +- 当前调试分支中:same-stream status reset、32-QP Prefetch/ReduceGrad、grouped Dispatch + route tiling/WQE/CQ 流控和相关测试。 +- 任务补丁或 MindSpeed 工作树中:MindSpeed adapter、external communication owner、RA + attach/ownership guard、历史 Combine V1 标量修复和通用 CQE cycle 修复。 + +开始新的调试任务时,应先检查这些修改是否已经进入当前目标分支和实际加载的二进制, +不能根据本文的历史状态假设代码已经包含修复。 diff --git a/integrations/moonep_torch/tilexr_moonep/abi.py b/integrations/moonep_torch/tilexr_moonep/abi.py index 0c2ac639..8d958f7f 100644 --- a/integrations/moonep_torch/tilexr_moonep/abi.py +++ b/integrations/moonep_torch/tilexr_moonep/abi.py @@ -15,6 +15,7 @@ TILEXR_MOONEP_FLAG_ZERO_COPY = 1 << 2 TILEXR_MOONEP_FLAG_COMBINE_PUBLISH_ONLY = 1 << 3 TILEXR_MOONEP_FLAG_COMBINE_CONSUME_ONLY = 1 << 4 +TILEXR_MOONEP_FLAG_RESET_STATUS = 1 << 5 TILEXR_MOONEP_REDUCE_GRAD_UDMA_THRESHOLD_BYTES = 1 << 20 diff --git a/integrations/moonep_torch/tilexr_moonep/runtime.py b/integrations/moonep_torch/tilexr_moonep/runtime.py index 6b1ac304..256d1b84 100644 --- a/integrations/moonep_torch/tilexr_moonep/runtime.py +++ b/integrations/moonep_torch/tilexr_moonep/runtime.py @@ -10,6 +10,7 @@ from .abi import ( TILEXR_MOONEP_ABI_VERSION, TILEXR_MOONEP_FLAG_NONE, + TILEXR_MOONEP_FLAG_RESET_STATUS, TILEXR_SUCCESS, TileXRMoonEPDType, TileXRMoonEPCombineArgsV1, @@ -686,7 +687,7 @@ def dispatch( args.routeWeightsNvs = ( ctypes.pointer(weights_nvs) if weights_nvs is not None else None ) - args.flags = TILEXR_MOONEP_FLAG_NONE + args.flags = TILEXR_MOONEP_FLAG_RESET_STATUS args.registeredWorkspace = void_p(registered_workspace) args.registeredWorkspaceBytes = int(registered_workspace_bytes) ret = self._moonep_lib.TileXRMoonEpDispatchV2( @@ -709,7 +710,29 @@ def prefetch_weight(self, context, plan, projections, stream_ptr: int) -> None: ret = self._moonep_lib.TileXRMoonEpPrefetchWeightV1( ctypes.byref(args), void_p(stream_ptr) ) - self._check("TileXRMoonEpPrefetchWeightV1", ret) + if int(ret) != TILEXR_SUCCESS: + backing = projections.backing + backing_ptr = 0 if backing is None else int(backing.data_ptr()) + projection_detail = [] + for name, tensor in ( + ("gate", projections.gate), + ("up", projections.up), + ("down", projections.down), + ): + projection_detail.append( + f"{name}=shape{tuple(tensor.shape)},bytes={tensor_nbytes(tensor)}," + f"offset={int(tensor.data_ptr()) - backing_ptr}" + ) + detail = ( + f"plan=(r={context.planner_group_size},e={context.expert_count}," + f"b={context.prefetch_slots},nvS={context.nv_s},k={context.topk}); " + + "; ".join(projection_detail) + + f"; backing_bytes={0 if backing is None else tensor_nbytes(backing)}" + + f"; active_udma_owner={self._active_udma_owner}" + + f"; active_udma_bytes={self._active_udma_bytes}" + + f"; udma_qp_count={self._udma_qp_count}" + ) + self._check("TileXRMoonEpPrefetchWeightV1", ret, detail) def udma_register(self, tensor) -> int: size = tensor_nbytes(tensor) diff --git a/integrations/moonep_torch/tilexr_moonep/torch_api.py b/integrations/moonep_torch/tilexr_moonep/torch_api.py index f422330e..a57303a3 100644 --- a/integrations/moonep_torch/tilexr_moonep/torch_api.py +++ b/integrations/moonep_torch/tilexr_moonep/torch_api.py @@ -1,6 +1,7 @@ from __future__ import annotations import os +import struct from dataclasses import dataclass, field from typing import Any, Callable @@ -9,6 +10,35 @@ _PREFETCH_WEIGHT_STATUS_SUCCESS = 4000 _UDMA_REGISTRATION_ALIGNMENT = 2 * 1024 * 1024 +_DISPATCH_COMPLETION_FLAG_RANKS = 512 +_DISPATCH_COMPLETION_QP_COUNT = 2 +_DISPATCH_COMPLETION_FLAGS_BYTES = ( + _DISPATCH_COMPLETION_FLAG_RANKS * _DISPATCH_COMPLETION_QP_COUNT * 8 +) +_DISPATCH_SIGNAL_BYTES = 64 * 64 +_DISPATCH_PROFILE_BYTES = 64 * 256 +_DISPATCH_DFX_BYTES = 64 * 128 +_DISPATCH_KERNEL_STATUS_BYTES = 64 +_DISPATCH_COMMON_TAIL_BYTES = ( + _DISPATCH_COMPLETION_FLAGS_BYTES + + _DISPATCH_SIGNAL_BYTES + + 2 * _DISPATCH_PROFILE_BYTES + + 2 * _DISPATCH_DFX_BYTES + + _DISPATCH_KERNEL_STATUS_BYTES +) +def _format_dispatch_completion_flags(data: bytes, rank_size: int) -> str: + rows = tuple(struct.iter_unpack(" 4 or @@ -479,6 +510,8 @@ def __init__( self._torch = torch_module or _torch() self._closed = False self._epoch = 0 + self._dispatch_call_count = 0 + self._dispatch_flag_before_snapshots: dict[int, tuple[Any, ...]] = {} self._bound_stream_ptr: int | None = None self._pending_refs: list[tuple[Any, ...]] = [] self._pending_plans: list[MoonEPPlan] = [] @@ -546,6 +579,158 @@ def _retain(self, plan: MoonEPPlan, *values: Any) -> None: def _expect_status(self, plan: MoonEPPlan, status: int) -> None: self._pending_statuses[id(plan)] = int(status) + def _check_plan_status(self, plan: MoonEPPlan) -> None: + expected = self._pending_statuses.get(id(plan)) + if expected is None: + return + actual = int(plan.status.item()) + if actual == expected: + self._dispatch_flag_before_snapshots.pop(id(plan), None) + return + if os.environ.get("TILEXR_MOONEP_FLAG_DUMP_MODE") == "failure": + self._dump_failed_dispatch_completion_flags(plan, actual) + if actual in (2005, 2006, 2007): + self.context.mark_poisoned() + dfx = self._dispatch_dfx_summary() + raise RuntimeError( + "MoonEP device status failures: " + f"epoch {plan.epoch}: actual {actual}, expected {expected}" + f"{'; dispatch_dfx=' + dfx if dfx else ''}" + ) + + def _dispatch_completion_flags_view(self): + raw = self.context._dispatch_workspace_owner + if raw is None: + raise RuntimeError("Dispatch workspace is unavailable") + aligned_offset = self.context._dispatch_workspace_ptr - int(raw.data_ptr()) + flags_start = ( + aligned_offset + + self.context._dispatch_workspace_bytes + - _DISPATCH_COMMON_TAIL_BYTES + ) + return raw.narrow(0, flags_start, _DISPATCH_COMPLETION_FLAGS_BYTES) + + def _dispatch_completion_flags_bytes(self) -> bytes: + return bytes( + self._dispatch_completion_flags_view().detach() + .cpu() + .tolist() + ) + + def _write_dispatch_completion_flags( + self, + data: bytes, + direction: str, + stage: str, + plan: MoonEPPlan, + call_count: int, + status: int, + ) -> None: + dump_dir = os.environ.get("TILEXR_MOONEP_FLAG_DUMP_DIR") + if not dump_dir: + return + os.makedirs(dump_dir, exist_ok=True) + stem = ( + f"rank{self.context.planner_group_rank}_pid{os.getpid()}_" + f"call{call_count:04d}_{direction}_" + f"epoch{plan.epoch}_{stage}_status{status}" + ) + with open(os.path.join(dump_dir, f"{stem}.bin"), "wb") as output: + output.write(data) + with open( + os.path.join(dump_dir, f"{stem}.txt"), "w", encoding="ascii" + ) as output: + output.write( + _format_dispatch_completion_flags( + data, self.context.planner_group_size + ) + ) + output.write("\n") + + def _dump_dispatch_completion_flags( + self, direction: str, stage: str, plan: MoonEPPlan + ) -> None: + self._write_dispatch_completion_flags( + self._dispatch_completion_flags_bytes(), + direction, + stage, + plan, + self._dispatch_call_count, + int(plan.status.item()), + ) + + def _capture_dispatch_completion_flags_before( + self, direction: str, plan: MoonEPPlan + ) -> None: + self._dispatch_flag_before_snapshots[id(plan)] = ( + self._dispatch_call_count, + direction, + self._dispatch_completion_flags_view().clone(), + ) + + def _dump_failed_dispatch_completion_flags( + self, plan: MoonEPPlan, actual_status: int + ) -> None: + snapshot = self._dispatch_flag_before_snapshots.get(id(plan)) + if snapshot is None: + return + call_count, direction, before = snapshot + before_data = bytes(before.detach().cpu().tolist()) + self._write_dispatch_completion_flags( + before_data, direction, "before", plan, call_count, 0 + ) + self._write_dispatch_completion_flags( + self._dispatch_completion_flags_bytes(), + direction, + "after", + plan, + call_count, + actual_status, + ) + + def _dispatch_dfx_summary(self) -> str: + if os.environ.get("TILEXR_MOONEP_DUMP_DFX_ON_ERROR", "0") != "1": + return "" + raw = self.context._dispatch_workspace_owner + if raw is None: + return "workspace-unavailable" + aligned_offset = self.context._dispatch_workspace_ptr - int(raw.data_ptr()) + workspace_bytes = self.context._dispatch_workspace_bytes + record = struct.Struct(" None: raise RuntimeError("check_pending_status requires a successful quiesce") try: statuses = [] + failed_dispatch_plans = [] for plan in self._pending_plans: expected = self._pending_statuses.get(id(plan)) if expected is None: continue - statuses.append( - ( - plan.epoch, - int(plan.status.item()), - expected, - ) - ) + actual = int(plan.status.item()) + statuses.append((plan.epoch, actual, expected)) + if actual == expected: + self._dispatch_flag_before_snapshots.pop(id(plan), None) + else: + failed_dispatch_plans.append((plan, actual, expected)) for plan in self._pending_reduce_plans: statuses.append((plan.epoch, int(plan.reduce_grad_status.item()), 0)) self._pending_refs.clear() @@ -1228,12 +1439,18 @@ def check_pending_status(self) -> None: if status != expected ] if failed: + if os.environ.get("TILEXR_MOONEP_FLAG_DUMP_MODE") == "failure": + for plan, actual, _ in failed_dispatch_plans: + self._dump_failed_dispatch_completion_flags(plan, actual) + dispatch_dfx = self._dispatch_dfx_summary() if any(status in (2005, 2006, 2007) for _, status, _ in failed): self.context.mark_poisoned() details = ", ".join( f"epoch {epoch}: actual {status}, expected {expected}" for epoch, status, expected in failed ) + if dispatch_dfx: + details += f"; dispatch_dfx={dispatch_dfx}" raise RuntimeError(f"MoonEP device status failures: {details}") finally: if self._reduce_grad_inflight: diff --git a/scripts/run_moonep.sh b/scripts/run_moonep.sh index 554fb3ea..8119843c 100644 --- a/scripts/run_moonep.sh +++ b/scripts/run_moonep.sh @@ -49,6 +49,8 @@ Available case IDs: 12 planning-64rank-single-route 64-rank, 8 nodes x 8 NPUs single route (rank_size=64, rank_per_dev=1, S=8, K=1, E=64, H=8, Hf=4, B=1, P=1) 13 planning-128rank-single-route 128-rank, 16 nodes x 8 NPUs single route (rank_size=128, rank_per_dev=1, S=8, K=1, E=128, H=8, Hf=4, B=1, P=1) 14 planning-16rank-16card-single-route 16-rank, 2 nodes x 8 NPUs single route (rank_size=16, rank_per_dev=1, S=8, K=1, E=16, H=8, Hf=4, B=1, P=1) + 15 dispatch-8rank-4k-ep8-grouped-urma 8-rank Dispatch-only grouped-URMA repro (rank_size=8, rank_per_dev=1, S=4096, K=8, E=32, H=7168, Hf=2048, B=4, P=1) + 16 flow-8rank-4k-ep8-grouped-urma-plan-reuse 8-rank full flow with Combine V2 then saved-plan backward Dispatch (rank_size=8, rank_per_dev=1, S=4096, K=8, E=32, H=7168, Hf=2048, B=4, P=1) Environment: ASCEND_RT_VISIBLE_DEVICES Legacy fallback when --visible-devices is omitted @@ -386,6 +388,32 @@ PY fi case_id="${resolved_case_id}" fi + +benchmark_kind="flow" +dispatch_modes=() +dispatch_repro_case_id="dispatch-8rank-4k-ep8-grouped-urma" +plan_reuse_repro_case_id="flow-8rank-4k-ep8-grouped-urma-plan-reuse" +if [[ "${case_id}" == "${dispatch_repro_case_id}" || + "${case_id}" == "${plan_reuse_repro_case_id}" ]]; then + if [[ "${mode}" != "benchmark" ]]; then + echo "${case_id} requires --mode benchmark" >&2 + exit 2 + fi + if [[ "${rank_size}" != "8" || "${node_count}" != "1" ]]; then + echo "${case_id} requires --rank-size 8 on one node" >&2 + exit 2 + fi + unset TILEXR_MOONEP_DISPATCH_TRANSPORT + export TILEXR_MOONEP_DISPATCH_PEER_MODE="group" + export TILEXR_MOONEP_DISPATCH_GROUP_WIDTH="16" +fi +if [[ "${case_id}" == "${dispatch_repro_case_id}" ]]; then + benchmark_kind="dispatch_hot_loop" + dispatch_modes=("hidden") +elif [[ "${case_id}" == "${plan_reuse_repro_case_id}" ]]; then + warmup="${warmup:-0}" + iterations="${iterations:-8}" +fi summary_file="${output_dir}/${case_id}/summary.json" if [[ "${aggregate_only}" == "true" ]]; then @@ -564,6 +592,7 @@ if (( node_count > 1 )); then fi launcher_args=( --mode "${mode}" + --benchmark-kind "${benchmark_kind}" --cases "${case_file}" --case-ids "${case_id}" --world-size "${rank_size}" @@ -573,6 +602,9 @@ launcher_args=( --output-dir "${output_dir}" --timeout-sec "${timeout_sec}" ) +if [[ "${benchmark_kind}" == "dispatch_hot_loop" ]]; then + launcher_args+=("--dispatch-modes" "${dispatch_modes[@]}") +fi if [[ "${dump_stage_tensors}" == "true" ]]; then launcher_args+=("--dump-stage-tensors") launcher_args+=("--tensor-preview-elements" "${tensor_preview_elements}") diff --git a/src/comm/udma/tilexr_udma_transport.cpp b/src/comm/udma/tilexr_udma_transport.cpp index 4633aa6c..f786839e 100644 --- a/src/comm/udma/tilexr_udma_transport.cpp +++ b/src/comm/udma/tilexr_udma_transport.cpp @@ -9,6 +9,7 @@ #include #include #include +#include #include #include #include @@ -21,6 +22,14 @@ namespace TileXR { namespace { +constexpr int TILEXR_HCCP_RA_ALREADY_INITIALIZED = 328002; + +bool AttachExistingRaEnabled() +{ + const char* value = std::getenv("TILEXR_UDMA_ATTACH_EXISTING_RA"); + return value != nullptr && std::strcmp(value, "1") == 0; +} + uint32_t Log2Uint64(uint64_t value) { uint32_t result = 0; @@ -188,6 +197,11 @@ int TileXRUDMATransport::Init(const TileXRUDMATransportOptions& options) Shutdown(); return ret; } + ret = AgreeRaOwnership(); + if (ret != TILEXR_SUCCESS) { + Shutdown(); + return ret; + } localStatus = BuildRoutes(); ret = AgreeInitStatus(localStatus); if (ret != TILEXR_SUCCESS) { @@ -250,6 +264,23 @@ int TileXRUDMATransport::AgreeInitStatus(int localStatus) const return TILEXR_SUCCESS; } +int TileXRUDMATransport::AgreeRaOwnership() const +{ + const int32_t local = raAttached_ ? 1 : 0; + std::array allAttached {}; + const int exchangeRet = options_.exchange->AllGather(&local, 1, allAttached.data()); + if (exchangeRet != TILEXR_SUCCESS) { + return exchangeRet; + } + for (int rank = 0; rank < options_.rankSize; ++rank) { + if (allAttached[rank] != local) { + TILEXR_LOG(ERROR) << "TileXR UDMA RA ownership mismatch at rank " << rank; + return TILEXR_ERROR_PARA_CHECK_FAIL; + } + } + return TILEXR_SUCCESS; +} + int TileXRUDMATransport::AgreeEidCount() const { if (options_.exchange == nullptr || options_.rankSize <= 0 || @@ -302,10 +333,16 @@ int TileXRUDMATransport::OpenDevice() initConfig.enableHdcAsync = 1; int ret = loader_.RaInit(&initConfig); if (ret != 0) { + if (ret == TILEXR_HCCP_RA_ALREADY_INITIALIZED && AttachExistingRaEnabled()) { + raAttached_ = true; + TILEXR_LOG(INFO) << "TileXR UDMA attaching to existing RA initialization"; + return TILEXR_SUCCESS; + } TILEXR_LOG(WARN) << "TileXR UDMA RaInit failed: " << ret; return TILEXR_ERROR_INTERNAL; } raInitialized_ = true; + raAttached_ = false; return TILEXR_SUCCESS; } @@ -2055,6 +2092,7 @@ void TileXRUDMATransport::CleanupContexts() tsdOpened_ = false; subPid_ = 0; } + raAttached_ = false; } void TileXRUDMATransport::Shutdown() diff --git a/src/comm/udma/tilexr_udma_transport.h b/src/comm/udma/tilexr_udma_transport.h index efecd812..af2afbfa 100644 --- a/src/comm/udma/tilexr_udma_transport.h +++ b/src/comm/udma/tilexr_udma_transport.h @@ -69,6 +69,7 @@ class TileXRUDMATransport { struct RegistrationState; int AgreeInitStatus(int localStatus) const; + int AgreeRaOwnership() const; int AgreeEidCount() const; int OpenDevice(); int BuildRoutes(); @@ -113,6 +114,7 @@ class TileXRUDMATransport { bool available_ = false; bool tsdOpened_ = false; bool raInitialized_ = false; + bool raAttached_ = false; pid_t subPid_ = 0; uint32_t logicDevId_ = 0; uint32_t deviceIdOffset_ = 0; diff --git a/src/include/tilexr_moonep.h b/src/include/tilexr_moonep.h index f89630b8..57f6136c 100644 --- a/src/include/tilexr_moonep.h +++ b/src/include/tilexr_moonep.h @@ -22,6 +22,7 @@ typedef void *TileXRCommPtr; #define TILEXR_MOONEP_FLAG_ZERO_COPY (UINT64_C(1) << 2) #define TILEXR_MOONEP_FLAG_COMBINE_PUBLISH_ONLY (UINT64_C(1) << 3) #define TILEXR_MOONEP_FLAG_COMBINE_CONSUME_ONLY (UINT64_C(1) << 4) +#define TILEXR_MOONEP_FLAG_RESET_STATUS (UINT64_C(1) << 5) typedef enum TileXRMoonEpStatus { TILEXR_MOONEP_SUCCESS = 0, diff --git a/src/include/tilexr_udma.h b/src/include/tilexr_udma.h index ae470f02..c7fd933e 100644 --- a/src/include/tilexr_udma.h +++ b/src/include/tilexr_udma.h @@ -244,22 +244,16 @@ __aicore__ inline uint32_t UDMAPollCQ(__gm__ UDMAInfo* udmaInfo, uint32_t pe, ui const uint32_t sqHead = ld_dev(reinterpret_cast<__gm__ uint32_t*>(wqCtxEntry->headAddr), 0); const uint32_t sqOutstanding = sqHead - sqTail; - if (sqOutstanding == 0U || sqOutstanding > wqCtxEntry->depth || - cqeAddr->entryIdx >= wqCtxEntry->depth) { + if (sqOutstanding == 0U || sqOutstanding > wqCtxEntry->depth) { return TILEXR_UDMA_STATUS_INVALID; } - __gm__ UDMASqeCtx* completedSqe = reinterpret_cast<__gm__ UDMASqeCtx*>( - wqCtxEntry->bufAddr + - (1U << wqCtxEntry->baseBkShift) * (sqTail % wqCtxEntry->depth)); - const UDMAOpcode opcode = static_cast(completedSqe->opcode); - if (opcode != UDMAOpcode::WRITE && opcode != UDMAOpcode::READ && - opcode != UDMAOpcode::WRITE_WITH_NOTIFY) { - return TILEXR_UDMA_STATUS_INVALID; - } - const uint32_t completedBb = UDMAWqeBBCnt(opcode); - const uint32_t expectedEntryIdx = - (sqTail + completedBb - 1U) % wqCtxEntry->depth; - if (completedBb > sqOutstanding || cqeAddr->entryIdx != expectedEntryIdx) { + const uint32_t tailIndex = sqTail % wqCtxEntry->depth; + const uint32_t completedEntryIndex = cqeAddr->entryIdx % wqCtxEntry->depth; + const uint32_t completedBb = + (completedEntryIndex + wqCtxEntry->depth - tailIndex) % + wqCtxEntry->depth + 1U; + if (completedBb > TILEXR_UDMA_MAX_SQE_BB_NUM || + completedBb > sqOutstanding) { return TILEXR_UDMA_STATUS_INVALID; } sqTail += completedBb; @@ -763,8 +757,9 @@ __aicore__ inline uint32_t UDMAFlushQpDoorbell( return TILEXR_UDMA_STATUS_SUCCESS; } -__aicore__ inline uint32_t UDMAQuietStatusOnQp( - const __gm__ CommArgs* args, int targetRank, uint32_t qpIdx) +__aicore__ inline uint32_t UDMAQuietStatusOnQpUntil( + const __gm__ CommArgs* args, int targetRank, uint32_t qpIdx, + uint32_t completionTarget) { if (!TILEXR_UDMA_ARCH_SUPPORTED) { return TILEXR_UDMA_STATUS_INVALID; @@ -776,13 +771,27 @@ __aicore__ inline uint32_t UDMAQuietStatusOnQp( if (udmaInfo->sqPtr == 0U || udmaInfo->scqPtr == 0U) { return TILEXR_UDMA_STATUS_INVALID; } + return UDMAPollCQ( + udmaInfo, static_cast(targetRank), qpIdx, completionTarget); +} + +__aicore__ inline uint32_t UDMAQuietStatusOnQp( + const __gm__ CommArgs* args, int targetRank, uint32_t qpIdx) +{ + if (!TILEXR_UDMA_ARCH_SUPPORTED || + !UDMAQueueOperationValid(args, targetRank, qpIdx)) { + return TILEXR_UDMA_STATUS_INVALID; + } + __gm__ UDMAInfo* udmaInfo = GetUDMAInfo(args); __gm__ UDMAWQCtx* qpCtxEntry = UDMAGetWQCtx(udmaInfo, static_cast(targetRank), qpIdx); - if (qpCtxEntry->wqeCntAddr == 0U) { + if (qpCtxEntry == nullptr || qpCtxEntry->wqeCntAddr == 0U) { return TILEXR_UDMA_STATUS_INVALID; } - uint32_t wqeCnt = ld_dev(reinterpret_cast<__gm__ uint32_t*>(qpCtxEntry->wqeCntAddr), 0); - return UDMAPollCQ(udmaInfo, static_cast(targetRank), qpIdx, wqeCnt); + const uint32_t completionTarget = ld_dev( + reinterpret_cast<__gm__ uint32_t*>(qpCtxEntry->wqeCntAddr), 0); + return UDMAQuietStatusOnQpUntil( + args, targetRank, qpIdx, completionTarget); } __aicore__ inline uint32_t UDMAQuietStatus(const __gm__ CommArgs* args, int targetRank) diff --git a/src/moonep/dispatch/urma/common/dispatch_wqe_batch.h b/src/moonep/dispatch/urma/common/dispatch_wqe_batch.h index 37ce952a..671fdce3 100644 --- a/src/moonep/dispatch/urma/common/dispatch_wqe_batch.h +++ b/src/moonep/dispatch/urma/common/dispatch_wqe_batch.h @@ -157,29 +157,30 @@ TILEXR_MOONEP_WQE_BATCH_INLINE uint32_t DispatchQpSelectedIndex( afterPrefix % 3U; } -TILEXR_MOONEP_WQE_BATCH_INLINE bool DispatchPeerWqesFitSq( +TILEXR_MOONEP_WQE_BATCH_INLINE bool DispatchPeerWqesStreamable( uint64_t routeCount, uint32_t sqEntryCount, uint32_t reserve = kDispatchSqPollReserve) { if (routeCount > UINT32_MAX || sqEntryCount <= reserve) { return false; } - const uint32_t count = static_cast(routeCount); - const uint32_t available = sqEntryCount - reserve; - for (uint32_t qpIdx = 0U; qpIdx < kDispatchQpCount; ++qpIdx) { - const uint32_t qpPayloadWqes = DispatchQpRouteCount( - count, 0U, qpIdx); - if (qpPayloadWqes >= available || qpPayloadWqes + 1U > available) { - return false; - } - } - return true; + return sqEntryCount - reserve >= kDispatchWqeBatchCapacity; } -TILEXR_MOONEP_WQE_BATCH_INLINE bool DispatchBatchNeedsCompletion( - bool finalBatchForPeer) +TILEXR_MOONEP_WQE_BATCH_INLINE bool DispatchGroupedBatchNeedsCompletion( + uint32_t batchCount) { - return finalBatchForPeer; + return batchCount != 0U; +} + +TILEXR_MOONEP_WQE_BATCH_INLINE uint32_t DispatchRouteTileCount( + uint32_t routeCount, uint32_t tileStart, uint32_t tileCapacity) +{ + if (tileStart >= routeCount || tileCapacity == 0U) { + return 0U; + } + const uint32_t remaining = routeCount - tileStart; + return remaining < tileCapacity ? remaining : tileCapacity; } } // namespace TileXRMoonEp diff --git a/src/moonep/dispatch/urma/host/dispatch_host.cpp b/src/moonep/dispatch/urma/host/dispatch_host.cpp index ad1169c8..6e1440eb 100644 --- a/src/moonep/dispatch/urma/host/dispatch_host.cpp +++ b/src/moonep/dispatch/urma/host/dispatch_host.cpp @@ -64,7 +64,8 @@ bool DispatchVectorBatchShapeSupported(uint64_t routeCount, destinationCapacity >= routeCount && destinationCapacity <= UINT32_MAX && (destinationCapacity & (destinationCapacity - 1U)) == 0U && - DispatchPeerWqesFitSq(routeCount, TileXR::TILEXR_UDMA_SQ_BB_COUNT); + DispatchPeerWqesStreamable(routeCount, + TileXR::TILEXR_UDMA_SQ_BB_COUNT); } int ValidateDispatchPeerConfig(const DispatchPeerConfig &config, @@ -314,7 +315,8 @@ static int RunDispatchUrma(const TileXRMoonEpDispatchArgsV1 *args, 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 || + args->hiddenNvsh == nullptr || + (args->flags & ~TILEXR_MOONEP_FLAG_RESET_STATUS) != 0 || !PlanValid(args->plan) || !TensorDescriptorValid(args->hiddenSh) || !TensorDescriptorValid(args->hiddenNvsh) || (args->routeWeightsSk == nullptr) != (args->routeWeightsNvs == nullptr) || @@ -440,11 +442,6 @@ static int RunDispatchUrma(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); @@ -458,18 +455,34 @@ static int RunDispatchUrma(const TileXRMoonEpDispatchArgsV1 *args, params.zeroFillRangeCount = args->plan->e + args->plan->b; params.layout = layout; + const bool statusResetEnqueued = resetStatus || + (args->flags & TILEXR_MOONEP_FLAG_RESET_STATUS) != 0; + if (statusResetEnqueued && aclrtMemsetAsync(args->plan->status, + sizeof(int32_t), 0, sizeof(int32_t), stream) != ACL_SUCCESS) { + return TILEXR_MOONEP_ERROR_INTERNAL; + } + params.input = args->hiddenSh->data; params.output = args->hiddenNvsh->data; params.mode = DispatchPayloadMode::Hidden; ret = MapLaunchStatus(TileXRMoonEpLaunchDispatchUrmaKernel(params)); if (ret != TILEXR_MOONEP_SUCCESS || args->routeWeightsSk == nullptr) { + if (ret != TILEXR_MOONEP_SUCCESS && statusResetEnqueued && + aclrtSynchronizeStream(stream) != ACL_SUCCESS) { + return TILEXR_MOONEP_ERROR_INTERNAL; + } return ret; } params.input = args->routeWeightsSk->data; params.output = args->routeWeightsNvs->data; params.mode = DispatchPayloadMode::RouteWeight; - return MapLaunchStatus(TileXRMoonEpLaunchDispatchUrmaKernel(params)); + ret = MapLaunchStatus(TileXRMoonEpLaunchDispatchUrmaKernel(params)); + if (ret != TILEXR_MOONEP_SUCCESS && statusResetEnqueued && + aclrtSynchronizeStream(stream) != ACL_SUCCESS) { + return TILEXR_MOONEP_ERROR_INTERNAL; + } + return ret; } int TileXRMoonEpRunDispatchUrmaV1(const TileXRMoonEpDispatchArgsV1 *args, diff --git a/src/moonep/dispatch/urma/kernels/tilexr_moonep_dispatch_kernel.cpp b/src/moonep/dispatch/urma/kernels/tilexr_moonep_dispatch_kernel.cpp index ed6b7d4d..2c7e3acb 100644 --- a/src/moonep/dispatch/urma/kernels/tilexr_moonep_dispatch_kernel.cpp +++ b/src/moonep/dispatch/urma/kernels/tilexr_moonep_dispatch_kernel.cpp @@ -61,6 +61,7 @@ struct alignas(32) DispatchWqeBatchContext { uint32_t rmtTokenValue; uint32_t selectedStart; uint32_t qpSelection; + uint32_t routePlanStart; }; constexpr uint32_t kDispatchWqeBatchContextBytes = @@ -115,7 +116,8 @@ inline void DispatchBuildWriteWqeBatchVf(__ubuf__ uint8_t *wqeBytes, const uint32_t routeId = static_cast( selectedRouteIndices[selectedIndex]); const uint64_t targetSlot = static_cast( - static_cast(dstValues[routeId])) & + static_cast( + dstValues[routeId - context->routePlanStart])) & context->routeCountMask; const uint32_t sourceRow = context->hiddenMode != 0U ? AscendC::Simt::UintDiv(routeId, context->topKMagic, @@ -469,6 +471,33 @@ __aicore__ inline uint32_t SelectDispatchPeerRoutes( return static_cast(selectedCount); } +__aicore__ inline bool PrepareDispatchRouteTile( + AscendC::GlobalTensor dstGlobal, + AscendC::LocalTensor routePlanLocal, + AscendC::LocalTensor routeRankLocal, + AscendC::LocalTensor routeIndexLocal, + uint32_t routeTileStart, uint32_t routeTileCount, uint32_t routeShift) +{ + if (routeTileCount == 0U || routeTileStart > INT16_MAX || + routeTileCount > static_cast(INT16_MAX) + 1U - routeTileStart) { + return false; + } + const uint32_t routePlanDataBytes = routeTileCount * sizeof(int32_t); + const AscendC::DataCopyExtParams copyIn { + 1U, routePlanDataBytes, 0U, 0U, 0U}; + const AscendC::DataCopyPadExtParams padIn { + false, 0U, 0U, 0U}; + AscendC::DataCopyPad(routePlanLocal, dstGlobal[routeTileStart], + copyIn, padIn); + SyncFunc(); + AscendC::ShiftRight(routeRankLocal, routePlanLocal, + static_cast(routeShift), static_cast(routeTileCount)); + AscendC::CreateVecIndex(routeIndexLocal, + static_cast(routeTileStart), routeTileCount); + AscendC::PipeBarrier(); + return true; +} + __aicore__ inline bool BuildDispatchWriteWqeBatch( AscendC::LocalTensor issueLocal, AscendC::LocalTensor selectedRouteIndices, @@ -477,7 +506,7 @@ __aicore__ inline bool BuildDispatchWriteWqeBatch( uint64_t rowBytes, uint64_t routeCountMask, uint32_t topKMagic, uint32_t topKShift, bool hiddenMode, uint32_t selectedStart, uint32_t tokenCount, bool appendSignal, uint64_t signalLocalAddr, - uint32_t sequencePhase) + uint32_t sequencePhase, uint32_t routePlanStart) { __ubuf__ uint8_t *issueAddr = reinterpret_cast<__ubuf__ uint8_t *>( issueLocal.GetPhyAddr()); @@ -510,6 +539,7 @@ __aicore__ inline bool BuildDispatchWriteWqeBatch( context->selectedStart = selectedStart; context->qpSelection = (state.qpIdx << 2U) | (sequencePhase & 3U); + context->routePlanStart = routePlanStart; #if defined(CATLASS_ARCH) && CATLASS_ARCH == 3510 AscendC::PipeBarrier(); @@ -676,7 +706,7 @@ __aicore__ inline bool SubmitDispatchWqeBatch(DispatchWqeBatchState &state, state.head = batchEndHead; state.completionCount = batchEndCompletionCount; - state.outstanding += batchCount; + state.outstanding = state.head - state.tail; state.batchCount = 0U; state.batchLimit = TileXRMoonEp::DispatchWqeBatchCount( UINT64_MAX, state.head, TileXR::TILEXR_UDMA_SQ_BB_COUNT); @@ -692,7 +722,7 @@ __aicore__ inline bool AppendDispatchWqes(DispatchWqeBatchState &state, uint32_t topKShift, bool hiddenMode, uint32_t selectedRouteCount, bool appendSignal, uint64_t signalLocalAddr, uint32_t phase, uint32_t &dfxFlags, uint32_t &firstQuietStatus, uint32_t &firstQuietPhase, - uint32_t sequencePhase) + uint32_t sequencePhase, uint32_t routePlanStart) { uint32_t selectedStart = 0U; bool signalPending = appendSignal; @@ -715,7 +745,8 @@ __aicore__ inline bool AppendDispatchWqes(DispatchWqeBatchState &state, selectedRouteIndices, dstValues, state, localSourceBase, rowBytes, routeCountMask, topKMagic, topKShift, hiddenMode, selectedStart, tokenCount, - appendSignalNow, signalLocalAddr, sequencePhase)) { + appendSignalNow, signalLocalAddr, sequencePhase, + routePlanStart)) { return false; } state.batchCount += tokenCount + (appendSignalNow ? 1U : 0U); @@ -804,7 +835,7 @@ __aicore__ inline bool DispatchBuildGroupedQpBatch( uint64_t rowBytes, uint64_t routeCountMask, uint32_t topKMagic, uint32_t topKShift, bool hiddenMode, uint32_t selectedRouteCount, uint32_t &selectedStart, bool &signalPending, uint64_t signalLocalAddr, - uint32_t sequencePhase, bool &finalBatch) + uint32_t sequencePhase, uint32_t routePlanStart, bool &finalBatch) { finalBatch = false; state.batchCount = 0U; @@ -827,7 +858,8 @@ __aicore__ inline bool DispatchBuildGroupedQpBatch( if (!BuildDispatchWriteWqeBatch(issueLocal, selectedRouteIndices, dstValues, state, localSourceBase, rowBytes, routeCountMask, topKMagic, topKShift, hiddenMode, selectedStart, tokenCount, - appendSignalNow, signalLocalAddr, sequencePhase)) { + appendSignalNow, signalLocalAddr, sequencePhase, + routePlanStart)) { return false; } state.batchCount = tokenCount + (appendSignalNow ? 1U : 0U); @@ -846,13 +878,22 @@ __aicore__ inline bool StageDispatchQpBatch( uint32_t &firstQuietPhase) { const uint32_t batchCount = state.batchCount; - if (batchCount == 0U || state.doorbellPending != 0U || + if (!TileXRMoonEp::DispatchGroupedBatchNeedsCompletion(batchCount) || + state.doorbellPending != 0U || batchCount > state.batchLimit || !DispatchEnsureSqBatchCapacity(state, batchCount, cqeLocal, phase, dfxFlags, firstQuietStatus, firstQuietPhase)) { return false; } + __ubuf__ uint8_t *issueAddr = reinterpret_cast<__ubuf__ uint8_t *>( + issueLocal.GetPhyAddr()); + __ubuf__ TileXR::UDMASqeCtx *lastSqe = + reinterpret_cast<__ubuf__ TileXR::UDMASqeCtx *>( + issueAddr + (batchCount - 1U) * kDispatchUdmaWqeBytes); + lastSqe->flag = static_cast(lastSqe->flag) | + TileXR::TILEXR_UDMA_SQE_FLAG_COMPLETION; + __gm__ uint8_t *wqeAddr = reinterpret_cast<__gm__ uint8_t *>( state.qpCtxEntry->bufAddr + kDispatchUdmaWqeBytes * (state.head % TileXR::TILEXR_UDMA_SQ_BB_COUNT)); @@ -866,9 +907,8 @@ __aicore__ inline bool StageDispatchQpBatch( SyncFunc(); state.head += batchCount; - state.completionCount += - TileXRMoonEp::DispatchBatchNeedsCompletion(finalBatch) ? 1U : 0U; - state.outstanding += batchCount; + state.completionCount += 1U; + state.outstanding = state.head - state.tail; state.stagedDoorbellHead = state.head; state.doorbellPending = 1U; if (finalBatch) { @@ -918,6 +958,22 @@ __aicore__ inline bool DispatchDrainPeerFinalCq( return true; } +__aicore__ inline void ProbeDispatchPeerFinalCq( + DispatchPreparedPeer &peer, AscendC::LocalTensor cqeLocal, + uint64_t &pollStatuses, uint64_t &remainingSqEntries) +{ + pollStatuses = 0U; + remainingSqEntries = 0U; + for (uint32_t qpIdx = 0U; + qpIdx < TileXRMoonEp::kDispatchQpCount; ++qpIdx) { + DispatchWqeBatchState &state = peer.qpState[qpIdx]; + const uint32_t pollStatus = DispatchPollCqBatch(state, cqeLocal); + const uint32_t remaining = state.finalHead - state.tail; + pollStatuses |= static_cast(pollStatus) << (qpIdx * 32U); + remainingSqEntries |= static_cast(remaining) << (qpIdx * 32U); + } +} + __aicore__ inline bool DecodeSendDst(int32_t encoded, uint64_t destinationCapacity, int32_t rankSize, int32_t &targetRank, uint64_t &targetSlot) { @@ -1033,31 +1089,50 @@ __aicore__ inline void CopyContiguousBytesGmToGmPipelined( outputCopyQueue.FreeTensor(pendingLocal); } -__aicore__ inline uint64_t LoadCompletionFlag(__gm__ uint64_t *flag, +__aicore__ inline void PublishDispatchSignalSource( + __gm__ uint64_t *signalSource, uint64_t expectedFlag, AscendC::LocalTensor relayLocal) { - AscendC::GlobalTensor flagGlobal; - flagGlobal.SetGlobalBuffer(reinterpret_cast<__gm__ uint8_t *>(flag)); - const AscendC::DataCopyExtParams copyIn { - 1U, static_cast(sizeof(uint64_t)), 0U, 0U, 0U}; - const AscendC::DataCopyPadExtParams padIn {false, 0U, 0U, 0U}; - AscendC::DataCopyPad(relayLocal, flagGlobal, copyIn, padIn); - AscendC::SetFlag(EVENT_ID0); - AscendC::WaitFlag(EVENT_ID0); - return relayLocal.ReinterpretCast().GetValue(0); + AscendC::LocalTensor signalLocal = + relayLocal.ReinterpretCast(); + for (uint32_t qpIdx = 0U; + qpIdx < TileXRMoonEp::kDispatchQpCount; ++qpIdx) { + signalLocal.SetValue(qpIdx, expectedFlag); + } + AscendC::GlobalTensor signalGlobal; + signalGlobal.SetGlobalBuffer(signalSource, + TileXRMoonEp::kDispatchQpCount); + const AscendC::DataCopyExtParams copyOut { + 1U, TileXRMoonEp::kDispatchQpCount * + static_cast(sizeof(uint64_t)), 0U, 0U, 0U}; + SyncFunc(); + AscendC::DataCopyPad(signalGlobal, signalLocal, copyOut); + SyncFunc(); + TileXR::UDMACleanCacheLines( + reinterpret_cast<__gm__ uint8_t *>(signalSource), + TileXRMoonEp::kDispatchQpCount * sizeof(uint64_t)); + AscendC::PipeBarrier(); +} + +__aicore__ inline uint64_t LoadCompletionFlag(__gm__ uint64_t *flag) +{ + TileXR::UDMACleanCacheLines( + reinterpret_cast<__gm__ uint8_t *>(flag), sizeof(uint64_t)); + AscendC::PipeBarrier(); + return flag[0]; } __aicore__ inline bool WaitCompletionFlag(__gm__ uint64_t *flag, uint64_t expected, uint64_t waitStart, uint64_t timeoutTicks, - uint64_t &observed, AscendC::LocalTensor relayLocal) + uint64_t &observed) { - observed = LoadCompletionFlag(flag, relayLocal); + observed = LoadCompletionFlag(flag); while (observed < expected) { if (static_cast(AscendC::GetSystemCycle()) - waitStart >= timeoutTicks) { return false; } - observed = LoadCompletionFlag(flag, relayLocal); + observed = LoadCompletionFlag(flag); } return true; } @@ -1069,6 +1144,10 @@ __aicore__ inline uint64_t LoadDispatchCredit( creditGlobal.SetGlobalBuffer(credit, TileXRMoonEp::kDispatchCreditStrideBytes / sizeof(uint64_t)); auto creditLocal = relayLocal.ReinterpretCast(); + TileXR::UDMACleanCacheLines( + reinterpret_cast<__gm__ uint8_t *>(credit), + TileXRMoonEp::kDispatchCreditStrideBytes); + SyncFunc(); AscendC::DataCopy(creditLocal, creditGlobal, TileXRMoonEp::kDispatchCreditStrideBytes / sizeof(uint64_t)); SyncFunc(); @@ -1162,7 +1241,7 @@ __aicore__ inline bool WaitDispatchIncomingPeerAndPublishCredit( uint64_t observed = 0U; if (!WaitCompletionFlag(receiveFlags + peerFlagBase + qpIdx, static_cast(magic), waitStart, timeoutTicks, - observed, relayLocal)) { + observed)) { dfxFlags |= TileXRMoonEp::kDispatchDfxCompletionTimeout; if (timeoutPeer == UINT32_MAX) { timeoutPeer = static_cast(incomingPeer); @@ -1203,7 +1282,10 @@ __aicore__ inline void WriteDfxRecord(__gm__ uint8_t *dfx, uint32_t firstQuietStatus, uint32_t firstQuietPhase, uint32_t timeoutPeer, uint32_t timeoutPhase, uint64_t expectedRouteCount, uint64_t processedRouteCount, uint64_t magic, uint64_t timeoutExpectedMagic, - uint64_t timeoutObservedFlag, AscendC::LocalTensor diagnosticLocal) + uint64_t timeoutObservedFlag, uint64_t localSignalObserved, + uint64_t completionFlagCount, uint64_t outgoingCqStatuses, + uint64_t outgoingRemainingSqEntries, + AscendC::LocalTensor diagnosticLocal) { __ubuf__ TileXRMoonEp::DispatchDfxRecord *record = reinterpret_cast<__ubuf__ TileXRMoonEp::DispatchDfxRecord *>( @@ -1227,9 +1309,10 @@ __aicore__ inline void WriteDfxRecord(__gm__ uint8_t *dfx, record->magic = magic; record->timeoutExpectedMagic = timeoutExpectedMagic; record->timeoutObservedFlag = timeoutObservedFlag; - for (uint32_t index = 0U; index < 4U; ++index) { - record->reserved[index] = 0U; - } + record->reserved[0] = localSignalObserved; + record->reserved[1] = completionFlagCount; + record->reserved[2] = outgoingCqStatuses; + record->reserved[3] = outgoingRemainingSqEntries; SyncFunc(); AscendC::GlobalTensor dfxGlobal; @@ -1350,6 +1433,21 @@ __aicore__ inline int32_t StatusFromDfxFlags(uint32_t flags) return TileXRMoonEp::kDispatchStatusSuccess; } +__aicore__ inline int32_t LoadDispatchPlanStatus(__gm__ int32_t *planStatus, + AscendC::LocalTensor relayLocal) +{ + AscendC::GlobalTensor statusGlobal; + statusGlobal.SetGlobalBuffer(reinterpret_cast<__gm__ uint8_t *>(planStatus), + sizeof(int32_t)); + const AscendC::DataCopyExtParams copyIn { + 1U, static_cast(sizeof(int32_t)), 0U, 0U, 0U}; + const AscendC::DataCopyPadExtParams padIn { + false, 0U, 0U, 0U}; + AscendC::DataCopyPad(relayLocal, statusGlobal, copyIn, padIn); + SyncFunc(); + return relayLocal.ReinterpretCast().GetValue(0U); +} + __aicore__ inline void DispatchPublishFirstStatus( __gm__ int32_t *planStatus, int32_t status) { @@ -1542,13 +1640,20 @@ extern "C" __global__ __aicore__ void tilexr_moonep_dispatch_urma_kernel( uint32_t routeIndexUbBytes = 0U; uint32_t compareMaskUbBytes = 0U; if (useVectorSlotSelect) { + const uint32_t vectorElementCount = groupedPeerMode ? + TileXRMoonEp::DispatchRouteTileCount( + static_cast(routeCount), 0U, + kRouteTileElements) : + static_cast(routeCount); routePlanUbBytes = static_cast(AlignUp( - routeCount * sizeof(int32_t), kUbAlignBytes)); + static_cast(vectorElementCount) * sizeof(int32_t), + kUbAlignBytes)); routeRankUbBytes = routePlanUbBytes; routeIndexUbBytes = static_cast(AlignUp( - routeCount * sizeof(int16_t), kUbAlignBytes)); + static_cast(vectorElementCount) * sizeof(int16_t), + kUbAlignBytes)); compareMaskUbBytes = static_cast(AlignUp( - CeilDiv(routeCount, 8U), kUbAlignBytes)); + CeilDiv(vectorElementCount, 8U), kUbAlignBytes)); const uint64_t fixedBytes = static_cast(routePlanUbBytes) + routeRankUbBytes + 2ULL * routeIndexUbBytes + compareMaskUbBytes + kDiagnosticUbBytes + @@ -1640,19 +1745,14 @@ extern "C" __global__ __aicore__ void tilexr_moonep_dispatch_urma_kernel( const uint64_t stagingCycles = stagingEndCycle - stagingStartCycle; #endif - const int32_t upstreamStatus = *planStatus; + const int32_t upstreamStatus = LoadDispatchPlanStatus(planStatus, relayLocal); auto currentScratch = workspace + scratchOffset + scratchIndex * scratchSlotBytes; auto receiveFlags = reinterpret_cast<__gm__ uint64_t *>( workspace + completionFlagsOffset); auto signalSource = reinterpret_cast<__gm__ uint64_t *>( workspace + signalOffset + blockIdx * TileXRMoonEp::kDispatchSignalStrideBytes); if (!localOnly) { - for (uint32_t qpIdx = 0U; - qpIdx < TileXRMoonEp::kDispatchQpCount; ++qpIdx) { - signalSource[qpIdx] = expectedFlag; - } - TileXR::UDMACleanCacheLines(reinterpret_cast<__gm__ uint8_t *>(signalSource), - TileXRMoonEp::kDispatchQpCount * sizeof(uint64_t)); + PublishDispatchSignalSource(signalSource, expectedFlag, relayLocal); } AscendC::GlobalTensor dstGlobal; @@ -1666,6 +1766,8 @@ extern "C" __global__ __aicore__ void tilexr_moonep_dispatch_urma_kernel( uint32_t timeoutPeer = UINT32_MAX; uint32_t timeoutPhase = UINT32_MAX; uint64_t timeoutObservedFlag = 0U; + uint64_t outgoingCqStatuses = UINT64_MAX; + uint64_t outgoingRemainingSqEntries = UINT64_MAX; uint64_t scannedRouteCount = 0U; uint64_t matchedRouteCount = 0U; uint64_t selectedRouteCount = 0U; @@ -1682,22 +1784,16 @@ extern "C" __global__ __aicore__ void tilexr_moonep_dispatch_urma_kernel( uint64_t quietCycles = 0U; #endif - if (useVectorSlotSelect && upstreamStatus == TileXRMoonEp::kDispatchStatusSuccess) { - const uint32_t routePlanDataBytes = static_cast( - routeCount * sizeof(int32_t)); - const AscendC::DataCopyExtParams copyIn {1U, routePlanDataBytes, 0U, 0U, 0U}; - const AscendC::DataCopyPadExtParams padIn {false, 0U, 0U, 0U}; - AscendC::DataCopyPad(routePlanLocal, dstGlobal, copyIn, padIn); - SyncFunc(); - uint32_t routeShift = 0U; - for (uint64_t value = destinationCapacity; value > 1U; value >>= 1U) { - ++routeShift; - } - AscendC::ShiftRight(routeRankLocal, routePlanLocal, - static_cast(routeShift), static_cast(routeCount)); - AscendC::CreateVecIndex(routeIndexLocal, static_cast(0), - static_cast(routeCount)); - AscendC::PipeBarrier(); + uint32_t routeShift = 0U; + for (uint64_t value = destinationCapacity; value > 1U; value >>= 1U) { + ++routeShift; + } + if (useVectorSlotSelect && !groupedPeerMode && + upstreamStatus == TileXRMoonEp::kDispatchStatusSuccess && + !PrepareDispatchRouteTile(dstGlobal, routePlanLocal, + routeRankLocal, routeIndexLocal, 0U, + static_cast(routeCount), routeShift)) { + dfxFlags |= TileXRMoonEp::kDispatchDfxInvalidConfig; } #if defined(TILEXR_MOONEP_DISPATCH_ENABLE_PROFILING) @@ -1748,82 +1844,107 @@ extern "C" __global__ __aicore__ void tilexr_moonep_dispatch_urma_kernel( break; } - const uint32_t selectedCount = SelectDispatchPeerRoutes( - compareMaskLocal, routeRankLocal, routeIndexLocal, - selectedRouteIndexLocal, peer, - static_cast(routeCount)); - scannedRouteCount += routeCount; - matchedRouteCount += selectedCount; - selectedRouteCount += selectedCount; - - uint32_t qpSelectedCount[TileXRMoonEp::kDispatchQpCount] = {}; - uint32_t qpSelectedStart[TileXRMoonEp::kDispatchQpCount] = {}; - bool qpSignalPending[TileXRMoonEp::kDispatchQpCount] = {true, true}; - for (uint32_t qpIdx = 0U; - qpIdx < TileXRMoonEp::kDispatchQpCount; ++qpIdx) { - qpSelectedCount[qpIdx] = TileXRMoonEp::DispatchQpRouteCount( - selectedCount, 0U, qpIdx); - } - bool firstLogicalBatch = true; bool sendOk = true; - while (sendOk && (qpSelectedStart[0] < qpSelectedCount[0] || - qpSelectedStart[1] < qpSelectedCount[1] || - qpSignalPending[0] || qpSignalPending[1])) { - bool finalBatch[TileXRMoonEp::kDispatchQpCount] = {}; - for (uint32_t qpIdx = 0U; - sendOk && qpIdx < TileXRMoonEp::kDispatchQpCount; - ++qpIdx) { - DispatchWqeBatchState &qpState = - preparedPeer.qpState[qpIdx]; - AscendC::LocalTensor issueLocal = qpIdx == 0U ? - udmaIssueQp0Local : udmaIssueQp1Local; - sendOk = DispatchBuildGroupedQpBatch(qpState, - issueLocal, selectedRouteIndexLocal, routePlanLocal, - reinterpret_cast(workspace), rowBytes, - destinationCapacity - 1U, topKMagic, topKShift, hiddenMode, - qpSelectedCount[qpIdx], qpSelectedStart[qpIdx], - qpSignalPending[qpIdx], reinterpret_cast( - signalSource + qpIdx), 0U, finalBatch[qpIdx]); + uint64_t peerSelectedCount = 0U; + for (uint32_t routeTileStart = 0U; + sendOk && routeTileStart < routeCount;) { + const uint32_t routeTileCount = + TileXRMoonEp::DispatchRouteTileCount( + static_cast(routeCount), routeTileStart, + kRouteTileElements); + if (!PrepareDispatchRouteTile(dstGlobal, routePlanLocal, + routeRankLocal, routeIndexLocal, routeTileStart, + routeTileCount, routeShift)) { + sendOk = false; + break; } + const uint32_t selectedCount = SelectDispatchPeerRoutes( + compareMaskLocal, routeRankLocal, routeIndexLocal, + selectedRouteIndexLocal, peer, routeTileCount); + scannedRouteCount += routeTileCount; + matchedRouteCount += selectedCount; + selectedRouteCount += selectedCount; + peerSelectedCount += selectedCount; + + uint32_t qpSelectedCount[TileXRMoonEp::kDispatchQpCount] = {}; + uint32_t qpSelectedStart[TileXRMoonEp::kDispatchQpCount] = {}; + const bool finalRouteTile = + routeTileStart + routeTileCount == routeCount; + bool qpSignalPending[TileXRMoonEp::kDispatchQpCount] = { + finalRouteTile, finalRouteTile}; for (uint32_t qpIdx = 0U; - sendOk && qpIdx < TileXRMoonEp::kDispatchQpCount; - ++qpIdx) { - if (preparedPeer.qpState[qpIdx].batchCount == 0U) { - continue; - } - sendOk = StageDispatchQpBatch( - preparedPeer.qpState[qpIdx], - qpIdx == 0U ? udmaIssueQp0Local : udmaIssueQp1Local, - relayLocal, finalBatch[qpIdx], group, dfxFlags, - firstQuietStatus, firstQuietPhase); - } - if (sendOk && firstLogicalBatch && previousPeerValid) { - sendOk = DispatchDrainPeerFinalCq(previousPeer, - relayLocal, completionTimeoutTicks, dfxFlags, - firstQuietStatus, firstQuietPhase, timeoutPeer, - timeoutPhase, timeoutObservedFlag); - previousPeerValid = false; + qpIdx < TileXRMoonEp::kDispatchQpCount; ++qpIdx) { + qpSelectedCount[qpIdx] = + TileXRMoonEp::DispatchQpRouteCount( + selectedCount, 0U, qpIdx); } - if (sendOk && firstLogicalBatch && creditPeerMode && - TileXRMoonEp::DispatchCreditRequired(group)) { - uint64_t observedCredit = 0U; - if (!WaitDispatchPeerCredit(args, rank, peer, group, - magic, completionTimeoutTicks, relayLocal, - observedCredit)) { - dfxFlags |= TileXRMoonEp::kDispatchDfxCreditTimeout; - if (timeoutPeer == UINT32_MAX) { - timeoutPeer = static_cast(peer); - timeoutPhase = group; - timeoutObservedFlag = observedCredit; + + while (sendOk && + (qpSelectedStart[0] < qpSelectedCount[0] || + qpSelectedStart[1] < qpSelectedCount[1] || + qpSignalPending[0] || qpSignalPending[1])) { + bool finalBatch[TileXRMoonEp::kDispatchQpCount] = {}; + for (uint32_t qpIdx = 0U; + sendOk && qpIdx < TileXRMoonEp::kDispatchQpCount; + ++qpIdx) { + DispatchWqeBatchState &qpState = + preparedPeer.qpState[qpIdx]; + AscendC::LocalTensor issueLocal = + qpIdx == 0U ? udmaIssueQp0Local : + udmaIssueQp1Local; + sendOk = DispatchBuildGroupedQpBatch(qpState, + issueLocal, selectedRouteIndexLocal, + routePlanLocal, + reinterpret_cast(workspace), rowBytes, + destinationCapacity - 1U, topKMagic, topKShift, + hiddenMode, qpSelectedCount[qpIdx], + qpSelectedStart[qpIdx], qpSignalPending[qpIdx], + reinterpret_cast(signalSource + qpIdx), + 0U, routeTileStart, finalBatch[qpIdx]); + } + for (uint32_t qpIdx = 0U; + sendOk && qpIdx < TileXRMoonEp::kDispatchQpCount; + ++qpIdx) { + if (preparedPeer.qpState[qpIdx].batchCount == 0U) { + continue; } - sendOk = false; + sendOk = StageDispatchQpBatch( + preparedPeer.qpState[qpIdx], + qpIdx == 0U ? udmaIssueQp0Local : + udmaIssueQp1Local, + relayLocal, finalBatch[qpIdx], group, dfxFlags, + firstQuietStatus, firstQuietPhase); } + if (sendOk && firstLogicalBatch && previousPeerValid) { + sendOk = DispatchDrainPeerFinalCq(previousPeer, + relayLocal, completionTimeoutTicks, dfxFlags, + firstQuietStatus, firstQuietPhase, timeoutPeer, + timeoutPhase, timeoutObservedFlag); + previousPeerValid = false; + } + if (sendOk && firstLogicalBatch && creditPeerMode && + TileXRMoonEp::DispatchCreditRequired(group)) { + uint64_t observedCredit = 0U; + if (!WaitDispatchPeerCredit(args, rank, peer, group, + magic, completionTimeoutTicks, relayLocal, + observedCredit)) { + dfxFlags |= + TileXRMoonEp::kDispatchDfxCreditTimeout; + if (timeoutPeer == UINT32_MAX) { + timeoutPeer = static_cast(peer); + timeoutPhase = group; + timeoutObservedFlag = observedCredit; + } + sendOk = false; + } + } + if (sendOk) { + RingDispatchPeerDoorbells(preparedPeer); + } + firstLogicalBatch = false; } - if (sendOk) { - RingDispatchPeerDoorbells(preparedPeer); - } - firstLogicalBatch = false; + routeTileStart += routeTileCount; } if (!sendOk || preparedPeer.qpState[0].finalStaged == 0U || @@ -1835,9 +1956,9 @@ extern "C" __global__ __aicore__ void tilexr_moonep_dispatch_urma_kernel( break; } - issuedPutCount += selectedCount; - issuedPutBytes += static_cast(selectedCount) * rowBytes; - processedRouteCount += selectedCount; + issuedPutCount += peerSelectedCount; + issuedPutBytes += peerSelectedCount * rowBytes; + processedRouteCount += peerSelectedCount; completionFlagCount += TileXRMoonEp::kDispatchQpCount; previousPeer = preparedPeer; previousPeerValid = true; @@ -1847,6 +1968,8 @@ extern "C" __global__ __aicore__ void tilexr_moonep_dispatch_urma_kernel( groupWidth, magic, creditPeerMode, completionTimeoutTicks, relayLocal, dfxFlags, timeoutPeer, timeoutPhase, timeoutObservedFlag)) { + ProbeDispatchPeerFinalCq(previousPeer, relayLocal, + outgoingCqStatuses, outgoingRemainingSqEntries); break; } #if defined(TILEXR_MOONEP_DISPATCH_ENABLE_PROFILING) @@ -1930,7 +2053,7 @@ extern "C" __global__ __aicore__ void tilexr_moonep_dispatch_urma_kernel( reinterpret_cast( signalSource + qpIdx), issuePhase, dfxFlags, firstQuietStatus, - firstQuietPhase, sequencePhase); + firstQuietPhase, sequencePhase, 0U); } if (batchOk) { issuedPutCount += selectedCount; @@ -2153,7 +2276,7 @@ extern "C" __global__ __aicore__ void tilexr_moonep_dispatch_urma_kernel( static_cast(AscendC::GetSystemCycle()); #endif if (upstreamStatus == TileXRMoonEp::kDispatchStatusSuccess) { - if (useVectorSlotSelect) { + if (useVectorSlotSelect && !groupedPeerMode) { const uint32_t localRouteCount = SelectDispatchPeerRoutes( compareMaskLocal, routeRankLocal, routeIndexLocal, selectedRouteIndexLocal, rank, @@ -2267,8 +2390,7 @@ extern "C" __global__ __aicore__ void tilexr_moonep_dispatch_urma_kernel( if (!WaitCompletionFlag( receiveFlags + peerFlagBase + qpIdx, expectedFlag, flagWaitStartCycle, - completionTimeoutTicks, observed, - relayLocal)) { + completionTimeoutTicks, observed)) { dfxFlags |= TileXRMoonEp:: kDispatchDfxCompletionTimeout; if (timeoutPeer == UINT32_MAX) { @@ -2301,11 +2423,17 @@ extern "C" __global__ __aicore__ void tilexr_moonep_dispatch_urma_kernel( const int32_t localExecutionStatus = StatusFromDfxFlags(dfxFlags); DispatchPublishFirstStatus(planStatus, localExecutionStatus); #if defined(TILEXR_MOONEP_DISPATCH_ENABLE_DFX) + uint64_t localSignalObserved = expectedFlag; + if (!localOnly && dfxFlags != 0U) { + localSignalObserved = LoadCompletionFlag(signalSource); + } WriteDfxRecord(workspace + dfxOffset, static_cast(payloadMode), static_cast(rank), static_cast(blockIdx), dfxFlags, firstInvalidRouteId, firstInvalidRawDst, firstQuietStatus, firstQuietPhase, timeoutPeer, timeoutPhase, routeCount, processedRouteCount, expectedFlag, - expectedFlag, timeoutObservedFlag, diagnosticLocal); + expectedFlag, timeoutObservedFlag, localSignalObserved, + completionFlagCount, outgoingCqStatuses, + outgoingRemainingSqEntries, diagnosticLocal); #endif #if defined(TILEXR_MOONEP_DISPATCH_ENABLE_PROFILING) const uint64_t dfxWriteEndCycle = @@ -2327,7 +2455,7 @@ extern "C" __global__ __aicore__ void tilexr_moonep_dispatch_urma_kernel( } const int32_t executionStatus = StatusFromDfxFlags(globalDfxFlags); #else - const int32_t executionStatus = *planStatus; + const int32_t executionStatus = LoadDispatchPlanStatus(planStatus, relayLocal); #endif auto kernelStatus = reinterpret_cast<__gm__ TileXRMoonEp::DispatchKernelStatus *>( workspace + kernelStatusOffset); @@ -2423,7 +2551,9 @@ extern "C" __global__ __aicore__ void tilexr_moonep_dispatch_urma_kernel( static_cast(rank), static_cast(blockIdx), dfxFlags, firstInvalidRouteId, firstInvalidRawDst, firstQuietStatus, firstQuietPhase, timeoutPeer, timeoutPhase, routeCount, processedRouteCount, expectedFlag, - expectedFlag, timeoutObservedFlag, diagnosticLocal); + expectedFlag, timeoutObservedFlag, localSignalObserved, + completionFlagCount, outgoingCqStatuses, + outgoingRemainingSqEntries, diagnosticLocal); #endif #if defined(TILEXR_MOONEP_DISPATCH_ENABLE_PROFILING) diff --git a/src/moonep/planner/host/planner_launch.cpp b/src/moonep/planner/host/planner_launch.cpp index 678f2df4..088fc2f2 100644 --- a/src/moonep/planner/host/planner_launch.cpp +++ b/src/moonep/planner/host/planner_launch.cpp @@ -35,7 +35,6 @@ int TileXRMoonEpLaunchKernel(const PlannerParams ¶ms, const PlannerLaunchCon if (ret != TileXR::TILEXR_SUCCESS) { return ret; } - const PlannerLayout &layout = context.layout; struct PlannerKernelArgs { GM_ADDR commArgs; diff --git a/src/moonep/prefetch_weight/kernels/tilexr_moonep_prefetch_weight_kernel.cpp b/src/moonep/prefetch_weight/kernels/tilexr_moonep_prefetch_weight_kernel.cpp index 1ebd4d9b..b7fd7f5f 100644 --- a/src/moonep/prefetch_weight/kernels/tilexr_moonep_prefetch_weight_kernel.cpp +++ b/src/moonep/prefetch_weight/kernels/tilexr_moonep_prefetch_weight_kernel.cpp @@ -9,7 +9,7 @@ namespace TileXRMoonEp { namespace Kernel { -constexpr uint32_t kMaxTrackedRankSize = 1024; +constexpr uint32_t kMaxTrackedRankSize = TileXR::TILEXR_MAX_RANK_SIZE; constexpr uint32_t kUsedPeerWordCount = kMaxTrackedRankSize / 64; class PrefetchWeightKernel { @@ -51,10 +51,16 @@ class PrefetchWeightKernel { AscendC::SyncAll(); uint64_t usedPeers[kUsedPeerWordCount] = {}; + uint64_t completionQueueIds[kMaxTrackedRankSize] = {}; + int32_t completionQueuePeers[kMaxTrackedRankSize] = {}; + uint32_t completionTargets[kMaxTrackedRankSize] = {}; + uint32_t completionQueueCount = 0U; uint32_t workerStatus = ValidateRuntime(); if (workerStatus == 0) { - SubmitReads(usedPeers, workerStatus); - CompleteReads(usedPeers, workerStatus); + SubmitReads(usedPeers, completionQueueIds, completionQueuePeers, + completionTargets, completionQueueCount, workerStatus); + CompleteReads(usedPeers, completionQueuePeers, completionTargets, + completionQueueCount, workerStatus); } if (workerStatus != 0) { (void)AscendC::AtomicCas(status_, static_cast(0), workerStatus); @@ -118,7 +124,11 @@ class PrefetchWeightKernel { } __aicore__ inline void SubmitReads( - uint64_t usedPeers[kUsedPeerWordCount], uint32_t &workerStatus) + uint64_t usedPeers[kUsedPeerWordCount], + uint64_t completionQueueIds[kMaxTrackedRankSize], + int32_t completionQueuePeers[kMaxTrackedRankSize], + uint32_t completionTargets[kMaxTrackedRankSize], + uint32_t &completionQueueCount, uint32_t &workerStatus) { auto wqeScratch = wqeBuf_.Get(); const int64_t globalExpertCount = expertsPerRank_ * rankSize_; @@ -145,6 +155,14 @@ class PrefetchWeightKernel { const int32_t localExpert = expert % static_cast(expertsPerRank_); MarkPeer(usedPeers, owner); + __gm__ TileXR::UDMAWQCtx *queue = TileXR::UDMAGetWQCtx( + TileXR::GetUDMAInfo(args_), static_cast(owner), worker_); + const uint32_t completionQueue = TrackCompletionQueue( + queue->wqeCntAddr, owner, completionQueueIds, completionQueuePeers, + completionTargets, completionQueueCount, workerStatus); + if (completionQueue >= kMaxTrackedRankSize) { + continue; + } for (uint32_t projection = 0; projection < 3; ++projection) { const uint64_t sourceOffset = offsets_[projection] + static_cast(localExpert) * rowBytes_[projection]; @@ -158,26 +176,60 @@ class PrefetchWeightKernel { workerStatus == 0) { workerStatus = kPrefetchWeightStatusSubmitErrorBase + (submitStatus & 0xFFU); + } else if (submitStatus == TileXR::TILEXR_UDMA_STATUS_SUCCESS) { + ++completionTargets[completionQueue]; } } } } + __aicore__ inline uint32_t TrackCompletionQueue( + uint64_t queueId, int32_t peer, + uint64_t completionQueueIds[kMaxTrackedRankSize], + int32_t completionQueuePeers[kMaxTrackedRankSize], + uint32_t completionTargets[kMaxTrackedRankSize], + uint32_t &completionQueueCount, uint32_t &workerStatus) const + { + for (uint32_t queue = 0U; queue < completionQueueCount; ++queue) { + if (completionQueueIds[queue] == queueId) { + return queue; + } + } + if (queueId == 0U || completionQueueCount >= kMaxTrackedRankSize) { + if (workerStatus == 0U) { + workerStatus = kPrefetchWeightStatusInvalidRuntime; + } + return kMaxTrackedRankSize; + } + const uint32_t queue = completionQueueCount++; + completionQueueIds[queue] = queueId; + completionQueuePeers[queue] = peer; + completionTargets[queue] = ld_dev( + reinterpret_cast<__gm__ uint32_t *>(queueId), 0); + return queue; + } + __aicore__ inline void CompleteReads( - const uint64_t usedPeers[kUsedPeerWordCount], uint32_t &workerStatus) + const uint64_t usedPeers[kUsedPeerWordCount], + const int32_t completionQueuePeers[kMaxTrackedRankSize], + const uint32_t completionTargets[kMaxTrackedRankSize], + uint32_t completionQueueCount, uint32_t &workerStatus) { - bool completedAny = false; - for (int32_t peer = 0; peer < rankSize_; ++peer) { + for (uint32_t queue = 0U; queue < completionQueueCount; ++queue) { + const int32_t peer = completionQueuePeers[queue]; if (!PeerUsed(usedPeers, peer)) { + if (workerStatus == 0U) { + workerStatus = kPrefetchWeightStatusInvalidRuntime; + } continue; } - completedAny = true; - const uint32_t cqStatus = TileXR::UDMAQuietStatusOnQp(args_, peer, worker_); + const uint32_t cqStatus = TileXR::UDMAQuietStatusOnQpUntil( + args_, peer, worker_, completionTargets[queue]); if (cqStatus != 0 && workerStatus == 0) { workerStatus = kPrefetchWeightStatusCqErrorBase + (cqStatus & 0xFFU); } } - if (completedAny) { + if (completionQueueCount != 0U) { AscendC::GlobalTensor cache; cache.SetGlobalBuffer(reinterpret_cast<__gm__ uint64_t *>(0), 1); AscendC::DataCacheCleanAndInvalid(target), qpIdx); + if (queue == nullptr || queue->wqeCntAddr == 0U) { + SetStatus(kReduceGradDeviceUdmaCqError); + return false; + } + completionTargets[qpIdx] = ld_dev( + reinterpret_cast<__gm__ uint32_t *>(queue->wqeCntAddr), 0); + } + return true; + } + __aicore__ inline bool WaitUdmaCompletion(int64_t target, uint32_t qpIdx, - uint64_t stage, uint64_t sequence) + uint64_t stage, uint64_t sequence, + uint32_t completionTargets[kReduceGradMaxUdmaQpCount]) { const uint64_t remoteOffset = UDMACompletionOffset(rank_, stage); GM_ADDR localScratch = workspace_ + UDMAPollScratchOffset(target, stage); @@ -275,8 +296,13 @@ class ReduceGradKernel { static_cast(target), qpIdx, reinterpret_cast<__gm__ uint8_t *>(localScratch), remoteOffset, static_cast(sizeof(uint64_t))); - if (getStatus != TileXR::TILEXR_UDMA_STATUS_SUCCESS || - TileXR::UDMAQuietStatusOnQp(args_, static_cast(target), qpIdx) != + if (getStatus != TileXR::TILEXR_UDMA_STATUS_SUCCESS) { + SetStatus(kReduceGradDeviceUdmaCqError); + return false; + } + const uint32_t completionTarget = ++completionTargets[qpIdx]; + if (TileXR::UDMAQuietStatusOnQpUntil(args_, + static_cast(target), qpIdx, completionTarget) != TileXR::TILEXR_UDMA_STATUS_SUCCESS) { SetStatus(kReduceGradDeviceUdmaCqError); return false; @@ -455,7 +481,8 @@ class ReduceGradKernel { } __aicore__ inline bool PostUdmaChunk(int64_t target, int64_t slot, - uint32_t projection, uint64_t chunk, GM_ADDR source, uint64_t bytes) + uint32_t projection, uint64_t chunk, GM_ADDR source, uint64_t bytes, + uint32_t completionTargets[kReduceGradMaxUdmaQpCount]) { const uint64_t stage = chunk & 1U; const uint64_t outboundOffset = UDMAOutboundOffset(target, stage); @@ -472,8 +499,13 @@ class ReduceGradKernel { reinterpret_cast<__gm__ uint8_t *>(outbound), UDMAInboundOffset(rank_, stage), static_cast(bytes), UDMAReadyOffset(rank_, stage), sequence); - if (putStatus != TileXR::TILEXR_UDMA_STATUS_SUCCESS || - TileXR::UDMAQuietStatusOnQp(args_, static_cast(target), qpIdx) != + if (putStatus != TileXR::TILEXR_UDMA_STATUS_SUCCESS) { + SetStatus(kReduceGradDeviceUdmaCqError); + return false; + } + const uint32_t completionTarget = ++completionTargets[qpIdx]; + if (TileXR::UDMAQuietStatusOnQpUntil(args_, + static_cast(target), qpIdx, completionTarget) != TileXR::TILEXR_UDMA_STATUS_SUCCESS) { SetStatus(kReduceGradDeviceUdmaCqError); return false; @@ -482,7 +514,8 @@ class ReduceGradKernel { } __aicore__ inline bool CompleteChunk(int64_t target, int64_t slot, - uint32_t projection, uint64_t chunk, GM_ADDR source, uint64_t bytes) + uint32_t projection, uint64_t chunk, GM_ADDR source, uint64_t bytes, + uint32_t completionTargets[kReduceGradMaxUdmaQpCount]) { bool complete = false; if (transports_[projection] == kReduceGradTransportPeer) { @@ -491,7 +524,8 @@ class ReduceGradKernel { } else { const uint32_t qpIdx = UDMAQpFor(projection, slot, chunk); complete = WaitUdmaCompletion( - target, qpIdx, chunk & 1U, Sequence(projection, slot, chunk)); + target, qpIdx, chunk & 1U, Sequence(projection, slot, chunk), + completionTargets); } if (!complete) { return false; @@ -506,6 +540,10 @@ class ReduceGradKernel { SetStatus(kReduceGradDeviceInvalidState); return false; } + uint32_t completionTargets[kReduceGradMaxUdmaQpCount] = {}; + if (!InitializeUdmaCompletionTargets(target, completionTargets)) { + return false; + } for (uint32_t projection = 0; projection < kReduceGradProjectionCount; ++projection) { const uint64_t chunkBytes = ChunkBytes(projection); const uint64_t chunks = ChunkCount(projection); @@ -534,7 +572,8 @@ class ReduceGradKernel { const uint64_t completedBytes = MinU64( rowBytes_[projection] - completedOffset, chunkBytes); if (!CompleteChunk(target, slot, projection, completedChunk, - sourceRow + completedOffset, completedBytes)) { + sourceRow + completedOffset, completedBytes, + completionTargets)) { return false; } } @@ -543,7 +582,8 @@ class ReduceGradKernel { const bool ok = transports_[projection] == kReduceGradTransportPeer ? SendPeerChunk(target, slot, projection, chunk, sourceRow + offset, bytes) : - PostUdmaChunk(target, slot, projection, chunk, sourceRow + offset, bytes); + PostUdmaChunk(target, slot, projection, chunk, + sourceRow + offset, bytes, completionTargets); if (!ok) { return false; } @@ -553,7 +593,7 @@ class ReduceGradKernel { const uint64_t offset = chunk * chunkBytes; const uint64_t bytes = MinU64(rowBytes_[projection] - offset, chunkBytes); if (!CompleteChunk(target, slot, projection, chunk, - sourceRow + offset, bytes)) { + sourceRow + offset, bytes, completionTargets)) { return false; } } diff --git a/tests/moonep/python/test_ffi_unittest.py b/tests/moonep/python/test_ffi_unittest.py index cd93687e..7131faa7 100644 --- a/tests/moonep/python/test_ffi_unittest.py +++ b/tests/moonep/python/test_ffi_unittest.py @@ -2,6 +2,7 @@ import ctypes import os +import struct import sys import unittest from pathlib import Path @@ -31,6 +32,7 @@ TileXRMoonEPTensorV1, ) from tilexr_moonep.runtime import TileXRMoonEPRuntime, _resolve_library +from tilexr_moonep.torch_api import _format_dispatch_completion_flags class FakeFunction: @@ -387,6 +389,18 @@ def tensor(shape, dtype): class FfiAbiTests(unittest.TestCase): + def test_dispatch_completion_flag_matrix_format(self): + flags = bytearray(512 * 2 * 8) + struct.pack_into(" None: cases = load_cases(root / "tools" / "moonep" / "cases" / "correctness.json") for case in cases: + if case.case_id in PRODUCTION_SCALE_REPRO_CASES: + continue world_size = RUNNER_WORLD_SIZES[case.case_id] dimensions_by_rank = [ MoonEPDimensions( @@ -312,9 +324,9 @@ def test_runner_cases_support_one_based_numeric_selection() -> None: for number, case in enumerate(cases, start=1): assert select_cases(cases, str(number)) == [case] - assert select_cases(cases, "1,14") == [cases[0], cases[13]] - for invalid in ("0", "15"): - with pytest.raises(ValueError, match=r"case number must be in \[1, 14\]"): + assert select_cases(cases, "1,16") == [cases[0], cases[15]] + for invalid in ("0", "17"): + with pytest.raises(ValueError, match=r"case number must be in \[1, 16\]"): select_cases(cases, invalid) @@ -322,7 +334,7 @@ def test_case_14_is_the_two_node_one_rank_per_device_case() -> None: root = Path(__file__).resolve().parents[3] cases = load_cases(root / "tools" / "moonep" / "cases" / "correctness.json") - assert len(cases) == 14 + assert len(cases) == 16 case = cases[13] assert case.case_id == "planning-16rank-16card-single-route" assert ( @@ -336,3 +348,69 @@ def test_case_14_is_the_two_node_one_rank_per_device_case() -> None: case.routing_pattern, case.route_distribution, ) == (8, 1, 16, 8, 4, 1, 1, "balanced", "rank_shifted_uniform") + + +def test_case_15_matches_the_4k_ep8_grouped_urma_dispatch_repro() -> None: + root = Path(__file__).resolve().parents[3] + cases = load_cases(root / "tools" / "moonep" / "cases" / "correctness.json") + + case = cases[14] + assert case.case_id == "dispatch-8rank-4k-ep8-grouped-urma" + assert ( + case.tokens_per_rank, + case.topk, + case.expert_count, + case.hidden_size, + case.intermediate_size, + case.prefetch_slots, + case.token_padding, + case.routing_pattern, + case.route_distribution, + case.warmup, + case.iterations, + ) == ( + 4096, + 8, + 32, + 7168, + 2048, + 4, + 1, + "unique_destinations", + "rank_shifted_uniform", + 0, + 8, + ) + + +def test_case_16_matches_the_4k_ep8_grouped_urma_plan_reuse_repro() -> None: + root = Path(__file__).resolve().parents[3] + cases = load_cases(root / "tools" / "moonep" / "cases" / "correctness.json") + + case = cases[15] + assert case.case_id == "flow-8rank-4k-ep8-grouped-urma-plan-reuse" + assert ( + case.tokens_per_rank, + case.topk, + case.expert_count, + case.hidden_size, + case.intermediate_size, + case.prefetch_slots, + case.token_padding, + case.routing_pattern, + case.route_distribution, + case.warmup, + case.iterations, + ) == ( + 4096, + 8, + 32, + 7168, + 2048, + 4, + 1, + "unique_destinations", + "rank_shifted_uniform", + 0, + 8, + ) diff --git a/tests/moonep/python/test_moonep_modes.py b/tests/moonep/python/test_moonep_modes.py index 02646b05..ace4a188 100644 --- a/tests/moonep/python/test_moonep_modes.py +++ b/tests/moonep/python/test_moonep_modes.py @@ -409,6 +409,27 @@ def test_launcher_forwards_mode_candidate_and_reference_overrides(tmp_path) -> N assert command[command.index("--tensor-preview-elements") + 1] == "5" +def test_single_node_launcher_builds_hidden_dispatch_hot_loop_command(tmp_path) -> None: + args = build_launcher_parser().parse_args( + [ + "--cases", + str(tmp_path / "cases.json"), + "--output-dir", + str(tmp_path / "out"), + "--benchmark-kind", + "dispatch_hot_loop", + "--dispatch-modes", + "hidden", + ] + ) + + command = _process_command(args) + assert "tools.moonep.dispatch_hot_loop" in command + assert "tools.moonep.benchmark" not in command + assert command[command.index("--dispatch-modes") + 1 :] == ["hidden"] + assert "--mode" not in command + + def test_candidate_factory_is_protocol_checked(monkeypatch) -> None: module = ModuleType("test_candidate_backend") diff --git a/tests/moonep/python/test_moonep_performance_report.py b/tests/moonep/python/test_moonep_performance_report.py index aeb702f8..da5ba376 100644 --- a/tests/moonep/python/test_moonep_performance_report.py +++ b/tests/moonep/python/test_moonep_performance_report.py @@ -12,6 +12,7 @@ FLOW_STAGE_ORDER, aggregate_distributed_artifacts, aggregate_rank_artifacts, + format_dispatch_performance, format_stage_performance, main as report_main, write_json, @@ -19,6 +20,51 @@ ) +def test_dispatch_hot_loop_summary_prints_dispatch_metrics( + tmp_path: Path, capsys +) -> None: + summary = { + "benchmark_kind": "dispatch_hot_loop", + "mode": "benchmark", + "case": { + "case_id": "dispatch-8rank-4k-ep8-grouped-urma", + "tokens_per_rank": 4096, + "topk": 8, + "expert_count": 32, + "hidden_size": 7168, + "intermediate_size": 2048, + "prefetch_slots": 4, + "token_padding": 1, + "dtype": "bfloat16", + "seed": 1234, + "routing_pattern": "unique_destinations", + "route_distribution": "rank_shifted_uniform", + "correctness": True, + "warmup": 0, + "iterations": 8, + }, + "dispatch_modes": ["hidden"], + "metrics_us": { + "hidden_host": {"mean": 40.0, "p50": 35.0, "p95": 80.0}, + "hidden_kernel": {"mean": 2200.0, "p50": 2150.0, "p95": 2300.0}, + }, + "tokens_per_second_by_mode": { + "hidden": {"mean": 14_000_000.0}, + }, + } + + table = format_dispatch_performance(summary) + assert "MoonEP Dispatch-only performance" in table + assert "hidden" in table + assert "2200.000" in table + assert "14000000.000" in table + + summary_path = tmp_path / "summary.json" + write_json(summary_path, summary) + assert report_main(["--summary", str(summary_path)]) == 0 + assert "MoonEP Dispatch-only performance" in capsys.readouterr().out + + RAW_TIMINGS = { "planning": 1.0, "dispatch_forward": 1.0, diff --git a/tests/moonep/python/test_run_moonep_script.py b/tests/moonep/python/test_run_moonep_script.py index 7195bcd5..e0870a35 100644 --- a/tests/moonep/python/test_run_moonep_script.py +++ b/tests/moonep/python/test_run_moonep_script.py @@ -149,7 +149,7 @@ def test_usage_describes_rank_size_and_rank_per_device_for_every_case() -> None: "\n\nEnvironment:", 1 )[0] descriptions = [line for line in usage_cases.splitlines() if line.strip()] - assert len(descriptions) == 14 + assert len(descriptions) == 16 assert all("rank_size=" in line for line in descriptions) assert all("rank_per_dev=" in line for line in descriptions) @@ -190,6 +190,69 @@ def test_usage_describes_eight_and_sixteen_rank_cases() -> None: SCRIPT, re.MULTILINE, ) + assert re.search( + r"^ 15\s+dispatch-8rank-4k-ep8-grouped-urma\s+8-rank Dispatch-only .+rank_size=8, rank_per_dev=1.+S=4096, K=8, E=32, H=7168, Hf=2048, B=4, P=1\)$", + SCRIPT, + re.MULTILINE, + ) + assert re.search( + r"^ 16\s+flow-8rank-4k-ep8-grouped-urma-plan-reuse\s+8-rank full flow with Combine V2 then saved-plan backward Dispatch .+rank_size=8, rank_per_dev=1.+S=4096, K=8, E=32, H=7168, Hf=2048, B=4, P=1\)$", + SCRIPT, + re.MULTILINE, + ) + + +def test_case_15_selects_hidden_only_grouped_urma_dispatch_hot_loop() -> None: + assert 'dispatch_repro_case_id="dispatch-8rank-4k-ep8-grouped-urma"' in SCRIPT + assert 'benchmark_kind="dispatch_hot_loop"' in SCRIPT + assert 'dispatch_modes=("hidden")' in SCRIPT + assert "unset TILEXR_MOONEP_DISPATCH_TRANSPORT" in SCRIPT + assert 'export TILEXR_MOONEP_DISPATCH_PEER_MODE="group"' in SCRIPT + assert 'export TILEXR_MOONEP_DISPATCH_GROUP_WIDTH="16"' in SCRIPT + assert '--benchmark-kind "${benchmark_kind}"' in SCRIPT + assert 'launcher_args+=("--dispatch-modes" "${dispatch_modes[@]}")' in SCRIPT + + +def test_case_16_selects_grouped_urma_full_flow_with_plan_reuse() -> None: + assert ( + 'plan_reuse_repro_case_id="flow-8rank-4k-ep8-grouped-urma-plan-reuse"' + in SCRIPT + ) + assert 'elif [[ "${case_id}" == "${plan_reuse_repro_case_id}" ]]' in SCRIPT + assert 'warmup="${warmup:-0}"' in SCRIPT + assert 'iterations="${iterations:-8}"' in SCRIPT + assert 'benchmark_kind="flow"' in SCRIPT + assert "unset TILEXR_MOONEP_DISPATCH_TRANSPORT" in SCRIPT + assert 'export TILEXR_MOONEP_DISPATCH_PEER_MODE="group"' in SCRIPT + assert 'export TILEXR_MOONEP_DISPATCH_GROUP_WIDTH="16"' in SCRIPT + + +def test_full_flow_reuses_the_plan_after_forward_combine() -> None: + benchmark = (ROOT / "tools" / "moonep" / "benchmark.py").read_text( + encoding="utf-8" + ) + flow = benchmark.split("def execute_iteration(", 1)[1].split( + "\ndef _tensor_row_bytes", 1 + )[0] + combine = flow.index('"combine_forward"') + backward = flow.index('"dispatch_backward"') + reuse = flow.index('buffer.dispatch(inputs["grad_output"], plan=plan)') + assert combine < backward < reuse + + runtime = ( + ROOT + / "integrations" + / "moonep_torch" + / "tilexr_moonep" + / "runtime.py" + ).read_text(encoding="utf-8") + assert "self._combine_v2_lib.TileXRMoonEpCombineStageV2(" in runtime + + +def test_single_node_launcher_accepts_dispatch_hot_loop_controls() -> None: + assert '"--benchmark-kind"' in LAUNCHER + assert 'choices=("flow", "dispatch_hot_loop")' in LAUNCHER + assert '"--dispatch-modes"' in LAUNCHER def test_script_exposes_managed_multinode_three_mode_launch() -> None: diff --git a/tests/moonep/unit/test_tilexr_moonep_dispatch_schedule.cpp b/tests/moonep/unit/test_tilexr_moonep_dispatch_schedule.cpp index 3ecc2e0d..6c77e2e2 100644 --- a/tests/moonep/unit/test_tilexr_moonep_dispatch_schedule.cpp +++ b/tests/moonep/unit/test_tilexr_moonep_dispatch_schedule.cpp @@ -248,11 +248,23 @@ void TestWqeBatchBoundaries() head += 7U * 128U; CHECK_EQ(TileXRMoonEp::DispatchWqeBatchCount(UINT64_MAX, head, 8192U), 121U); CHECK_EQ((1017U * sizeof(int16_t)) % 32U, 18U); - CHECK_TRUE(TileXRMoonEp::DispatchPeerWqesFitSq(8192U, 16384U)); - CHECK_TRUE(!TileXRMoonEp::DispatchPeerWqesFitSq(32768U, 16384U)); - CHECK_TRUE(!TileXRMoonEp::DispatchPeerWqesFitSq(UINT64_MAX, 16384U)); - CHECK_TRUE(TileXRMoonEp::DispatchBatchNeedsCompletion(true)); - CHECK_TRUE(!TileXRMoonEp::DispatchBatchNeedsCompletion(false)); + CHECK_TRUE(TileXRMoonEp::DispatchPeerWqesStreamable(8192U, 16384U)); + CHECK_TRUE(TileXRMoonEp::DispatchPeerWqesStreamable(32768U, 16384U)); + CHECK_TRUE(!TileXRMoonEp::DispatchPeerWqesStreamable(UINT64_MAX, 16384U)); + CHECK_TRUE(!TileXRMoonEp::DispatchPeerWqesStreamable(8192U, + TileXRMoonEp::kDispatchSqPollReserve + + TileXRMoonEp::kDispatchWqeBatchCapacity - 1U)); + CHECK_TRUE(!TileXRMoonEp::DispatchGroupedBatchNeedsCompletion(0U)); + CHECK_TRUE(TileXRMoonEp::DispatchGroupedBatchNeedsCompletion(1U)); + CHECK_TRUE(TileXRMoonEp::DispatchGroupedBatchNeedsCompletion( + TileXRMoonEp::kDispatchWqeBatchCapacity)); + + CHECK_EQ(TileXRMoonEp::DispatchRouteTileCount(32768U, 0U, 1024U), 1024U); + CHECK_EQ(TileXRMoonEp::DispatchRouteTileCount(32768U, 31744U, 1024U), + 1024U); + CHECK_EQ(TileXRMoonEp::DispatchRouteTileCount(32768U, 32768U, 1024U), 0U); + CHECK_EQ(TileXRMoonEp::DispatchRouteTileCount(1000U, 0U, 1024U), 1000U); + CHECK_EQ(TileXRMoonEp::DispatchRouteTileCount(1000U, 0U, 0U), 0U); } diff --git a/tests/moonep/unit/test_tilexr_moonep_host.cpp b/tests/moonep/unit/test_tilexr_moonep_host.cpp index 717da505..45a2df48 100644 --- a/tests/moonep/unit/test_tilexr_moonep_host.cpp +++ b/tests/moonep/unit/test_tilexr_moonep_host.cpp @@ -26,11 +26,13 @@ int dispatchUrmaV2Calls = 0; int dispatchWorkspaceQueryCalls = 0; int combineCalls = 0; int prefetchCalls = 0; +int memsetCalls = 0; uint64_t queryWorkspaceBytes = 512; int64_t queryNvS = 12; TileXR::CommArgs commArgs {}; const TileXRMoonEpDispatchArgsV1 *seenDispatch = nullptr; const TileXRMoonEpDispatchArgsV1 *seenDispatchUrma = nullptr; +uint64_t seenDispatchUrmaFlags = UINT64_MAX; const TileXRMoonEpCombineArgsV1 *seenCombine = nullptr; aclrtStream seenDispatchStream = nullptr; aclrtStream seenCombineStream = nullptr; @@ -74,6 +76,7 @@ void Reset() queryCalls = plannerCalls = dispatchCalls = dispatchUrmaCalls = dispatchUrmaV2Calls = dispatchWorkspaceQueryCalls = combineCalls = prefetchCalls = 0; + memsetCalls = 0; queryWorkspaceBytes = 512; queryNvS = 12; commArgs = TileXR::CommArgs {}; @@ -83,6 +86,7 @@ void Reset() plannerCall = PlannerCall {}; seenDispatch = nullptr; seenDispatchUrma = nullptr; + seenDispatchUrmaFlags = UINT64_MAX; seenCombine = nullptr; seenDispatchStream = nullptr; seenCombineStream = nullptr; @@ -285,20 +289,31 @@ void TestStageDelegation() Check(dispatchCalls == 1 && dispatchUrmaCalls == 1 && seenDispatchUrma == &dispatch, "URMA dispatch selection mismatch"); + TileXRMoonEpPlanV1 plan = ValidPlan(); + dispatch.plan = &plan; + dispatch.flags = TILEXR_MOONEP_FLAG_RESET_STATUS; + CheckStatus("URMA dispatch reset routing", TileXRMoonEpDispatchV1(&dispatch, stream), + dispatchUrmaReturn); + Check(memsetCalls == 0 && dispatchUrmaCalls == 2 && + seenDispatchUrma == &dispatch && + seenDispatchUrmaFlags == TILEXR_MOONEP_FLAG_RESET_STATUS, + "URMA dispatch reset flag must be delegated unchanged"); + dispatch.flags = TILEXR_MOONEP_FLAG_NONE; + 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, + Check(dispatchUrmaCalls == 2, "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 && + Check(dispatchUrmaCalls == 2 && dispatchUrmaV2Calls == 1 && seenDispatchUrma == &dispatchV2, "V2 dispatch must select the URMA implementation"); @@ -309,7 +324,6 @@ void TestStageDelegation() Check(combineCalls == 1 && seenCombine == &combine && seenCombineStream == stream, "combine delegation mismatch"); - TileXRMoonEpPlanV1 plan = ValidPlan(); TileXRMoonEpPrefetchWeightArgsV1 prefetch {}; prefetch.structSize = sizeof(prefetch); prefetch.abiVersion = TILEXR_MOONEP_ABI_VERSION_V1; @@ -370,6 +384,7 @@ extern "C" int TileXRMoonEpPlannerV3(const int32_t *topk, const int32_t *tpe, extern "C" aclError aclrtMemsetAsync( void *, size_t, int32_t, size_t, aclrtStream) { + ++memsetCalls; return ACL_SUCCESS; } @@ -379,11 +394,6 @@ extern "C" aclError aclrtMemcpyAsync( return ACL_SUCCESS; } -extern "C" aclError aclrtSynchronizeStream(aclrtStream) -{ - return ACL_SUCCESS; -} - namespace TileXRMoonEp { int TileXRMoonEpRunDispatchV1( const TileXRMoonEpDispatchArgsV1 *args, aclrtStream stream) @@ -408,6 +418,7 @@ int TileXRMoonEpRunDispatchUrmaV1( { ++dispatchUrmaCalls; seenDispatchUrma = args; + seenDispatchUrmaFlags = args->flags; return dispatchUrmaReturn; } diff --git a/tests/moonep/unit/test_tilexr_moonep_kernel_sources.cpp b/tests/moonep/unit/test_tilexr_moonep_kernel_sources.cpp index e4874560..b9626ddb 100644 --- a/tests/moonep/unit/test_tilexr_moonep_kernel_sources.cpp +++ b/tests/moonep/unit/test_tilexr_moonep_kernel_sources.cpp @@ -217,7 +217,9 @@ int main() Contains("prefetch kernel", prefetchKernel, "localExpert"); Excludes("prefetch kernel", prefetchKernel, "e_ + slot"); Contains("prefetch kernel", prefetchKernel, "UDMAGetNbiOnQp"); - Contains("prefetch kernel", prefetchKernel, "UDMAQuietStatusOnQp"); + Contains("prefetch kernel", prefetchKernel, "UDMAQuietStatusOnQpUntil"); + Contains("prefetch kernel", prefetchKernel, "completionQueueIds"); + Contains("prefetch kernel", prefetchKernel, "++completionTargets[completionQueue]"); Contains("prefetch kernel", prefetchKernel, "DataCacheCleanAndInvalid"); Excludes("prefetch kernel", prefetchKernel, "<<<"); Excludes("prefetch launch", prefetchLaunch, "launch_tilexr_moonep_prefetch_weight_kernel"); @@ -246,7 +248,11 @@ int main() Contains("reduce kernel", reduceKernel, "DataAsFlagCheckBatchCleared"); Contains("reduce kernel", reduceKernel, "AscendC::Add"); Contains("reduce kernel", reduceKernel, "UDMAPutRegisteredSignalNbiOnQp"); - Contains("reduce kernel", reduceKernel, "UDMAQuietStatusOnQp"); + Contains("reduce kernel", reduceKernel, "UDMAQuietStatusOnQpUntil"); + Contains("reduce kernel", reduceKernel, "InitializeUdmaCompletionTargets"); + Contains("reduce kernel", reduceKernel, "++completionTargets[qpIdx]"); + Excludes("reduce kernel", reduceKernel, + "UDMAQuietStatusOnQp(args_, static_cast(target), qpIdx)"); Excludes("reduce kernel", reduceKernel, "tilexr_moonep_reduce_grad_status_kernel"); Excludes("reduce kernel", reduceKernel, "kReduceGradDeviceStatusSuccess"); 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 9b84859e..5a259022 100644 --- a/tests/moonep/unit/test_tilexr_moonep_prefetch_weight_host.cpp +++ b/tests/moonep/unit/test_tilexr_moonep_prefetch_weight_host.cpp @@ -146,6 +146,24 @@ void TestLaunch() Reset(); launchRet = -77; Status("prefetch launch failure", TileXRMoonEp::TileXRMoonEpRunPrefetchWeightV1(&args, stream), -77); } + +void TestLargeQpCount() +{ + Reset(); + qpNum = 32; + auto plan = Plan(); + auto gate = Weight(0x100000, 4, 8); + auto up = Weight(0x101000, 4, 16); + auto down = Weight(0x102000, 8, 8); + auto args = Args(&plan, &gate, &up, &down); + auto stream = reinterpret_cast(uintptr_t {0x7000}); + Status("prefetch 32 QPs", + TileXRMoonEp::TileXRMoonEpRunPrefetchWeightV1(&args, stream), + TILEXR_MOONEP_SUCCESS); + Check(launchCalls == 1 && seenContext.layout.qpNum == 32 && + seenContext.layout.blockDim == 4, + "prefetch must cap workers independently of available QPs"); +} } extern "C" int TileXRGetCommArgsHost(TileXRCommPtr, TileXR::CommArgs *&out) { out = hostRet == 0 ? &commArgs : nullptr; return hostRet; } @@ -155,4 +173,4 @@ extern "C" int TileXRGetUDMARegistryHost(TileXRCommPtr, const TileXR::TileXRUDMA extern "C" int TileXRUDMAGetQpCount(TileXRCommPtr, uint32_t *out) { if (qpRet == 0) *out = qpNum; return qpRet; } namespace TileXRMoonEp { int TileXRMoonEpLaunchPrefetchWeightKernel(const PrefetchWeightParams &p, const PrefetchWeightLaunchContext &c) { ++launchCalls; seenParams = p; seenContext = c; return launchRet; } } -int main() { TestLaunch(); return failures == 0 ? EXIT_SUCCESS : EXIT_FAILURE; } +int main() { TestLaunch(); TestLargeQpCount(); return failures == 0 ? EXIT_SUCCESS : EXIT_FAILURE; } diff --git a/tests/moonep/unit/test_tilexr_moonep_reduce_grad_host.cpp b/tests/moonep/unit/test_tilexr_moonep_reduce_grad_host.cpp index 7ca99210..290b54ac 100644 --- a/tests/moonep/unit/test_tilexr_moonep_reduce_grad_host.cpp +++ b/tests/moonep/unit/test_tilexr_moonep_reduce_grad_host.cpp @@ -289,6 +289,29 @@ void TestMixedUdma() "invalid UDMA QP counts must be validated after every query"); } +void TestMaximumUdmaQpCount() +{ + Reset(); + g_commArgs.localRankSize = 1; + TileXRMoonEpPlanV1 plan = Plan(); + const uint64_t largeRow = + TILEXR_MOONEP_REDUCE_GRAD_UDMA_THRESHOLD_BYTES / sizeof(float) + 1; + TileXRMoonEpTensorV1 gate = Gradient( + reinterpret_cast(UINTPTR_C(0x400000)), largeRow); + TileXRMoonEpTensorV1 up = Gradient( + reinterpret_cast(UINTPTR_C(0x500000)), largeRow); + TileXRMoonEpTensorV1 down = Gradient( + reinterpret_cast(UINTPTR_C(0x600000)), largeRow); + + g_qpCount = 32; + const auto info = Query(&plan, &gate, &up, &down, TileXR::TILEXR_SUCCESS); + Check(info.workspaceBytes > 0, + "32-QP ReduceGrad query must return a UDMA workspace"); + + g_qpCount = 33; + (void)Query(&plan, &gate, &up, &down, TileXR::TILEXR_ERROR_NOT_INITIALIZED); +} + void TestSingleRankLargeRowsDoNotRequireUdma() { Reset(); @@ -487,6 +510,7 @@ int main() TestPeerOnly(); TestCompactPrefetchSlots(); TestMixedUdma(); + TestMaximumUdmaQpCount(); TestSingleRankLargeRowsDoNotRequireUdma(); TestLargeRankWorkspaceQuery(); TestLaunchFailureDrainsEnqueuedStatusReset(); diff --git a/tests/moonep/unit/test_tilexr_moonep_sources.cpp b/tests/moonep/unit/test_tilexr_moonep_sources.cpp index fcb875c6..1e36a5bf 100644 --- a/tests/moonep/unit/test_tilexr_moonep_sources.cpp +++ b/tests/moonep/unit/test_tilexr_moonep_sources.cpp @@ -131,6 +131,14 @@ int main() Contains("host", host, "aclrtMemcpyAsync"); Contains("host", host, "ACL_MEMCPY_DEVICE_TO_DEVICE"); Excludes("host", host, "aclrtSynchronizeStream"); + Contains("URMA dispatch Host", dispatchUrmaHost, + "TILEXR_MOONEP_FLAG_RESET_STATUS"); + CheckOrdered("URMA dispatch status reset", dispatchUrmaHost, { + "for (uint32_t lhs = 0; lhs < rangeCount; ++lhs)", + "DispatchUrmaLaunchParams params {}", + "aclrtMemsetAsync(args->plan->status", + "TileXRMoonEpLaunchDispatchUrmaKernel(params)", + }); Contains("CMake", cmake, "add_library(tilexr-moonep SHARED"); Contains("CMake", cmake, "SOVERSION 1"); @@ -175,6 +183,62 @@ int main() "ClearDispatchZeroFillRanges"); Contains("URMA dispatch kernel", dispatchUrmaKernel, "udmaIssueQp0Buf"); Contains("URMA dispatch kernel", dispatchUrmaKernel, "udmaIssueQp1Buf"); + Contains("URMA dispatch kernel", dispatchUrmaKernel, + "uint32_t completionCount;"); + Contains("URMA dispatch kernel", dispatchUrmaKernel, + "state.outstanding = state.head - state.tail;"); + Contains("URMA dispatch kernel", dispatchUrmaKernel, + "state.completionCount += 1U;"); + Contains("URMA dispatch kernel", dispatchUrmaKernel, + "DispatchGroupedBatchNeedsCompletion(batchCount)"); + Contains("URMA dispatch kernel", dispatchUrmaKernel, + "routeId - context->routePlanStart"); + Contains("URMA dispatch kernel", dispatchUrmaKernel, + "DispatchRouteTileCount("); + Contains("URMA dispatch kernel", dispatchUrmaKernel, + "useVectorSlotSelect && !groupedPeerMode"); + Contains("URMA dispatch kernel", dispatchUrmaKernel, + "LoadDispatchPlanStatus(planStatus, relayLocal)"); + Contains("URMA dispatch kernel", dispatchUrmaKernel, + "SyncFunc()"); + CheckOrdered("URMA dispatch signal publication", dispatchUrmaKernel, { + "PublishDispatchSignalSource(", + "signalLocal.SetValue(", + "SyncFunc()", + "AscendC::DataCopyPad(signalGlobal, signalLocal, copyOut)", + "SyncFunc()", + "TileXR::UDMACleanCacheLines(", + }); + CheckOrdered("URMA dispatch completion flag cache invalidation", + dispatchUrmaKernel, { + "LoadCompletionFlag(", + "TileXR::UDMACleanCacheLines(", + "AscendC::PipeBarrier()", + "return flag[0]", + }); + Excludes("URMA dispatch kernel", dispatchUrmaKernel, + "AscendC::DataCopyPad(relayLocal, flagGlobal, copyIn, padIn)"); + CheckOrdered("URMA dispatch credit cache invalidation", + dispatchUrmaKernel, { + "LoadDispatchCredit(", + "reinterpret_cast<__gm__ uint8_t *>(credit)", + "SyncFunc()", + "AscendC::DataCopy(creditLocal, creditGlobal,", + "SyncFunc()", + }); + Excludes("URMA dispatch kernel", dispatchUrmaKernel, + "signalSource[qpIdx] = expectedFlag"); + Excludes("URMA dispatch kernel", dispatchUrmaKernel, + "AppendDispatchSignalWqe"); + Excludes("URMA dispatch kernel", dispatchUrmaKernel, + "InvalidateDispatchInboundCache"); + Excludes("URMA dispatch kernel", dispatchUrmaKernel, "state.wqeCount"); + Excludes("URMA dispatch kernel", dispatchUrmaKernel, + "state.completionCount += batchCount"); + Contains("URMA dispatch Host", dispatchUrmaHost, + "DispatchPeerWqesStreamable"); + Excludes("URMA dispatch Host", dispatchUrmaHost, + "DispatchPeerWqesFitSq"); Excludes("URMA dispatch kernel", dispatchUrmaKernel, "<<<"); Contains("combine CMake", combineCmake, "add_library(tilexr-moonep-combine SHARED"); Contains("combine CMake", combineCmake, "TILEXR_MOONEP_COMBINE_KERNEL_EMBED_CPP"); diff --git a/tests/udma/unit/test_tilexr_udma_device_api.cpp b/tests/udma/unit/test_tilexr_udma_device_api.cpp index e35a5123..cf55d90d 100644 --- a/tests/udma/unit/test_tilexr_udma_device_api.cpp +++ b/tests/udma/unit/test_tilexr_udma_device_api.cpp @@ -360,6 +360,34 @@ void TestGetAndLegacyQp0Wrappers() "legacy PUT uses named ordered-completion flag"); } +void TestExplicitCompletionTargetDoesNotRereadPublishedCount() +{ + Fixture fixture; + constexpr uint32_t qp = 1U; + auto* local = fixture.localRegion.data(); + auto scratch = fixture.Scratch(); + uint32_t completionTarget = fixture.wqeCount[qp]; + + for (uint32_t read = 0U; read < 3U; ++read) { + Check(UDMAGetNbiOnQp(&fixture.args, scratch, 1, qp, + local + read * 8U, read * 8U, 8U) == + TILEXR_UDMA_STATUS_SUCCESS, + "tracked QP-aware GET succeeds"); + ++completionTarget; + Check(completionTarget == read + 1U, + "caller tracks the cumulative completion target locally"); + fixture.Cqe(qp, read)->owner = 1U; + fixture.Cqe(qp, read)->entryIdx = read; + } + + fixture.wqeCount[qp] = 0U; + Check(UDMAQuietStatusOnQpUntil(&fixture.args, 1, qp, completionTarget) == + TILEXR_UDMA_STATUS_SUCCESS, + "explicit completion target does not reread a stale published WQE count"); + Check(fixture.sqTail[qp] == 3U && fixture.cqTail[qp] == 3U, + "explicit completion target reclaims all three READs"); +} + void TestWriteNotifyWrapsWithinSqRing() { Fixture fixture; @@ -441,6 +469,33 @@ void TestWriteNotifyCompletionUsesFinalBbIndex() "WRITE_WITH_NOTIFY completion reclaims both basic blocks"); } +void TestWriteNotifyExplicitCompletionTargetDoesNotRereadPublishedCount() +{ + Fixture fixture; + constexpr uint32_t qp = 0U; + auto scratch = fixture.Scratch(); + UDMASignalParams signal = {}; + signal.sigAddr = reinterpret_cast( + fixture.remoteRegion.data() + 128U); + signal.signal = 9U; + uint32_t completionTarget = fixture.wqeCount[qp]; + + Check(UDMAWriteNotify(&fixture.args, scratch, fixture.remoteRegion.data(), + fixture.localRegion.data(), 1U, qp, 8U, &signal) == + TILEXR_UDMA_STATUS_SUCCESS, + "tracked WRITE_WITH_NOTIFY succeeds"); + ++completionTarget; + fixture.Cqe(qp, 0U)->owner = 1U; + fixture.Cqe(qp, 0U)->entryIdx = 1U; + fixture.wqeCount[qp] = 0U; + + Check(UDMAQuietStatusOnQpUntil(&fixture.args, 1, qp, completionTarget) == + TILEXR_UDMA_STATUS_SUCCESS, + "explicit WRITE_WITH_NOTIFY target does not reread a stale WQE count"); + Check(fixture.sqTail[qp] == 2U && fixture.cqTail[qp] == 1U, + "explicit WRITE_WITH_NOTIFY target reclaims both SQ BBs"); +} + void TestLegacyQuietWithoutRegistry() { Fixture fixture; @@ -491,6 +546,23 @@ void TestCqReclaimAcrossSqAndCqWrap() "quiet rings CQ doorbell with absolute consumer"); } +void TestCqeEntryIndexNormalizesSqCycle() +{ + Fixture fixture; + constexpr uint32_t qp = 0U; + fixture.sqTail[qp] = TILEXR_UDMA_SQ_BB_COUNT; + fixture.sqHead[qp] = TILEXR_UDMA_SQ_BB_COUNT + 1U; + fixture.wqeCount[qp] = 1U; + fixture.Cqe(qp, 0U)->owner = 1U; + fixture.Cqe(qp, 0U)->entryIdx = TILEXR_UDMA_SQ_BB_COUNT; + + Check(UDMAQuietStatusOnQp(&fixture.args, 1, qp) == + TILEXR_UDMA_STATUS_SUCCESS, + "CQE entry index carrying an SQ cycle bit is accepted"); + Check(fixture.sqTail[qp] == TILEXR_UDMA_SQ_BB_COUNT + 1U, + "cycle-bearing CQE reclaims the final BB in the current SQ cycle"); +} + void TestInvalidCqeDoesNotReclaim() { Fixture fixture; @@ -521,10 +593,13 @@ int main() TestDeferredPutAndFlushAreQpSpecific(); TestImmediatePutReclaimsCompletedFullSq(); TestGetAndLegacyQp0Wrappers(); + TestExplicitCompletionTargetDoesNotRereadPublishedCount(); TestWriteNotifyWrapsWithinSqRing(); TestWriteNotifyCompletionUsesFinalBbIndex(); + TestWriteNotifyExplicitCompletionTargetDoesNotRereadPublishedCount(); TestLegacyQuietWithoutRegistry(); TestCqReclaimAcrossSqAndCqWrap(); + TestCqeEntryIndexNormalizesSqCycle(); TestInvalidCqeDoesNotReclaim(); if (gFailures != 0) { diff --git a/tests/udma/unit/test_tilexr_udma_source_guard.cpp b/tests/udma/unit/test_tilexr_udma_source_guard.cpp index 66e3c0da..dab708de 100644 --- a/tests/udma/unit/test_tilexr_udma_source_guard.cpp +++ b/tests/udma/unit/test_tilexr_udma_source_guard.cpp @@ -340,6 +340,18 @@ void TestUDMAMultiQpHostTransportContract() "BuildUDMAInfoImage(reinterpret_cast(registration.infoDev), qpCount_"); CheckContains(transportPath, transport, "return IsAvailable() ? qpCount_ : 0U;"); CheckContains(transportPath, transport, "ret = AgreeInitStatus(localStatus);"); + CheckContains(transportPath, transport, "ret = AgreeRaOwnership();"); + CheckContains(transportPath, transport, "int TileXRUDMATransport::AgreeRaOwnership() const"); + CheckContains(transportPath, transport, + "options_.exchange->AllGather(&local, 1, allAttached.data())"); + CheckContains(transportPath, transport, "TILEXR_HCCP_RA_ALREADY_INITIALIZED = 328002"); + CheckContains(transportPath, transport, + "std::getenv(\"TILEXR_UDMA_ATTACH_EXISTING_RA\")"); + CheckContains(transportPath, transport, + "ret == TILEXR_HCCP_RA_ALREADY_INITIALIZED && AttachExistingRaEnabled()"); + CheckContains(transportPath, transport, + "TileXR UDMA attaching to existing RA initialization"); + CheckContains(transportPath, transport, "if (raInitialized_)"); CheckContains(transportPath, transport, "int TileXRUDMATransport::AgreeEidCount() const"); CheckContains(transportPath, transport, "std::array allEidCounts"); diff --git a/tools/moonep/cases/correctness.json b/tools/moonep/cases/correctness.json index bbd639d9..974baf87 100644 --- a/tools/moonep/cases/correctness.json +++ b/tools/moonep/cases/correctness.json @@ -208,5 +208,35 @@ "warmup": 0, "iters": 1, "correctness": true + }, + { + "id": "dispatch-8rank-4k-ep8-grouped-urma", + "S": 4096, + "K": 8, + "E": 32, + "H": 7168, + "Hf": 2048, + "B": 4, + "P": 1, + "routing": "unique_destinations", + "route_distribution": "rank_shifted_uniform", + "warmup": 0, + "iters": 8, + "correctness": true + }, + { + "id": "flow-8rank-4k-ep8-grouped-urma-plan-reuse", + "S": 4096, + "K": 8, + "E": 32, + "H": 7168, + "Hf": 2048, + "B": 4, + "P": 1, + "routing": "unique_destinations", + "route_distribution": "rank_shifted_uniform", + "warmup": 0, + "iters": 8, + "correctness": true } ] diff --git a/tools/moonep/launcher.py b/tools/moonep/launcher.py index dd388367..01868fed 100644 --- a/tools/moonep/launcher.py +++ b/tools/moonep/launcher.py @@ -140,6 +140,15 @@ def build_parser() -> argparse.ArgumentParser: parser.add_argument( "--mode", choices=("benchmark", "reference", "correctness"), default="benchmark" ) + parser.add_argument( + "--benchmark-kind", choices=("flow", "dispatch_hot_loop"), default="flow" + ) + parser.add_argument( + "--dispatch-modes", + nargs="+", + choices=("hidden", "weight", "pair"), + default=("hidden", "weight", "pair"), + ) parser.add_argument("--candidate-backend", default=None, metavar="MODULE:FACTORY") parser.add_argument( "--dump-stage-tensors", @@ -198,6 +207,8 @@ def main(argv: list[str] | None = None) -> int: ) if args.mode == "benchmark" and args.dump_stage_tensors: raise ValueError("--dump-stage-tensors is only valid in reference/correctness mode") + if args.mode != "benchmark" and args.benchmark_kind != "flow": + raise ValueError("--benchmark-kind dispatch_hot_loop requires --mode benchmark") topology = resolve_topology( physical_device_count=args.physical_device_count, ranks_per_device=args.ranks_per_device, diff --git a/tools/moonep/mindspeed/README.md b/tools/moonep/mindspeed/README.md new file mode 100644 index 00000000..900cc66b --- /dev/null +++ b/tools/moonep/mindspeed/README.md @@ -0,0 +1,33 @@ +# MindSpeed validation tools + +This directory contains reusable tools developed while validating the TileXR MoonEP +backend with MindSpeed on a single eight-device Ascend 950 node. + +- `tilexr_mindspeed_adapter.py`: MindSpeed backend adapter that keeps communication + buffers owned by TileXR. +- `mindspeed_external_comm_owner.patch`: MindSpeed patch that lets an external backend + own the communication runtime and buffer. +- `preflight_adapter.sh`: validates the adapter and the patched MindSpeed ownership + contract before an NPU model run. +- `grouped_urma_dispatch_oracle.py`: exact repeated grouped-URMA Dispatch/Combine + oracle at the production `S=4096`, `K=8`, EP8 route shape. +- `run_case15_32.sh`: 32-iteration case 15 runner. +- `run_grouped_oracle.sh`: eight-rank grouped-URMA oracle runner. + +The runners use the current repository by default and write results below +`run/moonep/mindspeed`. Override paths with `TILEXR_HOME`, +`TILEXR_INSTALL_PREFIX`, `TILEXR_CANN_ENV`, `TILEXR_CONDA_SH`, +`TILEXR_CONDA_ENV`, and `TILEXR_MOONEP_NATIVE_ENV`. + +Before using the adapter, apply the MindSpeed ownership patch and run its +preflight: + +```bash +git -C "${MINDSPEED_HOME}" apply \ + "${TILEXR_HOME}/tools/moonep/mindspeed/mindspeed_external_comm_owner.patch" +MINDSPEED_HOME="${MINDSPEED_HOME}" \ + bash "${TILEXR_HOME}/tools/moonep/mindspeed/preflight_adapter.sh" +``` + +The preflight installs `tilexr_mindspeed_adapter.py` into the checkout named by +`MINDSPEED_HOME`; use a disposable or task-owned MindSpeed checkout. diff --git a/tools/moonep/mindspeed/grouped_urma_dispatch_oracle.py b/tools/moonep/mindspeed/grouped_urma_dispatch_oracle.py new file mode 100644 index 00000000..5817122b --- /dev/null +++ b/tools/moonep/mindspeed/grouped_urma_dispatch_oracle.py @@ -0,0 +1,473 @@ +"""Exact 8-rank oracle for repeated grouped URMA MoonEP Dispatch.""" + +from __future__ import annotations + +import ctypes +import os +import struct +import sys +import time +from pathlib import Path + +import torch +import torch.distributed as dist + + +ROOT = Path(os.environ.get("TILEXR_ORACLE_SOURCE_ROOT", Path(__file__).resolve().parents[3])) +INTEGRATION = ROOT / "integrations" / "moonep_torch" +for path in (ROOT, INTEGRATION): + value = str(path) + if value not in sys.path: + sys.path.insert(0, value) + +from tilexr_moonep import Buffer, ProjectionBuffers + + +RANK_SIZE = 8 +TOKENS_PER_RANK = 4096 +HIDDEN_SIZE = int(os.environ.get("TILEXR_ORACLE_HIDDEN_SIZE", "32")) +TOPK = 8 +EXPERT_COUNT = 32 +ITERATIONS = int(os.environ.get("TILEXR_ORACLE_ITERATIONS", "16")) +ROUTE_MODE = os.environ.get("TILEXR_ORACLE_ROUTE_MODE", "balanced") +WITH_ROUTE_WEIGHTS = os.environ.get("TILEXR_ORACLE_WITH_ROUTE_WEIGHTS", "1") == "1" +WITH_COMBINE = os.environ.get("TILEXR_ORACLE_WITH_COMBINE", "1") == "1" +MAGIC_ADVANCE_AFTER_FIRST = int( + os.environ.get("TILEXR_ORACLE_MAGIC_ADVANCE_AFTER_FIRST", "0") +) +SWITCH_REGISTRATION = ( + os.environ.get("TILEXR_ORACLE_SWITCH_REGISTRATION", "0") == "1" +) +EXTRA_PLAN_COUNT = int(os.environ.get("TILEXR_ORACLE_EXTRA_PLAN_COUNT", "0")) +WITH_PREFETCH = os.environ.get("TILEXR_ORACLE_WITH_PREFETCH", "0") == "1" +PROJECTION_SIZE = int(os.environ.get("TILEXR_ORACLE_PROJECTION_SIZE", "256")) +WAIT_SECONDS = int(os.environ.get("TILEXR_ORACLE_WAIT_SECONDS", "120")) + + +def make_inputs(rank: int, device: str): + token = torch.arange(TOKENS_PER_RANK, dtype=torch.int64).view(-1, 1) + column = torch.arange(HIDDEN_SIZE, dtype=torch.int64).view(1, -1) + route = torch.arange(TOPK, dtype=torch.int64).view(1, -1) + hidden = ((token * 17 + column * 3 + rank * 29) % 251).to(torch.bfloat16) + weights = ( + ((token * TOPK + route * 11 + rank * 37) % 1024).to(torch.float32) + / 1024.0 + ) + if ROUTE_MODE == "balanced": + local_expert = (token + rank) % (EXPERT_COUNT // RANK_SIZE) + topk = (route * (EXPERT_COUNT // RANK_SIZE) + local_expert).to(torch.int32) + elif ROUTE_MODE == "model_skew": + generator = torch.Generator(device="cpu") + generator.manual_seed(20260811 + rank) + logits = torch.randn( + (TOKENS_PER_RANK, EXPERT_COUNT), generator=generator + ) + logits += torch.linspace(0.6, -0.6, EXPERT_COUNT).reshape(1, -1) + topk = torch.topk(logits, TOPK, dim=1).indices.to(torch.int32) + else: + raise ValueError(f"unknown TILEXR_ORACLE_ROUTE_MODE={ROUTE_MODE}") + tpe = torch.bincount(topk.flatten().to(torch.int64), minlength=EXPERT_COUNT) + return ( + hidden.contiguous().to(device), + weights.contiguous().to(device), + topk.contiguous().to(device), + tpe.to(torch.int32).contiguous().to(device), + ) + + +def wait_for(paths: list[Path], label: str) -> None: + deadline = time.monotonic() + WAIT_SECONDS + while not all(path.exists() for path in paths): + if time.monotonic() >= deadline: + missing = [str(path) for path in paths if not path.exists()] + raise TimeoutError(f"timed out waiting for {label}: {missing}") + time.sleep(0.05) + + +def publish_tensor(path: Path, value: torch.Tensor, rank: int) -> None: + temporary = path.with_name(f".{path.name}.rank{rank}.tmp") + torch.save(value, temporary) + os.replace(temporary, path) + + +def file_barrier(run_dir: Path, label: str, rank: int) -> None: + marker = run_dir / f"{label}.rank{rank}.ready" + marker.touch() + wait_for( + [run_dir / f"{label}.rank{peer}.ready" for peer in range(RANK_SIZE)], + label, + ) + + +def build_expected(run_dir: Path, target_rank: int, nv_s: int): + expected_hidden = torch.empty( + (nv_s, HIDDEN_SIZE), dtype=torch.bfloat16, device="cpu" + ) + expected_weights = torch.empty((nv_s,), dtype=torch.float32, device="cpu") + filled = torch.zeros((nv_s,), dtype=torch.bool, device="cpu") + route_ids = torch.arange(TOKENS_PER_RANK * TOPK, dtype=torch.int64) + + for source_rank in range(RANK_SIZE): + dst = torch.load(run_dir / f"dst.rank{source_rank}.pt", map_location="cpu") + dst = dst.to(torch.int64) + if tuple(dst.shape) != (TOKENS_PER_RANK * TOPK,): + raise AssertionError(f"rank {source_rank} dst shape is {tuple(dst.shape)}") + if bool(((dst < 0) | (dst >= RANK_SIZE * nv_s)).any().item()): + bad = int(torch.nonzero((dst < 0) | (dst >= RANK_SIZE * nv_s))[0].item()) + raise AssertionError( + f"rank {source_rank} dst[{bad}]={int(dst[bad])} is out of range" + ) + + selected = torch.div(dst, nv_s, rounding_mode="floor") == target_rank + slots = torch.remainder(dst[selected], nv_s) + selected_routes = route_ids[selected] + if bool(filled[slots].any().item()): + slot = int(slots[torch.nonzero(filled[slots])[0]].item()) + raise AssertionError(f"destination slot {slot} has multiple writers") + + source_hidden, source_weights, _, _ = make_inputs(source_rank, "cpu") + expected_hidden[slots] = source_hidden[selected_routes // TOPK] + expected_weights[slots] = source_weights.flatten()[selected_routes] + filled[slots] = True + + if not bool(filled.all().item()): + slot = int(torch.nonzero(~filled)[0].item()) + raise AssertionError(f"destination slot {slot} has no writer") + return expected_hidden, expected_weights + + +def compare_exact( + rank: int, + iteration: int, + name: str, + actual: torch.Tensor, + expected: torch.Tensor, +) -> None: + actual_cpu = actual.detach().cpu() + if torch.equal(actual_cpu, expected): + print( + f"[grouped-urma-oracle rank={rank}] iteration={iteration} " + f"{name}=exact", + flush=True, + ) + return + mismatch = actual_cpu != expected + first = tuple(int(value) for value in torch.nonzero(mismatch, as_tuple=False)[0]) + raise AssertionError( + f"rank {rank} iteration {iteration} {name} mismatch at {first}: " + f"actual={actual_cpu[first].item()} expected={expected[first].item()}" + ) + + +def advance_magic(buffer: Buffer, rank: int) -> None: + runtime = buffer._native_buffer.runtime + function = runtime._comm_lib.TileXRCommNextMagic + function.argtypes = [ctypes.c_void_p, ctypes.POINTER(ctypes.c_int64)] + function.restype = ctypes.c_int + values = [] + for _ in range(MAGIC_ADVANCE_AFTER_FIRST): + value = ctypes.c_int64() + ret = int(function(ctypes.c_void_p(runtime.comm_ptr), ctypes.byref(value))) + if ret != 0: + raise RuntimeError(f"TileXRCommNextMagic failed with ret={ret}") + values.append(int(value.value)) + if values: + print( + f"[grouped-urma-oracle rank={rank}] advanced_magic={values}", + flush=True, + ) + + +def activate_dummy_registration(buffer: Buffer, workspace: torch.Tensor) -> None: + native = buffer._native_buffer + native._synchronize_device() + native.runtime._activate_udma_region( + int(workspace.data_ptr()), + int(workspace.numel()), + "oracle_dummy", + f"workspace_bytes={int(workspace.numel())}", + ) + + +def dump_hidden_scratch( + buffer: Buffer, + plan, + rank: int, + iteration: int, + expected_hidden: torch.Tensor, +) -> None: + context = buffer._native_buffer.context + raw = context._dispatch_workspace_owner + aligned_offset = int(context._dispatch_workspace_ptr) - int(raw.data_ptr()) + source_bytes = TOKENS_PER_RANK * HIDDEN_SIZE * 2 + scratch_offset = (source_bytes + 63) // 64 * 64 + scratch_slot_bytes = TOKENS_PER_RANK * TOPK * HIDDEN_SIZE * 2 + raw_cpu = raw.detach().cpu() + print( + f"[grouped-urma-oracle rank={rank}] iteration={iteration} " + f"workspace_bytes={context._dispatch_workspace_bytes} " + f"aligned_offset={aligned_offset} hidden_scratch_offset={scratch_offset} " + f"hidden_scratch_slot_bytes={scratch_slot_bytes}", + flush=True, + ) + for scratch_index in range(2): + byte_offset = aligned_offset + scratch_offset + scratch_index * scratch_slot_bytes + scratch = ( + raw_cpu.narrow(0, byte_offset, scratch_slot_bytes) + .view(torch.bfloat16) + .reshape(TOKENS_PER_RANK * TOPK, HIDDEN_SIZE) + ) + mismatch_count = int((scratch != expected_hidden).sum().item()) + print( + f"[grouped-urma-oracle rank={rank}] iteration={iteration} " + f"scratch={scratch_index} mismatch_elements={mismatch_count} " + f"row0={scratch[0, :8].tolist()}", + flush=True, + ) + raw_bytes = raw_cpu.numpy().tobytes() + status_marker = struct.pack("= 0 and status_offset + 64 <= len(raw_bytes): + marker, version, record_bytes, status, payload_mode, magic = struct.unpack_from( + " int: + if os.environ.get("TILEXR_MOONEP_DISPATCH_TRANSPORT"): + raise RuntimeError("Dispatch transport override must be unset for this URMA oracle") + if os.environ.get("TILEXR_MOONEP_DISPATCH_PEER_MODE") != "group": + raise RuntimeError("TILEXR_MOONEP_DISPATCH_PEER_MODE must be group") + if os.environ.get("TILEXR_MOONEP_DISPATCH_GROUP_WIDTH") != "16": + raise RuntimeError("TILEXR_MOONEP_DISPATCH_GROUP_WIDTH must be 16") + + import torch_npu # noqa: F401 + + rank = int(os.environ["RANK"]) + local_rank = int(os.environ["LOCAL_RANK"]) + torch.npu.set_device(local_rank) + dist.init_process_group(backend="hccl") + if dist.get_world_size() != RANK_SIZE: + raise RuntimeError(f"expected {RANK_SIZE} ranks, got {dist.get_world_size()}") + + run_dir = Path(os.environ["TILEXR_ORACLE_RUN_DIR"]) + buffer = None + dummy_workspace = None + dummy_allocation = None + projections = None + try: + hidden, weights, topk, tpe = make_inputs(rank, "npu") + buffer = Buffer( + TOKENS_PER_RANK, + HIDDEN_SIZE, + TOPK, + EXPERT_COUNT, + RANK_SIZE, + B=EXPERT_COUNT // RANK_SIZE, + num_sms=32, + token_padding=1, + ) + if SWITCH_REGISTRATION: + dummy_workspace, dummy_allocation = ( + buffer._native_buffer._aligned_workspace(2 * 1024 * 1024, 2 * 1024 * 1024) + ) + dispatch_weights = weights if WITH_ROUTE_WEIGHTS else None + hidden_out, weight_out, _, plan = buffer.dispatch( + hidden, dispatch_weights, topk, tpe + ) + torch.npu.synchronize() + nv_s = int(hidden_out.shape[0]) + if nv_s != TOKENS_PER_RANK * TOPK: + raise AssertionError(f"expected NvS=32768, got {nv_s}") + + publish_tensor(run_dir / f"dst.rank{rank}.pt", plan.dst.detach().cpu(), rank) + file_barrier(run_dir, "plans", rank) + expected_hidden, expected_weights = build_expected(run_dir, rank, nv_s) + + expected_combine_hidden = (hidden.detach().cpu().to(torch.float32) * TOPK).to( + torch.bfloat16 + ) + expected_combine_weights = weights.detach().cpu() + + if WITH_PREFETCH: + local_experts = EXPERT_COUNT // RANK_SIZE + local_weights = tuple( + torch.zeros( + (local_experts, HIDDEN_SIZE, PROJECTION_SIZE), + dtype=torch.bfloat16, + device="npu", + ) + for _ in range(3) + ) + projections = ProjectionBuffers.from_local_weights( + buffer._context, *local_weights, torch_module=torch + ) + buffer._native_buffer.register_projection_buffers(projections) + buffer._native_buffer.prefetch_weight( + plan._require_native(), projections, async_finish=False + ) + + compare_exact(rank, 0, "hidden", hidden_out, expected_hidden) + if WITH_ROUTE_WEIGHTS: + compare_exact(rank, 0, "route_weights", weight_out, expected_weights) + if WITH_COMBINE: + combined_hidden, combined_weights, _ = buffer.combine( + plan=plan, + hidden_nvsh=hidden_out, + route_weights_nvs=weight_out, + ) + torch.npu.synchronize() + compare_exact( + rank, 0, "combined_hidden", combined_hidden, expected_combine_hidden + ) + if WITH_ROUTE_WEIGHTS: + compare_exact( + rank, + 0, + "combined_route_weights", + combined_weights, + expected_combine_weights, + ) + for extra_plan in range(EXTRA_PLAN_COUNT): + extra_topk = torch.remainder( + topk + (extra_plan + 1) * 3, EXPERT_COUNT + ).to(torch.int32) + extra_tpe = torch.bincount( + extra_topk.flatten().to(torch.int64), minlength=EXPERT_COUNT + ).to(torch.int32) + hidden_out, weight_out, _, plan = buffer.dispatch( + hidden, + dispatch_weights, + extra_topk.contiguous(), + extra_tpe.contiguous(), + ) + torch.npu.synchronize() + if WITH_PREFETCH: + buffer._native_buffer.prefetch_weight( + plan._require_native(), projections, async_finish=False + ) + if WITH_COMBINE: + combined_hidden, _, _ = buffer.combine( + plan=plan, + hidden_nvsh=hidden_out, + route_weights_nvs=weight_out, + ) + torch.npu.synchronize() + compare_exact( + rank, + -(extra_plan + 1), + "extra_plan_combined_hidden", + combined_hidden, + expected_combine_hidden, + ) + if SWITCH_REGISTRATION: + activate_dummy_registration(buffer, dummy_workspace) + + if EXTRA_PLAN_COUNT: + publish_tensor(run_dir / f"dst.rank{rank}.pt", plan.dst.detach().cpu(), rank) + file_barrier(run_dir, "extra_plans", rank) + expected_hidden, expected_weights = build_expected(run_dir, rank, nv_s) + + advance_magic(buffer, rank) + if SWITCH_REGISTRATION: + activate_dummy_registration(buffer, dummy_workspace) + file_barrier(run_dir, "iteration0", rank) + + reuse_status_mode = os.environ.get( + "TILEXR_ORACLE_SYNC_REUSE_STATUS", "0" + ) + if reuse_status_mode != "0": + native_plan = plan._require_native() + native_plan.status.zero_() + if reuse_status_mode == "sync": + torch.npu.synchronize() + buffer._native_buffer._pending_statuses.pop(id(native_plan), None) + + for iteration in range(1, ITERATIONS): + hidden_out, weight_out, _, returned_plan = buffer.dispatch( + hidden, dispatch_weights, plan=plan + ) + if returned_plan is not plan: + raise AssertionError("reused Dispatch did not return the same plan") + torch.npu.synchronize() + buffer._native_buffer._check_plan_status( + returned_plan._require_native() + ) + if WITH_PREFETCH: + buffer._native_buffer.prefetch_weight( + plan._require_native(), projections, async_finish=False + ) + try: + compare_exact(rank, iteration, "hidden", hidden_out, expected_hidden) + except AssertionError: + dump_hidden_scratch(buffer, plan, rank, iteration, expected_hidden) + raise + if WITH_ROUTE_WEIGHTS: + compare_exact( + rank, iteration, "route_weights", weight_out, expected_weights + ) + if WITH_COMBINE: + combined_hidden, combined_weights, _ = buffer.combine( + plan=plan, + hidden_nvsh=hidden_out, + route_weights_nvs=weight_out, + ) + torch.npu.synchronize() + compare_exact( + rank, + iteration, + "combined_hidden", + combined_hidden, + expected_combine_hidden, + ) + if WITH_ROUTE_WEIGHTS: + compare_exact( + rank, + iteration, + "combined_route_weights", + combined_weights, + expected_combine_weights, + ) + if SWITCH_REGISTRATION: + activate_dummy_registration(buffer, dummy_workspace) + file_barrier(run_dir, f"iteration{iteration}", rank) + + if rank == 0: + print( + f"[grouped-urma-oracle] PASS iterations={ITERATIONS} " + f"S=4096 K=8 H={HIDDEN_SIZE} rank_size=8 " + f"peer_mode=group group_width=16 weights={int(WITH_ROUTE_WEIGHTS)} " + f"combine={int(WITH_COMBINE)} " + f"route_mode={ROUTE_MODE} " + f"switch_registration={int(SWITCH_REGISTRATION)} " + f"extra_plans={EXTRA_PLAN_COUNT} " + f"prefetch={int(WITH_PREFETCH)} " + f"magic_advance={MAGIC_ADVANCE_AFTER_FIRST}", + flush=True, + ) + return 0 + finally: + if buffer is not None and not buffer.destroyed: + buffer.destroy() + if dist.is_initialized(): + dist.destroy_process_group() + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/tools/moonep/mindspeed/mindspeed_external_comm_owner.patch b/tools/moonep/mindspeed/mindspeed_external_comm_owner.patch new file mode 100644 index 00000000..0a7f6f5b --- /dev/null +++ b/tools/moonep/mindspeed/mindspeed_external_comm_owner.patch @@ -0,0 +1,113 @@ +diff --git a/mindspeed/core/transformer/moe/moonep_model_arena.py b/mindspeed/core/transformer/moe/moonep_model_arena.py +--- a/mindspeed/core/transformer/moe/moonep_model_arena.py ++++ b/mindspeed/core/transformer/moe/moonep_model_arena.py +@@ -133,11 +133,17 @@ class MoonEPModelArenaLayout: + ).total_bytes + +- def owned_requests(self) -> tuple["ModelArenaAllocationRequest", ...]: +- return ( ++ def owned_requests( ++ self, *, include_communication_buffer: bool = True ++ ) -> tuple["ModelArenaAllocationRequest", ...]: ++ if not isinstance(include_communication_buffer, bool): ++ raise TypeError("include_communication_buffer must be a boolean") ++ communication = ( + ModelArenaAllocationRequest( + "communication.boundary_pool", "communication", + (self.communication_boundary_pool_bytes,), torch.uint8, + ), ++ ) if include_communication_buffer else () ++ return communication + ( + ModelArenaAllocationRequest( + "global.weight_slot.fc1", "weight_slot", + (self.B, self.H, 2 * self.F), torch.bfloat16, +@@ -322,4 +328,5 @@ class MoonEPModelArenaOrchestrator: + native_library_path: str | None = None, + include_legacy_share_handles: bool = True, + pass_mem_handles: bool = True, ++ install_communication_runtime: bool = True, + ) -> None: +@@ -345,10 +352,13 @@ class MoonEPModelArenaOrchestrator: + raise TypeError("include_legacy_share_handles must be a boolean") + if not isinstance(pass_mem_handles, bool): + raise TypeError("pass_mem_handles must be a boolean") ++ if not isinstance(install_communication_runtime, bool): ++ raise TypeError("install_communication_runtime must be a boolean") + if include_legacy_share_handles and not pass_mem_handles: + raise ValueError( + "legacy share handles require caller-provided memory handles" + ) + self._include_legacy_share_handles = include_legacy_share_handles + self._pass_mem_handles = pass_mem_handles ++ self._install_communication_runtime = install_communication_runtime + self._arena = MoonEPSymmetricArena( +@@ -399,5 +409,7 @@ class MoonEPModelArenaOrchestrator: + ) + + owned = [] +- for request in self.layout.owned_requests(): ++ for request in self.layout.owned_requests( ++ include_communication_buffer=self._install_communication_runtime ++ ): + binding = self._allocation_provider(request) +@@ -432,11 +444,16 @@ class MoonEPModelArenaOrchestrator: + manifest_gatherer=self._manifest_gatherer + ) + self._state = ModelArenaState.FROZEN +- communication_ids = ("communication.boundary_pool",) +- self._runtime_installer( +- self._arena, +- communication_buffer_ids=communication_ids, +- data_op_engine_type=self._data_op_engine_type, +- library_path=self._native_library_path, ++ communication_ids = tuple( ++ binding.request.buffer_id ++ for binding in owned ++ if binding.request.role == "communication" + ) ++ if self._install_communication_runtime: ++ self._runtime_installer( ++ self._arena, ++ communication_buffer_ids=communication_ids, ++ data_op_engine_type=self._data_op_engine_type, ++ library_path=self._native_library_path, ++ ) + self._installed = InstalledModelArena( +diff --git a/mindspeed/core/transformer/moe/moonep_unified_arena_wiring.py b/mindspeed/core/transformer/moe/moonep_unified_arena_wiring.py +--- a/mindspeed/core/transformer/moe/moonep_unified_arena_wiring.py ++++ b/mindspeed/core/transformer/moe/moonep_unified_arena_wiring.py +@@ -51,4 +51,9 @@ from .moonep_model_arena import ( + from .moonep_registration import get_current_expert_registry + + ++EXTERNAL_COMMUNICATION_OWNER_ATTRIBUTE = ( ++ "__mindspeed_external_communication_owner__" ++) ++ ++ + class UnifiedArenaWiringError(ModelArenaContractError): +@@ -438,5 +443,18 @@ class MoonEPUnifiedArenaFeatureState: + if not collected: + raise UnifiedArenaWiringError("get_model produced no expert external buffers") +- layout = layout_from_args(self._training_args(), deps.world_size) ++ training_args = self._training_args() ++ layout = layout_from_args(training_args, deps.world_size) ++ from .moonep_backend import resolve_native_backend_factory ++ ++ backend_factory = resolve_native_backend_factory( ++ getattr(training_args, "moonep_backend_factory", None) ++ ) ++ external_communication_owner = getattr( ++ backend_factory, EXTERNAL_COMMUNICATION_OWNER_ATTRIBUTE, False ++ ) ++ if not isinstance(external_communication_owner, bool): ++ raise UnifiedArenaWiringError( ++ f"{EXTERNAL_COMMUNICATION_OWNER_ATTRIBUTE} must be a boolean" ++ ) + + def allocate(request: ModelArenaAllocationRequest): +@@ -480,3 +498,4 @@ class MoonEPUnifiedArenaFeatureState: + # only share handle. 950 uses Fabric V2; 910B uses legacy V1. + pass_mem_handles=True, ++ install_communication_runtime=not external_communication_owner, + ) diff --git a/tools/moonep/mindspeed/preflight_adapter.sh b/tools/moonep/mindspeed/preflight_adapter.sh new file mode 100644 index 00000000..b6be6a59 --- /dev/null +++ b/tools/moonep/mindspeed/preflight_adapter.sh @@ -0,0 +1,80 @@ +#!/usr/bin/env bash + +set -euo pipefail + +script_dir=$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd) +tilexr_home=${TILEXR_HOME:-$(cd "${script_dir}/../../.." && pwd)} +mindspeed_home=${MINDSPEED_HOME:?MINDSPEED_HOME must point to the MindSpeed checkout} +adapter_dir=${mindspeed_home}/mindspeed/core/transformer/moe +adapter=${adapter_dir}/tilexr_mindspeed_adapter.py +install_prefix=${TILEXR_INSTALL_PREFIX:-${tilexr_home}/install} + +if [[ ! -d "${adapter_dir}" ]]; then + printf 'MindSpeed MoE package directory not found: %s\n' "${adapter_dir}" >&2 + exit 1 +fi + +install -m 0644 "${script_dir}/tilexr_mindspeed_adapter.py" "${adapter}" +if grep -Eq 'from moonep import Buffer|aclshmem' "${adapter}"; then + echo "forbidden SHMEM Buffer dependency in TileXR adapter" >&2 + exit 1 +fi + +export TILEXR_INSTALL_PREFIX=${install_prefix} +export PYTHONPATH="${mindspeed_home}:${tilexr_home}/integrations/moonep_torch:${tilexr_home}${PYTHONPATH:+:${PYTHONPATH}}" + +python - <<'PY' +import inspect +from types import SimpleNamespace + +import torch + +from mindspeed.core.transformer.moe import tilexr_mindspeed_adapter as adapter +from mindspeed.core.transformer.moe.moonep_model_arena import MoonEPModelArenaLayout + +assert adapter.MindSpeedTileXRBuffer.__mro__[1].__module__ == "tilexr_moonep.compat" +assert "hidden_buffer" in inspect.signature(adapter.MindSpeedTileXRBuffer.dispatch).parameters +assert "hidden_buffer" in inspect.signature(adapter.MindSpeedTileXRBuffer.combine).parameters +assert adapter.MOONEP_NATIVE_NPU_CAPABILITY in ( + adapter.create_tilexr_moonep_backend.__mindspeed_capabilities__ +) +assert ( + adapter.create_tilexr_moonep_backend.__mindspeed_external_communication_owner__ + is True +) + +layout = MoonEPModelArenaLayout(S=8, H=64, K=2, E=8, R=2, B=4, F=128) +default_requests = layout.owned_requests() +tilexr_requests = layout.owned_requests(include_communication_buffer=False) +assert any(request.role == "communication" for request in default_requests) +assert all(request.role != "communication" for request in tilexr_requests) +assert [request for request in default_requests if request.role != "communication"] == list( + tilexr_requests +) +assert adapter.upstream_moonep.Buffer._tilexr_mindspeed_guard is True + +owner = object() +plan = object() +probe = SimpleNamespace(_plan_owner_token=owner, _dispatch_generation=0) +hidden = torch.empty((4, 8)) +route = torch.empty((4,)) +hidden_alias, route_alias = adapter.MindSpeedTileXRBuffer._bind_zero_copy_views( + probe, plan, hidden, route +) +assert hidden_alias is not hidden and hidden_alias.data_ptr() == hidden.data_ptr() +assert route_alias is not route and route_alias.data_ptr() == route.data_ptr() +assert hidden_alias._moonep_buffer_owner is owner +assert hidden_alias._moonep_dispatch_plan is plan +assert hidden_alias._moonep_dispatch_generation == 1 +assert probe._dispatch_generation == 1 + +try: + adapter.upstream_moonep.Buffer.__init__(object()) +except RuntimeError as exc: + assert "forbids construction" in str(exc) +else: + raise AssertionError("upstream SHMEM Buffer constructor guard did not fire") + +print(f"adapter={adapter.__file__}") +print("tilexr_mindspeed_adapter_preflight=PASS") +PY diff --git a/tools/moonep/mindspeed/run_case15_32.sh b/tools/moonep/mindspeed/run_case15_32.sh new file mode 100644 index 00000000..afeddc80 --- /dev/null +++ b/tools/moonep/mindspeed/run_case15_32.sh @@ -0,0 +1,40 @@ +#!/usr/bin/env bash + +set -euo pipefail + +script_dir=$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd) +tilexr_home=${TILEXR_HOME:-$(cd "${script_dir}/../../.." && pwd)} +install_prefix=${TILEXR_INSTALL_PREFIX:-${tilexr_home}/install} +cann_env=${TILEXR_CANN_ENV:-/home/pkg/b131/cann/set_env.sh} +conda_sh=${TILEXR_CONDA_SH:-/home/miniconda3/etc/profile.d/conda.sh} +conda_env=${TILEXR_CONDA_ENV:-ai_moe_test} +output_root=${TILEXR_MOONEP_OUTPUT_ROOT:-${tilexr_home}/run/moonep/mindspeed} + +export LD_LIBRARY_PATH= +export PYTHONPATH= +unset ASCEND_HOME ASCEND_HOME_PATH ASCEND_AICPU_PATH ASCEND_OPP_PATH TOOLCHAIN_HOME +source "${cann_env}" +source "${conda_sh}" +conda activate "${conda_env}" +if [[ -n "${TILEXR_MOONEP_NATIVE_ENV:-}" ]]; then + source "${TILEXR_MOONEP_NATIVE_ENV}" +fi + +export TILEXR_INSTALL_PREFIX=${install_prefix} +export TILEXR_MOONEP_CONDA_ENV=${conda_env} +export TILEXR_UDMA_QP_ROUTE_SPEC=${TILEXR_UDMA_QP_ROUTE_SPEC:-port_count:6,port_count:2} +export TILEXR_UDMA_ATTACH_EXISTING_RA=${TILEXR_UDMA_ATTACH_EXISTING_RA:-1} +unset TILEXR_MOONEP_DISPATCH_TRANSPORT +export TILEXR_MOONEP_DISPATCH_PEER_MODE=group +export TILEXR_MOONEP_DISPATCH_GROUP_WIDTH=16 +export TILEXR_MOONEP_DUMP_DFX_ON_ERROR=${TILEXR_MOONEP_DUMP_DFX_ON_ERROR:-1} +export TILEXR_MOONEP_OUTPUT_DIR="${output_root}/case15_iter32_$(date +%Y%m%d-%H%M%S)" + +cd "${tilexr_home}" +exec bash scripts/run_moonep.sh \ + --mode benchmark \ + --rank-size 8 \ + --case-id 15 \ + --visible-devices 0,1,2,3,4,5,6,7 \ + --warmup 0 \ + --iterations 32 diff --git a/tools/moonep/mindspeed/run_grouped_oracle.sh b/tools/moonep/mindspeed/run_grouped_oracle.sh new file mode 100644 index 00000000..1d7660fc --- /dev/null +++ b/tools/moonep/mindspeed/run_grouped_oracle.sh @@ -0,0 +1,58 @@ +#!/usr/bin/env bash + +set -euo pipefail + +script_dir=$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd) +tilexr_home=${TILEXR_HOME:-$(cd "${script_dir}/../../.." && pwd)} +install_prefix=${TILEXR_INSTALL_PREFIX:-${tilexr_home}/install} +cann_env=${TILEXR_CANN_ENV:-/home/pkg/b131/cann/set_env.sh} +conda_sh=${TILEXR_CONDA_SH:-/home/miniconda3/etc/profile.d/conda.sh} +conda_env=${TILEXR_CONDA_ENV:-ai_moe_test} +output_root=${TILEXR_MOONEP_OUTPUT_ROOT:-${tilexr_home}/run/moonep/mindspeed} +run_dir=${output_root}/grouped_oracle_$(date +%Y%m%d-%H%M%S) + +export LD_LIBRARY_PATH= +export PYTHONPATH= +unset ASCEND_HOME ASCEND_HOME_PATH ASCEND_AICPU_PATH ASCEND_OPP_PATH TOOLCHAIN_HOME +source "${cann_env}" +source "${conda_sh}" +conda activate "${conda_env}" +if [[ -n "${TILEXR_MOONEP_NATIVE_ENV:-}" ]]; then + source "${TILEXR_MOONEP_NATIVE_ENV}" +fi + +export PYTHONPATH="${tilexr_home}/integrations/moonep_torch:${tilexr_home}${PYTHONPATH:+:${PYTHONPATH}}" +export LD_LIBRARY_PATH="${install_prefix}/lib64${LD_LIBRARY_PATH:+:${LD_LIBRARY_PATH}}" +export TILEXR_INSTALL_PREFIX=${install_prefix} +export TILEXR_UDMA_QP_ROUTE_SPEC=${TILEXR_UDMA_QP_ROUTE_SPEC:-port_count:6,port_count:2} +export TILEXR_UDMA_ATTACH_EXISTING_RA=${TILEXR_UDMA_ATTACH_EXISTING_RA:-1} +unset TILEXR_MOONEP_DISPATCH_TRANSPORT +export TILEXR_MOONEP_DISPATCH_PEER_MODE=group +export TILEXR_MOONEP_DISPATCH_GROUP_WIDTH=16 +export TILEXR_MOONEP_DUMP_DFX_ON_ERROR=${TILEXR_MOONEP_DUMP_DFX_ON_ERROR:-1} +export TILEXR_ORACLE_SOURCE_ROOT=${tilexr_home} +export TILEXR_ORACLE_RUN_DIR=${run_dir} +export TILEXR_ORACLE_ITERATIONS=${TILEXR_ORACLE_ITERATIONS:-20} +export TILEXR_ORACLE_HIDDEN_SIZE=${TILEXR_ORACLE_HIDDEN_SIZE:-7168} +export TILEXR_ORACLE_ROUTE_MODE=${TILEXR_ORACLE_ROUTE_MODE:-model_skew} +export TILEXR_ORACLE_WITH_ROUTE_WEIGHTS=${TILEXR_ORACLE_WITH_ROUTE_WEIGHTS:-0} +export TILEXR_ORACLE_WITH_COMBINE=${TILEXR_ORACLE_WITH_COMBINE:-1} +export TILEXR_ORACLE_SWITCH_REGISTRATION=${TILEXR_ORACLE_SWITCH_REGISTRATION:-1} +export TILEXR_ORACLE_EXTRA_PLAN_COUNT=${TILEXR_ORACLE_EXTRA_PLAN_COUNT:-5} +export TILEXR_ORACLE_WITH_PREFETCH=${TILEXR_ORACLE_WITH_PREFETCH:-1} +export TILEXR_ORACLE_PROJECTION_SIZE=${TILEXR_ORACLE_PROJECTION_SIZE:-256} +export HCCL_CONNECT_TIMEOUT=${HCCL_CONNECT_TIMEOUT:-120} +export HCCL_EXEC_TIMEOUT=${HCCL_EXEC_TIMEOUT:-120} +export CUDA_DEVICE_MAX_CONNECTIONS=${CUDA_DEVICE_MAX_CONNECTIONS:-1} +export TASK_QUEUE_ENABLE=${TASK_QUEUE_ENABLE:-2} +export PYTORCH_NPU_ALLOC_CONF=${PYTORCH_NPU_ALLOC_CONF:-expandable_segments:True} +export STREAMS_PER_DEVICE=${STREAMS_PER_DEVICE:-32} + +mkdir -p "${run_dir}" +cd "${tilexr_home}" +exec timeout --signal=TERM --kill-after=20s 120s python -m torch.distributed.launch \ + --nproc_per_node 8 \ + --master_addr 127.0.0.1 \ + --master_port "${TILEXR_ORACLE_MASTER_PORT:-45678}" \ + --use-env \ + "${script_dir}/grouped_urma_dispatch_oracle.py" diff --git a/tools/moonep/mindspeed/tilexr_mindspeed_adapter.py b/tools/moonep/mindspeed/tilexr_mindspeed_adapter.py new file mode 100644 index 00000000..f17373b0 --- /dev/null +++ b/tools/moonep/mindspeed/tilexr_mindspeed_adapter.py @@ -0,0 +1,449 @@ +from __future__ import annotations + +import os + +import moonep as upstream_moonep +from tilexr_moonep import Buffer as TileXRBuffer +from tilexr_moonep import ProjectionBuffers + +from .moonep_backend import ( + MOONEP_NATIVE_NPU_CAPABILITY, + MoonEPBufferFlexBackend, +) + + +def _reject_upstream_buffer_init(*args, **kwargs): + del args, kwargs + raise RuntimeError( + "TileXR MindSpeed mode forbids construction of the SHMEM communication Buffer" + ) + + +if not getattr(upstream_moonep.Buffer, "_tilexr_mindspeed_guard", False): + upstream_moonep.Buffer.__init__ = _reject_upstream_buffer_init + upstream_moonep.Buffer._tilexr_mindspeed_guard = True + + +class MindSpeedTileXRBuffer(TileXRBuffer): + """MR3832 Buffer surface backed by TileXR communication resources.""" + + def __init__(self, *args, token_buffer_count=1, **kwargs): + if isinstance(token_buffer_count, bool) or int(token_buffer_count) <= 0: + raise ValueError("token_buffer_count must be a positive integer") + super().__init__(*args, **kwargs) + self.token_buffer_count = int(token_buffer_count) + device = f"npu:{self._context.device_index}" + self.token_buffers = tuple( + self._torch.empty( + (self._context.nv_s, self.H), + dtype=self._torch.bfloat16, + device=device, + ) + for _ in range(self.token_buffer_count) + ) + self.route_buffers = tuple( + self._torch.zeros( + (self._context.nv_s,), + dtype=self._torch.float32, + device=device, + ) + for _ in range(self.token_buffer_count) + ) + self._route_by_token_ptr = { + int(token.data_ptr()): route + for token, route in zip(self.token_buffers, self.route_buffers) + } + self._packed_projections = None + self._packed_projection_signature = None + self._reduce_dummy = None + self._reduce_dummy_buffer = None + self._tilexr_remote_prefetches = 0 + self._plan_owner_token = object() + self._dispatch_generation = 0 + self._ctx = self._require_ctx() + if os.environ.get("TILEXR_MINDSPEED_TRACE", "0") == "1": + print( + f"[TileXR MindSpeed rank {self._context.planner_group_rank}] " + f"buffer_owner=TileXRComm boundaries={self.token_buffer_count}", + flush=True, + ) + + def _require_ctx(self): + context = super()._require_ctx() + context.update( + rank=int(self._context.planner_group_rank), + num_sms=int(self.num_sms), + sparse_epoch_prefetch_batch2=False, + direct_local_reduce_batch3=False, + ) + return context + + def _validate_boundary(self, hidden_buffer): + if hidden_buffer is None: + return None + if int(hidden_buffer.data_ptr()) not in self._route_by_token_ptr: + raise RuntimeError("hidden_buffer is not owned by this TileXR Buffer") + if ( + tuple(hidden_buffer.shape) != (self._context.nv_s, self.H) + or hidden_buffer.dtype != self._torch.bfloat16 + or not hidden_buffer.is_contiguous() + ): + raise RuntimeError("hidden_buffer has an incompatible TileXR boundary layout") + return self._route_by_token_ptr[int(hidden_buffer.data_ptr())] + + def _bind_zero_copy_views(self, plan, hidden, route_weights): + self._dispatch_generation += 1 + hidden_alias = hidden.view_as(hidden) + hidden_alias._moonep_buffer_owner = self._plan_owner_token + hidden_alias._moonep_dispatch_generation = self._dispatch_generation + hidden_alias._moonep_dispatch_plan = plan + route_alias = ( + route_weights.view_as(route_weights) + if route_weights is not None + else None + ) + return hidden_alias, route_alias + + def _dump_native_plan_once(self, plan): + dump_dir = os.environ.get("TILEXR_MINDSPEED_PLAN_DUMP_DIR") + if not dump_dir: + return + native = plan._require_native() + os.makedirs(dump_dir, exist_ok=True) + path = os.path.join( + dump_dir, + f"rank{self._context.planner_group_rank}_epoch{native.epoch}.pt", + ) + if os.path.exists(path): + return + names = ( + "dst", + "experts_to_copy", + "zero_fill_ranges", + "remote_stats", + "dup_groups", + "dup_loffs", + "dup_counts", + "status", + "reduce_grad_status", + "workspace", + ) + self._torch.save( + { + "epoch": int(native.epoch), + "tensors": { + name: getattr(native, name).detach().cpu() + for name in names + }, + }, + path, + ) + + def dispatch(self, *args, hidden_buffer=None, **kwargs): + async_finish = bool(kwargs.pop("async_finish", False)) + zero_copy = bool(kwargs.pop("zero_copy", False)) + result = super().dispatch( + *args, + async_finish=False, + zero_copy=False, + **kwargs, + ) + hidden, route_weights, cu_seqlens, plan = result + if os.environ.get("TILEXR_MINDSPEED_TRACE", "0") == "1": + native_plan = plan._require_native() + print( + f"[TileXR MindSpeed rank " + f"{self._context.planner_group_rank}] " + f"dispatch_sync_begin epoch={native_plan.epoch}", + flush=True, + ) + try: + self._native_buffer.synchronize() + except BaseException as exc: + print( + f"[TileXR MindSpeed rank " + f"{self._context.planner_group_rank}] " + f"dispatch_device_diag epoch={native_plan.epoch} " + f"exception={type(exc).__name__} message={exc!s}", + flush=True, + ) + raise + print( + f"[TileXR MindSpeed rank " + f"{self._context.planner_group_rank}] " + f"dispatch_sync_end epoch={native_plan.epoch}", + flush=True, + ) + self._dump_native_plan_once(plan) + route_boundary = self._validate_boundary(hidden_buffer) + if hidden_buffer is not None: + hidden_buffer.copy_(hidden) + hidden = hidden_buffer + if route_weights is not None: + route_boundary.copy_(route_weights) + route_weights = route_boundary + if zero_copy: + hidden, route_weights = self._bind_zero_copy_views( + plan, hidden, route_weights + ) + self._zero_copy_aliases = ( + plan._require_native(), + hidden, + route_weights, + ) + event = self._record_event() if async_finish else None + public_result = (hidden, route_weights, cu_seqlens, plan) + return (*public_result, event) if async_finish else public_result + + def _stage_route_weights(self, route_weights, *, hidden_buffer): + route_boundary = self._validate_boundary(hidden_buffer) + if route_weights is None: + return None + flattened = route_weights.reshape(-1) + if int(flattened.numel()) != int(route_boundary.numel()): + raise RuntimeError("route weight staging does not match the TileXR boundary") + route_boundary.copy_(flattened) + return route_boundary + + def combine(self, *args, hidden_buffer=None, **kwargs): + zero_copy = bool(kwargs.pop("zero_copy", False)) + if hidden_buffer is not None: + self._validate_boundary(hidden_buffer) + if os.environ.get("TILEXR_MINDSPEED_TRACE", "0") == "1": + plan = kwargs.get("plan") + if plan is None and args: + plan = args[0] + native_plan = plan._require_native() + limit = int(plan.R) * int(plan.NvS) + dst_min = int(plan.dst.min().item()) + dst_max = int(plan.dst.max().item()) + invalid = int( + ((plan.dst >= limit) | (plan.dst < -limit)).sum().item() + ) + dup_counts = plan.dup_counts.cpu().tolist() + print( + f"[TileXR MindSpeed rank {self._context.planner_group_rank}] " + f"combine_plan epoch={native_plan.epoch} dst_min={dst_min} " + f"dst_max={dst_max} limit={limit} invalid={invalid} " + f"dup_counts={dup_counts}", + flush=True, + ) + if invalid: + raise RuntimeError( + f"TileXR Combine plan has {invalid} out-of-range routes" + ) + try: + result = super().combine(*args, zero_copy=False, **kwargs) + if os.environ.get("TILEXR_MINDSPEED_TRACE", "0") == "1": + try: + self._native_buffer.synchronize() + except RuntimeError: + print( + f"[TileXR MindSpeed rank " + f"{self._context.planner_group_rank}] " + f"combine_device_diag epoch={native_plan.epoch} " + f"status={int(native_plan.status.item())} " + f"encoded_and_limit={plan.dup_counts.cpu().tolist()}", + flush=True, + ) + raise + return result + finally: + if zero_copy: + self._zero_copy_aliases = None + + def _ensure_packed_projections(self, local_fc1, local_fc2): + signature = ( + tuple(local_fc1.shape), + local_fc1.dtype, + tuple(local_fc2.shape), + local_fc2.dtype, + ) + if self._packed_projections is None: + local = self._context.experts_per_rank + dummy = self._torch.zeros( + (local, 32), + dtype=self._torch.bfloat16, + device=local_fc1.device, + ) + projections = ProjectionBuffers.from_local_weights( + self._context, + local_fc1, + dummy, + local_fc2, + torch_module=self._torch, + ) + self._native_buffer.register_projection_buffers(projections) + self._packed_projections = projections + self._packed_projection_signature = signature + elif signature != self._packed_projection_signature: + raise RuntimeError("packed projection shape changed after TileXR registration") + local = self._context.experts_per_rank + self._packed_projections.gate.narrow(0, 0, local).copy_(local_fc1) + self._packed_projections.up.narrow(0, 0, local).zero_() + self._packed_projections.down.narrow(0, 0, local).copy_(local_fc2) + return self._packed_projections + + def _prefetch_weight_packed_batch2( + self, + plan, + *, + full_weights, + local_experts, + local_slots, + source_vas, + async_finish=False, + ): + del local_slots, source_vas + full_fc1, full_fc2 = full_weights + local_fc1, local_fc2 = local_experts + projections = self._ensure_packed_projections(local_fc1, local_fc2) + native_plan = plan._require_native() + trace = os.environ.get("TILEXR_MINDSPEED_TRACE", "0") == "1" + if trace: + print( + f"[TileXR MindSpeed rank {self._context.planner_group_rank}] " + f"packed_prefetch_begin epoch={native_plan.epoch}", + flush=True, + ) + try: + self._native_buffer.prefetch_weight( + native_plan, projections, async_finish=False + ) + except BaseException as exc: + if trace: + print( + f"[TileXR MindSpeed rank " + f"{self._context.planner_group_rank}] " + f"packed_prefetch_exception epoch={native_plan.epoch} " + f"exception={type(exc).__name__} message={exc!s}", + flush=True, + ) + raise + if trace: + print( + f"[TileXR MindSpeed rank {self._context.planner_group_rank}] " + f"packed_prefetch_end epoch={native_plan.epoch}", + flush=True, + ) + local = self._context.experts_per_rank + rank = self._context.planner_group_rank + active_remote = 0 + experts = plan.experts_to_copy[rank] + for slot in range(self.B): + expert = int(experts[slot].item()) + if expert < 0: + continue + full_fc1[self.E + slot].copy_(projections.gate[local + slot]) + full_fc2[self.E + slot].copy_(projections.down[local + slot]) + if expert // local != rank: + active_remote += 1 + self._tilexr_remote_prefetches += active_remote + if os.environ.get("TILEXR_MINDSPEED_TRACE", "0") == "1": + print( + f"[TileXR MindSpeed rank {rank}] packed_prefetch_remote=" + f"{active_remote} total={self._tilexr_remote_prefetches}", + flush=True, + ) + return self._record_event() if async_finish else None + + def _ensure_reduce_dummy(self, full_fc1, reduce_fc1): + dummy_width = ( + 262160 + if os.environ.get("TILEXR_MINDSPEED_FORCE_DUMMY_UDMA", "0") == "1" + else 16 + ) + full_shape = (self.E + self.B, dummy_width) + reduce_shape = (self.R, self.B, dummy_width) + if self._reduce_dummy is None: + self._reduce_dummy = self._torch.zeros( + full_shape, dtype=self._torch.float32, device=full_fc1.device + ) + self._reduce_dummy_buffer = self._torch.zeros( + reduce_shape, dtype=self._torch.float32, device=reduce_fc1.device + ) + else: + self._reduce_dummy.zero_() + self._reduce_dummy_buffer.zero_() + return self._reduce_dummy, self._reduce_dummy_buffer + + def _reduce_grad_packed_batch2( + self, + plan, + *, + full_grads, + reduce_buffers, + local_grads, + local_slots, + local_staging, + source_vas, + ): + del local_slots, local_staging, source_vas + full_fc1, full_fc2 = full_grads + reduce_fc1, reduce_fc2 = reduce_buffers + local_fc1, local_fc2 = local_grads + dummy, dummy_reduce = self._ensure_reduce_dummy(full_fc1, reduce_fc1) + trace = os.environ.get("TILEXR_MINDSPEED_TRACE", "0") == "1" + native_plan = plan._require_native() + if trace: + print( + f"[TileXR MindSpeed rank {self._context.planner_group_rank}] " + f"packed_reduce_grad_begin epoch={native_plan.epoch} " + f"full_shapes={[tuple(value.shape) for value in (full_fc1, dummy, full_fc2)]} " + f"reduce_shapes={[tuple(value.shape) for value in (reduce_fc1, dummy_reduce, reduce_fc2)]} " + f"experts_to_copy={native_plan.experts_to_copy.cpu().tolist()}", + flush=True, + ) + try: + super().reduce_grad( + plan=plan, + full_gate_grad=full_fc1, + full_up_grad=dummy, + full_down_grad=full_fc2, + gate_reduce_buffer=reduce_fc1, + up_reduce_buffer=dummy_reduce, + down_reduce_buffer=reduce_fc2, + ) + except BaseException as exc: + if trace: + info = self._native_buffer.reduce_grad_info + print( + f"[TileXR MindSpeed rank {self._context.planner_group_rank}] " + f"packed_reduce_grad_exception epoch={native_plan.epoch} " + f"exception={type(exc).__name__} message={exc!s} " + f"status={int(native_plan.reduce_grad_status.item())} " + f"info={info}", + flush=True, + ) + raise + local = self._context.experts_per_rank + begin = self._context.planner_group_rank * local + local_fc1.copy_(full_fc1.narrow(0, begin, local)) + local_fc2.copy_(full_fc2.narrow(0, begin, local)) + if os.environ.get("TILEXR_MINDSPEED_TRACE", "0") == "1": + print( + f"[TileXR MindSpeed rank {self._context.planner_group_rank}] " + f"packed_reduce_grad=1 info={self._native_buffer.reduce_grad_info}", + flush=True, + ) + return local_fc1, local_fc2 + + def destroy(self): + self._packed_projections = None + self._reduce_dummy = None + self._reduce_dummy_buffer = None + self.token_buffers = () + self.route_buffers = () + self._route_by_token_ptr = {} + super().destroy() + + +def create_tilexr_moonep_backend(**kwargs): + kwargs["buffer_cls"] = MindSpeedTileXRBuffer + return MoonEPBufferFlexBackend(**kwargs) + + +create_tilexr_moonep_backend.__mindspeed_capabilities__ = { + MOONEP_NATIVE_NPU_CAPABILITY +} +create_tilexr_moonep_backend.__mindspeed_external_communication_owner__ = True diff --git a/tools/moonep/report.py b/tools/moonep/report.py index 977db39a..4331f3ed 100644 --- a/tools/moonep/report.py +++ b/tools/moonep/report.py @@ -393,6 +393,66 @@ def format_row(row: tuple[str, ...]) -> str: ) +def format_dispatch_performance(summary: Mapping[str, object]) -> str: + if summary.get("benchmark_kind") != "dispatch_hot_loop": + raise ValueError("summary is not a Dispatch hot-loop report") + modes = summary.get("dispatch_modes") + metrics = summary.get("metrics_us") + throughput = summary.get("tokens_per_second_by_mode") + if not isinstance(modes, list) or not modes: + raise ValueError("Dispatch hot-loop summary has no modes") + if not isinstance(metrics, dict) or not isinstance(throughput, dict): + raise ValueError("Dispatch hot-loop summary has incomplete metrics") + + headers = ( + "Mode", + "Host mean us", + "Kernel mean us", + "Kernel P50 us", + "Kernel P95 us", + "Mean token/s", + ) + rows = [] + for mode in modes: + host = metrics.get(f"{mode}_host") + kernel = metrics.get(f"{mode}_kernel") + mode_throughput = throughput.get(mode) + if not all(isinstance(value, dict) for value in (host, kernel, mode_throughput)): + raise ValueError(f"Dispatch hot-loop summary is missing {mode} metrics") + rows.append( + ( + str(mode), + _format_number(float(host["mean"]), precision=3), + _format_number(float(kernel["mean"]), precision=3), + _format_number(float(kernel["p50"]), precision=3), + _format_number(float(kernel["p95"]), precision=3), + _format_number(float(mode_throughput["mean"]), precision=3), + ) + ) + + widths = [ + max(len(headers[index]), *(len(row[index]) for row in rows)) + for index in range(len(headers)) + ] + + def format_row(row: tuple[str, ...]) -> str: + return " ".join( + value.ljust(widths[index]) if index == 0 else value.rjust(widths[index]) + for index, value in enumerate(row) + ) + + return "\n".join( + ( + _format_benchmark_inputs(summary), + "", + "MoonEP Dispatch-only performance (global critical rank per iteration)", + format_row(headers), + format_row(tuple("-" * width for width in widths)), + *(format_row(row) for row in rows), + ) + ) + + def _write_stage_csv( path: Path, stage_performance: Mapping[str, object], @@ -920,7 +980,12 @@ def main(argv: list[str] | None = None) -> int: if args.summary is not None: with Path(args.summary).open("r", encoding="utf-8") as handle: summary = json.load(handle) - print(format_stage_performance(summary)) + formatter = ( + format_dispatch_performance + if summary.get("benchmark_kind") == "dispatch_hot_loop" + else format_stage_performance + ) + print(formatter(summary)) return 0 if not args.case_ids or args.node_count is None or args.world_size is None: raise ValueError(