diff --git a/docs/superpowers/specs/2026-07-28-tilexr-host-transport-routing-design.md b/docs/superpowers/specs/2026-07-28-tilexr-host-transport-routing-design.md new file mode 100644 index 00000000..35ee91b2 --- /dev/null +++ b/docs/superpowers/specs/2026-07-28-tilexr-host-transport-routing-design.md @@ -0,0 +1,1075 @@ +# TileXR EP Memory / Direct URMA 统一路由设计 + +## 1. 文档定位 + +本文是 TileXR EP 通信在 `MEMORY` 与 `DIRECT_URMA` 两条数据面之间进行统一路由的权威设计与 +实现说明。内容以 2026-07-28 的 `feature/tilexr-auto-route` 实际代码和 Ascend950PR 硬件验证结果 +为准。 + +本文同时回答以下问题: + +- Host API 如何在不改变公开调用入口的情况下选择数据面; +- `auto`、`memory`、`direct_urma` 三种模式的精确语义; +- 同机和跨机分别使用什么阈值; +- selector 中的 `bytes` 在当前 EP 实现里具体表示什么; +- Memory 与 Direct URMA 各自使用什么地址空间、同步方式和完成语义; +- Direct URMA workspace 如何计算、注册和校验; +- 同机 Direct URMA 如何确保设备端真正走 UDMA; +- 哪些能力已经在硬件上验证,哪些仍属于剩余风险。 + +本文只覆盖 EP dispatch/combine 的统一路由。collectives、独立 UDMA demo、SDMA,以及 EP 之外的 +其他模块不自动继承本文策略。 + +## 2. 目标与非目标 + +### 2.1 目标 + +保留同一组 EP Host API,由 Host 在每次 dispatch 或 combine 调用前选择以下一条真实数据面: + +```text +MEMORY + peerMems[] + IPC_DATA_OFFSET + AICore DataCopyPad + +DIRECT_URMA + ordinary device workspace + TileXRUDMARegister + UDMA put/get/signal +``` + +设计必须满足: + +1. Memory 和 Direct URMA 均可强制选择并独立验证。 +2. 同机和跨机 Memory 复用同一套 peer-memory/DataCopy 实现。 +3. 同机和跨机 Direct URMA 复用同一套 registered-workspace kernel。 +4. 一次 Host API 调用只解析一次 route,校验和 kernel launch 使用同一个结果。 +5. `auto` 只做策略选择,不执行数据搬运,也不隐式申请或注册 workspace。 +6. Direct URMA 不可用时,`auto` 可选择 Memory;强制 Direct URMA 不允许静默回退。 +7. TCP socket 只用于 communicator rendezvous、注册信息交换和测试 barrier,不承载 EP payload。 + +### 2.2 非目标 + +- 不引入 host-staging 作为第三种生产 transport。 +- 不把 Memory 和 UDMA 强行封装成一个缺少地址上下文的 device put/get primitive。 +- 不修改 EP 公共 C API 签名。 +- 不修改 UDMA registered-memory 的 offset 地址模型。 +- 不修改 Ascend950 IPC buffer 的 malloc policy。 +- 不使用 host-only、simulator 或 host-staging 结果替代真实 Ascend950 数据面验证。 +- 不声称当前 rank-2 验证已经覆盖任意 rank 数、混合节点拓扑和 TP 组合。 + +## 3. 核心结论 + +### 3.1 最终 Auto 策略 + +在 communicator 元数据有效且 Direct URMA capability/registry 可用时: + +```text +跨机:routeBytes < 128 KiB -> MEMORY +跨机:routeBytes >= 128 KiB -> DIRECT_URMA + +同机:routeBytes < 4 MiB -> MEMORY +同机:routeBytes >= 4 MiB -> DIRECT_URMA +``` + +若 Direct URMA 不可用,则不论大小均选择 `MEMORY`。边界使用 `>=`,因此 `128 KiB` 和 `4 MiB` +本身属于 Direct URMA。 + +### 3.2 当前 `routeBytes` 的精确定义 + +当前 EP Host 实现传给 `TileXRSelectAutoTransport` 的不是输入 tensor 字节数,也不是某一个 peer 的 +实际发送字节数,而是: + +```cpp +routeBytes = static_cast(context.window.totalBytes); +``` + +`window.totalBytes` 是单个 EP window 的总容量,包含 window header 和全部 rank slot 的最坏情况 +容量。本文统一称其为 `routeBytes`。 + +这一定义非常重要: + +- `routeBytes` 是 Host 在 launch 前可稳定计算的 operation footprint; +- 它与当前 Memory window 容量校验使用同一口径; +- 它不是运行时实际命中的 route 数,也不是单条 UDMA put 的字节数; +- CSV 中的 `bytes` 是单向单流 benchmark 的实际传输大小,与 `routeBytes` 不是完全相同的统计口径。 + +因此,当前两个阈值已经完成路由正确性和端到端数据正确性验证,但不能据此声称每个 EP shape 都 +实现了严格的性能最优选择。若后续要求更精细的性能路由,应评估改用 `slotBytes`、最大远端 peer +payload 或运行时 route histogram;该变化属于新的策略版本,必须重新采样和验证。 + +## 4. 术语与拓扑前提 + +| 术语 | 精确定义 | +|---|---| +| `rankSize` | communicator 的全局 rank 数 | +| `localRankSize` | communicator 认为同一 full-mesh 本地节点内的 rank 数 | +| 同机 | `localRankSize == rankSize` | +| 跨机 | `0 < localRankSize < rankSize` | +| `peerMems[]` | communicator 建立并上传到 `CommArgs` 的 IPC peer window 地址 | +| registered workspace | 应用申请的普通 device memory,经 `TileXRUDMARegister` 注册后形成的 UDMA region | +| `routeBytes` | 当前实现中的 `EpWindowConfig::totalBytes` | +| operation | 一次 dispatch 或一次 combine Host API 调用 | + +当前 device helper 使用 `rank / localRankSize` 判断节点归属,因此实现假设: + +- `localRankSize` 大于 0; +- 每个节点包含相同数量的 rank; +- 同一节点的 global rank 连续排列。 + +不满足这些前提的非均匀拓扑不在当前设计验证范围内。 + +## 5. 总体架构 + +路由只在 Host 层解析,设备 kernel 不再读取环境变量,也不重新执行 auto selector。 + +```mermaid +flowchart TD + A["EP dispatch/combine API"] --> B["校验基础参数并构建 EpWindowConfig"] + B --> C["读取 TILEXR_TRANSPORT_MODE"] + C --> D["ResolveTransport(mode, CommArgs, window.totalBytes)"] + D --> E["将结果保存到 EpHostLaunchContext.transport"] + E --> F["按 route 校验 peerMems / registry / workspace"] + F --> G{"transport"} + G -->|"MEMORY"| H["ordinary peer-memory kernel"] + G -->|"DIRECT_URMA"| I["registered-workspace UDMA kernel"] +``` + +关键边界如下: + +- `TileXREpPrepareLaunchContext` 负责 dispatch 的 window、route 和资源校验; +- `TileXREpPrepareCombineLaunchContext` 对 combine 独立执行同样流程; +- `EpHostLaunchContext.transport` 是一次调用内唯一可信 route; +- `ep_kernel_launch.cpp` 只读取 `context.transport`,不得再次读取 env 或 selector; +- dispatch 和 combine 是两个独立 operation,会分别解析 route;调用方不应在并发或配对调用之间 + 动态修改进程级 `TILEXR_TRANSPORT_MODE`。 + +### 5.1 代码模块与职责 + +| 模块 | 关键文件 | 职责 | +|---|---|---| +| 公共 selector | `src/include/tilexr_transport.h` | 定义 transport kind、capability 判定、同机/跨机阈值和 auto 决策 | +| EP mode adapter | `src/ep/host/ep_transport_route.{h,cpp}` | 解析环境变量,将 `auto/memory/direct_urma` 转换为确定 route | +| EP window/layout | `src/ep/host/ep_layout.cpp`、`src/ep/common/ep_window.h` | 计算 `slotBytes`、`totalBytes`、UDMA operation/workspace 布局 | +| Host launch context | `src/ep/host/ep_launch_context.cpp` | 一次性完成参数、route、peer mapping、registry 和 workspace 校验 | +| Kernel launch | `src/ep/host/ep_kernel_launch.cpp` | 只依据 `context.transport` 选择 Memory 或 Direct kernel | +| Device data plane | `src/ep/kernels/tilexr_ep_*` | 实现 IPC DataCopy、UDMA put/quiet、ready/status 和 drain | +| Communicator | `src/comm/tilexr_comm.cpp` | 发现 `rankSize/localRankSize`、建立 `peerMems[]`、初始化 UDMA context | +| UDMA context | `src/comm/udma/tilexr_udma_context.cpp` | 对外管理 register/unregister、registry 和 `CommArgs` 状态 | +| UDMA transport | `src/comm/udma/tilexr_udma_transport.cpp` | 建 route/context/QP、注册本地 MR、导入远端 MR、生成 device UDMA info | + +```mermaid +flowchart LR + API["EP public API"] --> LC["ep_launch_context.cpp"] + LC --> LAYOUT["ep_layout.cpp"] + LC --> ROUTE["ep_transport_route.cpp"] + ROUTE --> SELECTOR["tilexr_transport.h"] + LC --> LAUNCH["ep_kernel_launch.cpp"] + LAUNCH --> MEMK["Memory kernels"] + LAUNCH --> URMAK["Direct URMA kernels"] + COMM["tilexr_comm.cpp"] --> IPC["peerMems IPC mappings"] + COMM --> UCTX["TileXRUDMAContext"] + UCTX --> UTRANS["TileXRUDMATransport"] + IPC --> MEMK + IPC --> URMAK + UTRANS --> INFO["CommArgs.udmaInfoPtr"] + UCTX --> REG["CommArgs.udmaRegistryPtr"] + INFO --> URMAK + REG --> URMAK +``` + +## 6. EP Window 与路由大小计算 + +### 6.1 基础常量 + +```text +kEpWindowAlignmentBytes = 32 +kEpWindowHeaderBytes = 64 +kEpSrcSlotHeaderBytes = 64 +kEpAssistTupleInts = 4 +sizeof(EpAssistTuple) = 16 bytes +IPC_BUFF_MAX_SIZE = 100 MiB +IPC_DATA_OFFSET = 2 MiB +``` + +### 6.2 Dispatch window + +设: + +```text +routesPerToken = topK + sharedExpertNum +maxRoutesPerSrc = bs * routesPerToken +dtypeBytes = 1 for INT8, 2 for FP16/BFP16 +rowBytes = h * dtypeBytes +payloadExtra = 4 bytes when quantMode == per-token dynamic, otherwise 0 +``` + +则: + +```text +payloadBytesPerSlot = align32(maxRoutesPerSrc * (rowBytes + payloadExtra)) +assistBytesPerSlot = align32(maxRoutesPerSrc * 16) +slotBytes = align32(64 + payloadBytesPerSlot + assistBytesPerSlot) +totalBytes = 64 + rankSize * slotBytes +routeBytes = totalBytes +``` + +所有乘加操作在 Host 侧进行有符号 64 位溢出检查。`totalBytes > 100 MiB` 时直接返回参数错误。 + +TP 模式还要求: + +```text +align32(totalBytes) * (effectiveTpWorldSize + 2) <= 100 MiB +``` + +该检查用于 TP IPC 辅助窗口,不改变 auto selector 仍以 `totalBytes` 为输入的事实。 + +### 6.3 Combine window + +combine 使用同一基础 builder,但当前 API 不传 `sharedExpertNum` 和动态量化额外 scale bytes,因此其 +window 以 `routesPerToken = topK` 计算。dispatch 与 combine 的 `routeBytes` 可能不同,并可独立落在 +不同阈值区间;这是当前实现允许的行为。 + +## 7. Mode 解析与决策规则 + +### 7.1 模式来源 + +当前使用进程环境变量: + +```text +TILEXR_TRANSPORT_MODE=auto +TILEXR_TRANSPORT_MODE=memory +TILEXR_TRANSPORT_MODE=direct_urma +``` + +未设置、空字符串或显式 `auto` 均解析为 `AUTO`。其他未知值返回 +`TILEXR_ERROR_PARA_CHECK_FAIL`,不使用默认值掩盖配置错误。 + +环境变量是当前不修改公开 API 的临时 override,不是线程局部配置。应用应在创建通信工作负载前 +设置一次,不应在多线程 launch 期间修改。 + +### 7.2 Direct URMA capability 判定 + +`TileXRDirectUrmaAvailable(args)` 必须同时满足: + +```text +args != nullptr +args.extraFlag 包含 ExtraFlag::UDMA +args.udmaInfoPtr != nullptr +args.udmaRegistryPtr != nullptr +``` + +其中 `udmaRegistryPtr` 只有成功调用 `TileXRUDMARegister` 并同步 `CommArgs` 后才存在。因此: + +- `auto` 不会自动申请或注册 workspace; +- 大包希望进入 Direct URMA 时,调用方必须在 EP API 前完成 workspace 注册; +- 未注册时 `auto` 会把 Direct URMA 判定为不可用并选择 Memory; +- 为验证阈值本身,EP auto demo 即使测试小包也先注册 workspace,使 selector 能在两条路之间按大小 + 选择,而不是因为 capability 缺失被迫选择 Memory。 + +### 7.3 精确决策表 + +| Mode | 条件 | 结果 | +|---|---|---| +| `MEMORY` | 不检查 UDMA capability | `MEMORY` | +| `DIRECT_URMA` | capability 完整 | `DIRECT_URMA` | +| `DIRECT_URMA` | capability 缺失 | `TILEXR_ERROR_NOT_INITIALIZED` | +| `AUTO` | `routeBytes == 0` | `MEMORY` | +| `AUTO` | capability 缺失 | `MEMORY` | +| `AUTO` | 跨机且 `routeBytes < 128 KiB` | `MEMORY` | +| `AUTO` | 跨机且 `routeBytes >= 128 KiB` | `DIRECT_URMA` | +| `AUTO` | 同机且 `routeBytes < 4 MiB` | `MEMORY` | +| `AUTO` | 同机且 `routeBytes >= 4 MiB` | `DIRECT_URMA` | + +`routeBytes == 0` 是通用 selector 的防御行为;合法 EP shape 的 `totalBytes` 实际大于 0。 + +```mermaid +flowchart TD + A["输入 mode、CommArgs、routeBytes"] --> B{"mode"} + B -->|"memory"| M["选择 MEMORY"] + B -->|"direct_urma"| C{"Direct capability 完整"} + C -->|"否"| E["返回 NOT_INITIALIZED"] + C -->|"是"| U["选择 DIRECT_URMA"] + B -->|"auto"| Z{"routeBytes > 0 且 capability 完整"} + Z -->|"否"| M + Z -->|"是"| T{"localRankSize < rankSize"} + T -->|"跨机"| X{"routeBytes >= 128 KiB"} + T -->|"同机"| Y{"routeBytes >= 4 MiB"} + X -->|"是"| U + X -->|"否"| M + Y -->|"是"| U + Y -->|"否"| M +``` + +兼容常量 `TILEXR_AUTO_DIRECT_URMA_THRESHOLD_BYTES` 仍保留,并等于同机 `4 MiB` 阈值。新代码应优先 +使用显式的 SAME_NODE/CROSS_NODE 常量,避免误解。 + +### 7.4 Auto 回退边界 + +Auto 的回退只发生在 selector 阶段:Direct URMA capability 不完整时选择 Memory。 + +若 selector 已选择 Direct URMA,但随后发现以下问题,则返回错误,不再二次回退: + +- workspace 为空; +- registry 结构无效; +- 任一 rank 注册区不足; +- 当前 rank 注册 base 与传入 workspace 不一致; +- 本地 peer mapping 缺失。 + +这是为了避免“按 Direct 校验到一半后静默改走 Memory”造成 route、资源和 kernel 不一致。 + +## 8. Host Launch Context 与资源校验 + +### 8.1 固定处理顺序 + +dispatch 和 combine 均按以下顺序处理: + +1. 获取 Host `CommArgs`。 +2. 获取 Device `CommArgs`。 +3. 校验 API 参数并构建 `EpWindowConfig`。 +4. 以 `window.totalBytes` 解析 route。 +5. 保存到 `EpHostLaunchContext.transport`。 +6. 按 route 校验 peer mapping。 +7. 按 route 校验 Memory window 或 registered workspace。 +8. 启动与 route 对应的 kernel。 + +任一步失败都会清空 launch context 并返回错误。 + +```mermaid +sequenceDiagram + participant App as "调用方" + participant API as "TileXRMoeEpDispatch/Combine" + participant Ctx as "PrepareLaunchContext" + participant Route as "ResolveTransportFromEnv" + participant Comm as "TileXR communicator" + participant Launch as "Kernel launcher" + participant Device as "Ascend kernel" + + App->>API: "params + comm + workspace" + API->>Ctx: "构建本次 launch context" + Ctx->>Comm: "获取 Host/Device CommArgs" + Ctx->>Ctx: "计算 EpWindowConfig" + Ctx->>Route: "mode + CommArgs + totalBytes" + Route-->>Ctx: "resolved transport" + Ctx->>Ctx: "按 transport 校验资源" + Ctx-->>API: "EpHostLaunchContext" + API->>Launch: "context + params" + Launch->>Device: "启动唯一对应 kernel" + Device-->>App: "stream completion / status" +``` + +### 8.2 MEMORY 校验 + +Memory route: + +- 要求所有 global rank 的 `peerMems[rank]` 非空; +- 不要求 `ExtraFlag::UDMA`; +- 不读取 UDMA registry; +- 不要求 `params.workspace`; +- 要求 `0 < totalBytes <= IPC_BUFF_MAX_SIZE`。 + +Memory dispatch launch 显式向 ordinary kernel 传 `workspace=nullptr`。即使调用方传入了 workspace, +Memory kernel 也不能因为该指针非空而误进入历史 UDMA 分支。 + +### 8.3 DIRECT_URMA 校验 + +Direct route 首先要求 UDMA capability,然后检查: + +- `params.workspace != nullptr`; +- Host registry 的 `rankSize` 与 communicator 一致; +- 每个 global rank 的 region 都覆盖 `[0, requiredWorkspaceBytes)`; +- 当前 rank registry base 与 `params.workspace` 完全相等; +- 当前节点范围内的 `peerMems[]` 非空。 + +Direct route 只要求本节点 peer mappings: + +```text +beginRank = floor(rank / localRankSize) * localRankSize +endRank = min(beginRank + localRankSize, rankSize) +``` + +跨节点 peer payload 使用 UDMA,不要求远端 `peerMems[]` 映射;本节点 IPC 仍用于混合路径和 TP 辅助 +交换,因此本节点 mapping 不能缺失。同机场景中本节点就是全部 rank,所以仍会检查全部 `peerMems[]`。 + +## 9. MEMORY 数据面 + +### 9.1 地址模型 + +Memory 数据区地址为: + +```text +peerMems[peer] + IPC_DATA_OFFSET + window-relative offset +``` + +每个 IPC allocation 的布局是: + +```text +[ 2 MiB reusable flag area ][ 100 MiB data area ] +^ peerMems[rank] ^ peerMems[rank] + IPC_DATA_OFFSET +``` + +因此 EP 数据容量判断是 `totalBytes <= 100 MiB`,不能再次把 `IPC_DATA_OFFSET` 计入 100 MiB 数据区 +并错误扣减容量。 + +### 9.2 Communicator 建立 peer window + +进程模式 communicator 的主要步骤为: + +1. 当前 rank 申请本地 `peerMem_[rank]`。 +2. `rtIpcSetMemoryName` 导出 IPC memory name。 +3. Ascend950/SuperPod 路径通过 `rtSetIpcMemorySuperPodPid(name, sdid, pid)` 授权 peer。 +4. 其他 rank 使用 `rtIpcOpenMemory` 打开该 window。 +5. `SyncCommArgs` 将 `peerMems[]` 上传到 device。 + +跨机 Memory 因此仍是设备侧 peer-memory DataCopy,不是 TCP host-staging。 + +```mermaid +sequenceDiagram + participant R0 as "Rank 0 process" + participant Sock as "Socket exchange" + participant R1 as "Rank 1 process" + participant D0 as "Rank 0 NPU" + participant D1 as "Rank 1 NPU" + + R0->>D0: "aclrtMalloc peerMem[0]" + R1->>D1: "aclrtMalloc peerMem[1]" + R0->>Sock: "发布 IPC name / pid / sdid" + R1->>Sock: "发布 IPC name / pid / sdid" + Sock-->>R0: "peer metadata" + Sock-->>R1: "peer metadata" + R0->>D0: "rtIpcOpenMemory(peer 1)" + R1->>D1: "rtIpcOpenMemory(peer 0)" + R0->>D0: "上传 CommArgs.peerMems[]" + R1->>D1: "上传 CommArgs.peerMems[]" + D0->>D1: "AICore DataCopyPad payload" + D1->>D0: "AICore DataCopyPad payload" +``` + +### 9.3 搬运与同步 + +ordinary kernel 使用 64 KiB UB,其中 4 KiB 保留给同步,剩余区域分块完成: + +```text +source GM -> UB scratch -> peer GM +``` + +搬运使用 `DataCopyPad`,并通过 MTE event 与 `PipeBarrier` 保证块间顺序。同步 flag 使用每次调用从 +`TileXRCommNextMagic` 获取的新 magic,不通过批量清零 flag 区开始新一轮。 + +同机和跨机 Memory 使用相同 kernel、地址模型和完成逻辑。TCP diagnostic fallback 不属于该路径。 + +## 10. DIRECT_URMA Workspace + +### 10.1 注册模型 + +Direct URMA workspace 是应用拥有的一段普通 device memory: + +```text +aclrtMalloc ordinary device memory + -> TileXRUDMARegister(comm, workspace, bytes, &handle) + -> allgather each rank {base, bytes} + -> build one-region-per-rank TileXRUDMARegistry + -> upload registry and update CommArgs.udmaRegistryPtr +``` + +当前 registry 每个 rank 只有一个连续 region,因此 dispatch、combine、relay 和 status 必须位于同一 +连续注册 workspace 内。`TileXRUDMARegister` 需要 live socket exchange,不支持 `InitThread` 模式。 + +```mermaid +sequenceDiagram + participant App as "应用 / EP demo" + participant API as "TileXRUDMARegister" + participant Ctx as "TileXRUDMAContext" + participant Trans as "TileXRUDMATransport" + participant HCCP as "HCCP RA" + participant Sock as "Socket AllGather" + participant Dev as "Device CommArgs" + + App->>App: "aclrtMalloc ordinary workspace" + App->>API: "comm, workspace, bytes" + API->>Ctx: "RegisterMemory" + Ctx->>Trans: "RegisterMemoryOnContexts" + Trans->>HCCP: "RaCtxLmemRegister per required EID" + HCCP-->>Trans: "local key / token / segment" + Trans->>Sock: "AllGather MR metadata" + Trans->>HCCP: "RaCtxRmemImport for UDMA peers" + Trans->>Sock: "AllGather UDMAMemInfo / EID" + Trans->>Dev: "上传 UDMAInfo image" + Ctx->>Sock: "AllGather each rank base/bytes" + Ctx->>Dev: "上传 TileXRUDMARegistry" + Ctx->>Dev: "更新 udmaInfoPtr / udmaRegistryPtr / UDMA flag" + API-->>App: "registration handle" +``` + +### 10.1.1 多节点 UDMA 资源范围 + +UDMA transport 初始化必须接收 communicator 的 `localRankSize`。在多节点 communicator 中,EP 的 +Direct 数据面是 IPC/UDMA 混合模式,同节点 peer 不会执行 UDMA put,因此 Host 只为跨节点 peer +建立 UDMA route、context/QP 映射和 remote MR import: + +```text +peer == self + -> 不分配 peer UDMA 资源 + +localRankSize < rankSize && floor(peer / localRankSize) == floor(rank / localRankSize) + -> 同节点 peer,使用 peerMems[],不分配 peer UDMA 资源 + +其他 peer + -> 跨节点 peer,建立 UDMA route/QP/MR +``` + +纯同机 communicator 的 `localRankSize == rankSize`,此时除 self 外所有 peer 仍保留 UDMA 资源, +以支持同机大报文 Direct URMA。该规则只裁剪多节点 hybrid 模式中的冗余资源,不改变同机 Direct +语义,也不改变 socket AllGather 的全局 rank 参与范围。 + +`localRankSize` 的 Host 传递链是: + +```text +TileXRComm::GetDev + -> TileXRComm::localRankSize_ + -> TileXRUDMAContextOptions.localRankSize + -> TileXRUDMATransportOptions.localRankSize + -> TileXRUDMATransport::UsesUDMAPeer +``` + +`UsesUDMAPeer` 的等价逻辑如下: + +```cpp +if (peer is invalid or peer == rank) { + return false; +} +if (localRankSize >= rankSize) { + return true; // 纯同机 Direct:所有非 self peer 使用 UDMA +} +return peer / localRankSize != rank / localRankSize; +``` + +该判定同时用于 `BuildRoutes` 和 remote MR import。被判定为同节点的 peer 不进入 +`peerLocalEid_/peerRemoteEid_`,也不会出现在 `remoteMemHandles_` 中;`RefreshUDMAInfo` 对这些 rank +保留不被 Direct kernel 使用的占位项,从而保持 device table 仍按 global rank 索引。 + +2-host/16-rank 初次扩展时,每个 rank 曾为同节点 IPC peer 也创建 UDMA context 并注册 MR,导致 +部分 rank 的 `RaCtxLmemRegister` 返回 `528101 (ROCE_EOPENSRC)`,随后其他 rank 的注册 AllGather +因连接断开连锁失败。传递 `localRankSize` 并只保留跨节点 UDMA peer 后,本地 MR 注册和 16-rank +dispatch/combine 均通过。 + +UDMA 初始化本身按以下顺序执行。`UsesUDMAPeer(peer)` 是资源规模控制点,必须在 route 解析阶段就 +过滤同节点 peer,而不是创建完 QP/MR 后再在 kernel 中忽略: + +```mermaid +flowchart TD + A["TileXRComm::InitUDMA"] --> B["传入 rankSize、localRankSize、devId、exchange"] + B --> C["TileXRUDMATransport::Init"] + C --> D["OpenDevice"] + D --> E["BuildRoutes"] + E --> F{"UsesUDMAPeer(peer)"} + F -->|"多节点同机 peer"| G["跳过 UDMA route"] + F -->|"跨节点 peer"| H["记录 localEid / remoteEid"] + F -->|"纯同机场景非 self peer"| H + H --> I["CreateContexts"] + I --> J["CreateQueues"] + J --> K["ImportQueues"] + K --> L["RefreshUDMAInfo"] + L --> M["设置 ExtraFlag::UDMA"] +``` + +以 2-host/16-rank、每节点 8 rank 为例,rank 0 的资源关系为: + +```mermaid +flowchart LR + subgraph N0["Node 0: rank 0-7"] + R0["Rank 0"] + L1["Rank 1"] + L2["Rank 2"] + L7["Rank 7"] + end + subgraph N1["Node 1: rank 8-15"] + U8["Rank 8"] + U9["Rank 9"] + U15["Rank 15"] + end + R0 -->|"IPC peerMems"| L1 + R0 -->|"IPC peerMems"| L2 + R0 -->|"IPC peerMems"| L7 + R0 -->|"UDMA QP + imported MR"| U8 + R0 -->|"UDMA QP + imported MR"| U9 + R0 -->|"UDMA QP + imported MR"| U15 +``` + +### 10.2 精确布局 + +定义: + +```text +A = align32(totalBytes) +R = rankSize +S = slotBytes +C = 64,UDMA cache-line bytes +``` + +单个 operation 区布局: + +```text +sendWindowOffset = 0 +recvWindowOffset = A +readyOffset = 2 * A +readyRankOffset(r) = readyOffset + r * C +readyBytes = R * C +relaySlotsOffset = align64(readyOffset + readyBytes) +relaySlotsBytes = R * R * S +relayReadyOffset = align64(relaySlotsOffset + relaySlotsBytes) +relayReadyRankOffset(r) = relayReadyOffset + r * C +relayReadyBytes = R * C +operationBytes = align64(relayReadyOffset + relayReadyBytes) +``` + +完整 workspace: + +```text +dispatchOperationOffset = 0 +combineOperationOffset = operationBytes +statusOffset = 2 * operationBytes +requiredWorkspaceBytes = align64(statusOffset + 8) +``` + +图示: + +```text +workspace +|-- dispatch operation +| |-- send window +| |-- recv window +| |-- ready[R * 64B] +| |-- relay slots[R * R] +| `-- relay ready[R * 64B] +|-- combine operation +| |-- send window +| |-- recv window +| |-- ready[R * 64B] +| |-- relay slots[R * R] +| `-- relay ready[R * 64B] +`-- status[1] +``` + +```mermaid +flowchart TB + W["Registered workspace"] --> D["Dispatch operation"] + W --> C["Combine operation"] + W --> S["Status: 8 bytes aligned to 64B"] + D --> DS["send window"] + D --> DR["recv window"] + D --> DY["ready: R x 64B"] + D --> DL["relay slots: R x R x slotBytes"] + D --> DLR["relay ready: R x 64B"] + C --> CS["send window"] + C --> CR["recv window"] + C --> CY["ready: R x 64B"] + C --> CL["relay slots: R x R x slotBytes"] + C --> CLR["relay ready: R x 64B"] +``` + +每个 ready slot 只使用首 8B 存放 `uint64_t`,其余字节为 cache-line padding。不能把 ready +压缩成连续的 `R * 8B`:当每节点存在多个 rank 时,多个远端 rank 会并发写同一条 64B cache line, +后到写入可能覆盖先到的 ready 值。实际 2-host/4-rank 验证曾表现为 rank 2 在 combine 阶段等待 +rank 0 的 step 74 ready 超时;改为每 rank 独占 cache line 后,4-rank 和 8-rank Direct URMA 均通过。 + +`requiredWorkspaceBytes` 可能显著大于 `totalBytes`,主要原因是 `R * R * slotBytes` relay 区。 +Host 和 Device 必须使用相同公式,不能用历史的 `alignedTotal * (tpWorldSize + 2)` 代替 UDMA +workspace 大小。 + +## 11. DIRECT_URMA 数据面 + +### 11.1 Kernel 选择 + +Host gate 为: + +```text +context.transport == DIRECT_URMA +hostArgs != nullptr +rankSize > 1 +slotBytes > 0 +``` + +满足时 dispatch/combine 均启动当前名为 `*_cross_node_kernel` 的 registered-workspace kernel。 +函数名保留了历史命名,但实现同时支持同机和跨机。单 rank 不属于 Direct URMA 硬件验收范围。 + +### 11.2 同机与跨机分类 + +Device kernel 计算: + +```cpp +useUdmaForAllPeers = (localRankSize == rankSize); +effectiveLocalRankSize = useUdmaForAllPeers ? 1 : localRankSize; +``` + +语义: + +- 纯同机 Direct:除 self 外的所有 EP slot peer 都视为 UDMA peer; +- 跨机 Direct:本节点 peer 继续走 IPC,跨节点 peer 走 UDMA; +- self slot 始终在本地 send/recv window 间复制,不发 UDMA; +- `effectiveLocalRankSize` 只改变 EP slot 的 transport 分类,不修改 communicator 的真实 + `localRankSize`。 + +```mermaid +flowchart LR + subgraph A["Node A"] + A0["Rank A0"] + A1["Rank A1"] + end + subgraph B["Node B"] + B0["Rank B0"] + B1["Rank B1"] + end + A0 <-->|"IPC peer window"| A1 + B0 <-->|"IPC peer window"| B1 + A0 <-->|"registered-memory UDMA"| B0 + A0 <-->|"registered-memory UDMA"| B1 + A1 <-->|"registered-memory UDMA"| B0 + A1 <-->|"registered-memory UDMA"| B1 +``` + +### 11.3 Dispatch + +#### 同机 Direct + +1. `sendWindow = workspace + dispatchOperationOffset`。 +2. `recvWindow = sendWindow + align32(totalBytes)`。 +3. `dispatchIpcWindow = nullptr`,强制 route payload 写入 registered `sendWindow`。 +4. self slot 从 `sendWindow` 本地复制到 `recvWindow`。 +5. 每个非 self peer 使用 UDMA put/signal,把本 rank 的目标 slot 写入对端 `recvWindow` 中以 src rank + 为索引的 slot。 +6. 等待所有 peer ready。 +7. 从本地 `recvWindow` 按 src rank drain 数据。 + +#### 跨机 Direct + +1. 本节点非 self peer 的 route 写入当前 rank 的 IPC window。 +2. 跨节点 peer 的 route 写入 registered `sendWindow`。 +3. self slot本地复制到 `recvWindow`。 +4. 跨节点 slot 使用 UDMA put,ready value 使用 UDMA 写入。 +5. drain 时,本节点 source 从对应 `peerMems[srcRank]` 读取,跨节点 source 从 `recvWindow` 读取。 + +```mermaid +sequenceDiagram + participant Src as "Source rank kernel" + participant Local as "Same-node peer IPC window" + participant Send as "Source registered sendWindow" + participant Remote as "Remote registered recvWindow" + participant Dst as "Destination rank kernel" + + Src->>Src: "按 destination rank 构造 slot" + Src->>Local: "DataCopyPad 本节点 slot" + Src->>Send: "写入跨节点 slot并 clean cache" + Src->>Remote: "UDMAPutNbi payload" + Src->>Remote: "UDMAQuiet 后写 ready[rank]" + Dst->>Remote: "invalidate 并等待 ready" + Dst->>Local: "drain 本节点 source slot" + Dst->>Remote: "drain 跨节点 source slot" + Dst->>Dst: "生成 expandX/counts/assist" +``` + +### 11.4 Combine + +combine 使用第二个 operation 区: + +```text +sendWindow = workspace + operationBytes +recvWindow = sendWindow + align32(totalBytes) +``` + +#### 同机 Direct + +1. expert output 按目标 source rank scatter 到 registered `sendWindow`。 +2. self slot复制到本地 `recvWindow`。 +3. 非 self slot 使用 UDMA put 写入目标 rank 的 `recvWindow`。 +4. ready 同步成功后,Host 读取 status。 +5. Host 启动 drain kernel,从 `recvWindow` 聚合回 `yOut`。 + +#### 跨机 Direct + +1. 本节点目标 slot 通过 `DataCopyPad` 写入目标 rank 的 IPC window。 +2. 跨节点目标 slot 使用 UDMA put 写入目标 rank 的 `recvWindow`。 +3. drain kernel 对本节点 expert rank 读取 IPC window,对跨节点 expert rank 读取 `recvWindow`。 + +```mermaid +sequenceDiagram + participant Expert as "Expert rank send phase" + participant Local as "Same-node target IPC window" + participant Send as "Combine sendWindow" + participant Remote as "Remote combine recvWindow" + participant Status as "Local workspace status" + participant Host as "Host status check" + participant Drain as "Target rank drain kernel" + + Expert->>Expert: "按 source rank scatter expertOut" + Expert->>Local: "DataCopyPad 本节点 target slot" + Expert->>Send: "写跨节点 target slot并 clean cache" + Expert->>Remote: "UDMAPutNbi payload + ready" + Expert->>Expert: "等待 remote ready" + Expert->>Status: "写 OK 或 timeout status" + Status-->>Host: "stream sync 后 D2H 读取" + Host->>Drain: "启动 combine drain kernel" + Drain->>Local: "读取本节点 expert slot" + Drain->>Remote: "invalidate并读取跨节点 expert slot" + Drain->>Drain: "聚合 yOut" +``` + +### 11.5 TP 辅助路径的边界 + +同机 Direct 的“所有 peer 走 UDMA”特指 EP dispatch/combine slot 数据面。TP group aggregation 当前仍 +要求 TP peers 位于同一节点,并使用 `peerMems[]` 中 `TileXREpTpWindowOffset(totalBytes)` 后的 IPC +辅助窗口交换 TP rows/counts。 + +因此不能把当前实现描述为“同机 Direct 完全不访问 IPC”。更准确的说法是: + +- 普通 EP peer slot payload 使用 UDMA; +- self slot本地复制; +- TP 辅助聚合仍可携带 rows/counts 经过 IPC; +- Direct Host 校验仍要求本节点 peer mappings。 + +### 11.6 Cache 与完成语义 + +UDMA source window 在 put 前执行 cache clean;receiver 在轮询 ready、slot header 和 status 前执行 +cache invalidate。put 后调用 `UDMAQuiet`,ready value 由 magic 和 step 组合,避免复用旧轮次状态。 + +Host 行为: + +- Direct dispatch kernel 启动后,`TileXREpCheckUdmaStatus` 会同步 stream 并读取 status,因此 Direct + dispatch API 当前包含一次 Host 侧同步; +- Direct combine 的发送阶段也会同步并读取 status,成功后再异步启动 drain kernel; +- Memory 路径不读取 UDMA status,也不增加这次 Direct 专用同步; +- 调用方若需要 combine drain 完成,仍应按公开 stream 语义进行同步。 + +## 12. Status 与错误传播 + +Device status 位于 `workspace + statusOffset`: + +| 值 | 常量 | 含义 | Host 返回 | +|---:|---|---|---| +| 0 | `kEpStatusOk` | 正常 | `TILEXR_SUCCESS` | +| 1 | `kEpStatusRemoteReadyTimeout` | combine remote ready 超时 | `TILEXR_ERROR_TIMEOUT` | +| 2 | `kEpStatusDispatchReadyTimeout` | dispatch ready 超时 | `TILEXR_ERROR_TIMEOUT` | +| 3 | `kEpStatusDispatchSlotTimeout` | dispatch source slot 或 TP slot 超时 | `TILEXR_ERROR_TIMEOUT` | + +Direct kernel 启动后先把 status 写为 0;超时路径写非零值并停止当前阶段。Host 当前把所有非零 device +status 统一映射为 `TILEXR_ERROR_TIMEOUT`,详细阶段需结合 device 日志和 status 原值定位。 + +其他主要错误: + +| 场景 | 行为 | +|---|---| +| mode 字符串未知 | `TILEXR_ERROR_PARA_CHECK_FAIL` | +| Memory peer mapping 缺失 | `TILEXR_ERROR_NOT_INITIALIZED` | +| Memory window 越过 100 MiB | `TILEXR_ERROR_PARA_CHECK_FAIL` | +| forced Direct capability 缺失 | `TILEXR_ERROR_NOT_INITIALIZED` | +| Direct workspace 为空 | `TILEXR_ERROR_NOT_INITIALIZED` | +| registry 无效、区域不足或 base 不匹配 | `TILEXR_ERROR_PARA_CHECK_FAIL` | +| stream sync / D2H status copy 失败 | `TILEXR_ERROR_INTERNAL` | + +任何错误都不会切换到 host-staging。 + +## 13. 内存分配策略 + +### 13.1 peerMems IPC buffer + +当前 `TileXRComm::InitMem` 仅在 `CHIP_310P3` 使用 `ACL_MEM_MALLOC_HUGE_FIRST_P2P`。Ascend950/950PR +继续使用: + +```cpp +aclrtMalloc(..., ACL_MEM_MALLOC_HUGE_FIRST) +``` + +本次实现没有把 Ascend950 peer buffer 改成 P2P malloc。现有跨机 Memory 已在该策略下通过真实 +DataCopy 验证,因此不为本功能引入额外 allocator 变更。 + +### 13.2 Direct workspace + +Direct workspace 同样可以来自普通 `aclrtMalloc` device memory,随后调用 `TileXRUDMARegister`。 +EP demo 为提高注册稳定性,会过量申请并把传入注册地址对齐到 2 MiB;这是 demo 的准备策略,当前 +`TileXRUDMARegister` Host API 本身只显式校验非空地址和非零大小,不能把 demo 对齐写成通用 API +强制契约。 + +### 13.3 两类内存不能混为一谈 + +```text +peerMems[] + communicator-owned IPC window + MEMORY payload 和本节点/TP 辅助访问 + +registered workspace + application-owned ordinary device memory + TileXRUDMARegister 后供 Direct URMA offset addressing +``` + +Direct URMA target 必须是注册 workspace,不把 `peerMems[]` 直接当作 UDMA registry region,也不在 +`InitThread` 中注册。 + +## 14. 性能数据依据 + +### 14.1 同机阈值:4 MiB + +数据文件: + +```text +D:\TileXR\full_memory_unidir_bd1.csv +D:\TileXR\full_direct_urma_unidir_bd1.csv +``` + +关键共同点位: + +| bytes | Memory GB/s | Direct URMA GB/s | URMA / Memory | 结论 | +|---:|---:|---:|---:|---| +| 1 MiB | 44.198 | 38.688 | 0.875 | Memory 更快 | +| 2 MiB | 46.835 | 44.648 | 0.953 | Memory 更快 | +| 4 MiB | 48.871 | 48.909 | 1.001 | 交叉点,基本持平 | +| 8 MiB | 49.665 | 50.675 | 1.020 | Direct URMA 开始领先 | +| 64 MiB | 39.647 | 52.796 | 1.332 | Direct URMA 明显领先 | + +因此同机阈值定为 `4 MiB`,边界归 Direct URMA。 + +### 14.2 跨机阈值:128 KiB + +数据文件: + +```text +tilexr_cross_unidir_direct_urma_bd1_8b_1g_20260728_110322.csv +tilexr_cross_unidir_memory_bd1_8b_64m_20260728_110835.csv +``` + +关键共同点位: + +| bytes | Memory GB/s | Direct URMA GB/s | URMA / Memory | 结论 | +|---:|---:|---:|---:|---| +| 64 KiB | 7.294 | 7.274 | 0.997 | Memory 略快,基本持平 | +| 128 KiB | 10.720 | 13.558 | 1.265 | Direct URMA 首次明确领先 | +| 256 KiB | 13.981 | 24.581 | 1.758 | Direct URMA 领先 | +| 1 MiB | 17.358 | 58.278 | 3.357 | Direct URMA 明显领先 | +| 4 MiB | 18.682 | 87.424 | 4.680 | Direct URMA 明显领先 | +| 64 MiB | 18.478 | 104.790 | 5.671 | Direct URMA 明显领先 | + +以上点位均为 `status=0, errors=0`,因此跨机阈值定为 `128 KiB`。 + +### 14.3 数据口径限制 + +两组 CSV 的 `bytes` 是 benchmark 单流 payload;当前 EP selector 的 `routeBytes` 是聚合 window +capacity。以 rank-2 demo 为例: + +```text +BS=1024: totalBytes=131264, slotBytes=65600 +BS=32768: totalBytes=4194496, slotBytes=2097216 +``` + +因此硬件用例跨过的是 `totalBytes` 阈值,而不是单个 remote slot payload 阈值。当前策略是明确、 +确定且已正确执行的,但其性能映射是近似的。后续若优化路由精度,这一口径差异是首要改进点。 + +## 15. 测试与验收结果 + +### 15.1 Host 与源码测试 + +已覆盖: + +- `auto/memory/direct_urma` 解析; +- 未知 mode 错误; +- 同机和跨机阈值前一字节、边界值; +- UDMA capability 缺失时 auto 回退; +- forced Direct 缺能力时报错; +- forced Memory 不依赖 registry; +- 同机 Direct launch context 校验 registry; +- Host launch 只使用 resolved `context.transport`; +- Memory launch 强制传空 workspace; +- 同机 Direct device source 包含 `useUdmaForAllPeers`; +- device auto put/get 空操作 wrapper 不再出现; +- Host runtime link block 不包含 CANN `devlib`。 +- 多节点 hybrid UDMA 只为跨节点 peer 建立 transport 资源,同节点 peer 保持 IPC。 + +本地和 141.61.49.192/223 Linux EP 单测均为 `5/5 PASS`,两端完整 Host/Bisheng 构建成功。 + +### 15.2 Auto 硬件矩阵 + +| 拓扑 | 设备 | Demo BS | routeBytes | 选择 | 结果 | +|---|---|---:|---:|---|---| +| 跨机 2-host/2-rank | 192 card0 + 223 card0 | 4 | 704 B | Memory | 两 rank dispatch/combine PASS | +| 跨机 2-host/2-rank | 192 card0 + 223 card0 | 1024 | 131264 B | Direct URMA | 两 rank dispatch/combine PASS | +| 跨机 2-host/4-rank,每节点 2 rank | 两端 card0/card1 | 4 | 1344 B | Memory | 四 rank dispatch/combine PASS | +| 跨机 2-host/4-rank,每节点 2 rank | 两端 card0/card1 | 512 | 131392 B | Direct URMA | 四 rank dispatch/combine PASS | +| 跨机 2-host/8-rank,每节点 4 rank | 两端 card0/card1/card4/card5 | 4 | 2624 B | Memory | 八 rank dispatch/combine PASS | +| 跨机 2-host/8-rank,每节点 4 rank | 两端 card0/card1/card4/card5 | 256 | 131648 B | Direct URMA | 八 rank dispatch/combine PASS | +| 跨机 2-host/16-rank,每节点 8 rank | 两端 card0-card7 | 4 | 5184 B | Memory | 十六 rank dispatch/combine PASS | +| 跨机 2-host/16-rank,每节点 8 rank | 两端 card0-card7 | 128 | 132160 B | Direct URMA | 十六 rank dispatch/combine PASS | +| 同机 1-host/2-rank | 223 card0/card1 | 4 | 704 B | Memory | 两 rank dispatch/combine PASS | +| 同机 1-host/2-rank | 223 card0/card1 | 32768 | 4194496 B | Direct URMA | 两 rank dispatch/combine PASS | + +其他已完成证据: + +- 跨机 Memory demo 默认启动 AICore push/collect,结果 `seg0=2000 seg1=2001`; +- 统一 EP Host API 强制 Memory,dispatch/combine PASS; +- 独立 registered-memory UDMA put/signal PASS; +- 统一 EP Host API 强制 Direct URMA,dispatch/combine PASS; +- 2-host/4-rank Direct 首轮在 combine step 74 暴露 ready cache-line false sharing;修复为 64B stride 后通过; +- 2-host/8-rank 使用非连续空闲卡 `0,1,4,5`,避免占用当时由其他任务使用的 192 card2/card3; +- 2-host/16-rank 初轮在 MR 注册阶段暴露冗余同节点 UDMA resource 扩展问题;按 `localRankSize` + 裁剪为仅跨节点 peer 后,forced Direct 和 auto Direct 均为 16/16 PASS; +- 2-host/16-rank auto Memory 使用 `BS=4`、`routeBytes=5184 B`,16/16 PASS;auto Direct 使用 + `BS=128`、`routeBytes=132160 B`,16/16 PASS; +- demo 的 `libascend_hal.so` 解析到真实 driver 路径; +- 测试结束后,本轮进程和端口无残留,其他用户任务未被终止。 + +### 15.3 验证结论的边界 + +已证明: + +- 两条真实数据面均可工作; +- 四个 auto 分支均按当前 `routeBytes` 契约选择正确 kernel; +- rank-2 同机与跨机、rank-4/rank-8/rank-16 跨机 dispatch/combine 数据正确; +- 每节点 2 rank、4 rank 和 8 rank 时,同机 IPC 与跨机 UDMA 混合数据面可工作; +- ready flag 使用 64B stride 后,多远端写者不会共享同一 UDMA cache line。 +- 多节点 communicator 只为跨节点 peer 建立 UDMA transport 资源,可避免 16-rank 下的冗余 MR/QP + 资源增长。 + +尚未由本矩阵证明: + +- 超过 2 台主机、每节点 rank 数不一致或全局 rank 未按节点连续排列的拓扑; +- 多轮循环、长时间 soak 和并发 communicator 的稳定性; +- 2-host/16-rank 的长时间 soak 和故障恢复行为; +- TP > 1 与同机 Direct 的完整硬件矩阵; +- 动态量化、shared expert 与阈值组合; +- 当前 `totalBytes` 指标对所有 EP shape 的性能最优性; +- 910B fallback 上的 UDMA 数据面,910B 本身不能替代 Ascend950 UDMA 验证。 + +## 16. Runtime 链接约束 + +Host library 和测试可执行文件的 RPATH/RUNPATH 禁止包含: + +```text +${ASCEND_HOME_PATH}/${ARCH}-linux/devlib +``` + +该目录中的 `libascend_hal.so` 是开发 stub,错误进入 Host runtime search path 可能导致 +`aclInit` 返回 `500000`。Host 必须解析真实 driver: + +```text +${ASCEND_DRIVER_PATH}/lib64/driver/libascend_hal.so +``` + +Ascend C kernel 的 Bisheng 链接命令可以使用 `-L.../devlib`,但该路径不能传播到 Host target 的 +RPATH/RUNPATH。构建后使用 `readelf -d` 和 `ldd` 分别检查路径声明和实际解析结果。 + +## 17. 运维与诊断约定 + +- `TILEXR_EP_DEMO_BS` 只用于 demo 调整 window 大小,不是生产路由配置。 +- demo 输出 `dispatchWindowBytes` 和 selector 结果,用于确认阈值侧。 +- `TILEXR_MEMORY_DEMO_HOST_STAGING=0` 和 `TILEXR_MEMORY_DEMO_HOST_COPY=0` 才是 Memory 数据面验收条件。 +- host-staging 日志必须明确标记 diagnostic,不计入 Memory PASS。 +- 硬件测试前后运行 `npu-smi info`;只使用空闲卡。 +- 测试必须使用独立 `TILEXR_COMM_ID` 和 barrier 端口。 +- 超时只清理本轮记录的 TileXR PID/端口,不按进程名批量终止其他任务。 +- 本地与远端源码同步只使用 Mutagen,并在构建前 flush、构建后核对关键文件 SHA-256。 + +## 18. 后续演进建议 + +按优先级排列: + +1. 统一性能数据和 selector 的大小口径,比较 `totalBytes`、`slotBytes`、最大 remote payload 三种指标。 +2. 增加 route 决策可观测性,由生产 Host 直接记录 mode、topology、routeBytes、threshold、capability + 和 resolved transport,而不是只依赖 demo 推导。 +3. 补做 2-host/16-rank soak,并扩展非均匀节点 rank 数矩阵。 +4. 增加 TP、shared expert、动态量化与 Direct URMA 的组合测试。 +5. 将环境变量 override 收敛到正式 config API,避免进程级可变状态。 +6. 保持 allocator 改动与 transport 路由解耦;只有独立证据证明 Ascend950 必须使用 P2P malloc 时, + 才单独提出并验证 allocator 变更。 diff --git a/scripts/common_env.sh b/scripts/common_env.sh index 83e08e85..14cec064 100755 --- a/scripts/common_env.sh +++ b/scripts/common_env.sh @@ -79,6 +79,19 @@ if [ ! -r "${ASCEND_DRIVER_PATH}/kernel/inc" ] && [ -d "${ASCEND_HOME_PATH}/${TI export ASCEND_DRIVER_PATH=${TILEXR_DRIVER_SHIM_HOME} fi +_tilexr_prepend_path_if_dir() { + if [ -d "$1" ]; then + case ":${PATH}:" in + *":$1:"*) ;; + *) export PATH="$1:${PATH}" ;; + esac + fi +} + +_tilexr_prepend_path_if_dir "${ASCEND_HOME_PATH}/tools/bisheng_compiler/bin" +_tilexr_prepend_path_if_dir "${ASCEND_HOME_PATH}/${TILEXR_OS_ARCH}-linux/bin" +_tilexr_prepend_path_if_dir "/usr/local/Ascend/cann-${TILEXR_CANN_VER}/tools/bisheng_compiler/bin" + export PATH=${MPI_HOME}/bin:${PATH} export PATH=${TILEXR_UTIL_HOME}/cmake/bin:${PATH} export PATH=${TILEXR_UTIL_HOME}/ccache:${TILEXR_UTIL_HOME}/ripgrep:${TILEXR_UTIL_HOME}/sshpass/bin:${PATH} diff --git a/src/comm/CMakeLists.txt b/src/comm/CMakeLists.txt index e0fd5253..2c231a93 100644 --- a/src/comm/CMakeLists.txt +++ b/src/comm/CMakeLists.txt @@ -174,6 +174,7 @@ install(FILES ${CMAKE_CURRENT_SOURCE_DIR}/../include/tilexr_sync.h ${CMAKE_CURRENT_SOURCE_DIR}/../include/tilexr_data_as_flag.h ${CMAKE_CURRENT_SOURCE_DIR}/../include/tilexr_perf_trace.h + ${CMAKE_CURRENT_SOURCE_DIR}/../include/tilexr_transport.h ${CMAKE_CURRENT_SOURCE_DIR}/../include/tilexr_udma.h ${CMAKE_CURRENT_SOURCE_DIR}/../include/tilexr_udma_reg.h ${CMAKE_CURRENT_SOURCE_DIR}/../include/tilexr_udma_types.h diff --git a/src/comm/tilexr_comm.cpp b/src/comm/tilexr_comm.cpp index 69736d0d..c5cc54d7 100644 --- a/src/comm/tilexr_comm.cpp +++ b/src/comm/tilexr_comm.cpp @@ -137,6 +137,7 @@ int TileXRComm::InitUDMA() TileXRUDMAContextOptions options {}; options.rank = rank_; options.rankSize = rankSize_; + options.localRankSize = static_cast(localRankSize_); options.devId = devId_; options.exchange = socketExchange_; options.threadMode = !uid_.empty(); diff --git a/src/comm/tools/socket/tilexr_sock_exchange.h b/src/comm/tools/socket/tilexr_sock_exchange.h index 56096b37..6ae17486 100644 --- a/src/comm/tools/socket/tilexr_sock_exchange.h +++ b/src/comm/tools/socket/tilexr_sock_exchange.h @@ -10,6 +10,7 @@ #ifndef TILEXR_SOCK_EXCHANGE_H #define TILEXR_SOCK_EXCHANGE_H +#include #include #include #include @@ -97,8 +98,14 @@ class TileXRSockExchange { template int Send(int fd, const T *sendBuf, size_t sendSize, int flag) const { - do { - auto ret = send(fd, sendBuf, sendSize, flag); + const auto *sendBytes = reinterpret_cast(sendBuf); + size_t sentBytes = 0; + while (sentBytes < sendSize) { + auto ret = send(fd, sendBytes + sentBytes, sendSize - sentBytes, flag); + if (ret > 0) { + sentBytes += static_cast(ret); + continue; + } if (ret < 0) { if (CheckErrno(errno)) { TILEXR_LOG(ERROR) << "send failed: " << strerror(errno); @@ -107,13 +114,20 @@ class TileXRSockExchange { TILEXR_LOG(DEBUG) << "Send failed: " << strerror(errno); } return ret; - } while (true); + } + return static_cast(sentBytes); } template int Recv(int fd, T *recvBuf, size_t recvSize, int flag) const { - do { - auto ret = recv(fd, recvBuf, recvSize, flag); + auto *recvBytes = reinterpret_cast(recvBuf); + size_t receivedBytes = 0; + while (receivedBytes < recvSize) { + auto ret = recv(fd, recvBytes + receivedBytes, recvSize - receivedBytes, flag); + if (ret > 0) { + receivedBytes += static_cast(ret); + continue; + } if (ret < 0) { if (CheckErrno(errno)) { TILEXR_LOG(ERROR) << "recv failed: " << strerror(errno); @@ -122,7 +136,8 @@ class TileXRSockExchange { TILEXR_LOG(DEBUG) << "recv failed: " << strerror(errno); } return ret; - } while (true); + } + return static_cast(receivedBytes); } template int ClientSendRecv(const T *sendBuf, size_t sendSize, T *recvBuf) diff --git a/src/comm/udma/tilexr_udma_context.cpp b/src/comm/udma/tilexr_udma_context.cpp index 5c1eeca5..db672add 100644 --- a/src/comm/udma/tilexr_udma_context.cpp +++ b/src/comm/udma/tilexr_udma_context.cpp @@ -58,6 +58,7 @@ int TileXRUDMAContext::Init(const TileXRUDMAContextOptions& options) TileXRUDMATransportOptions transportOptions {}; transportOptions.rank = options_.rank; transportOptions.rankSize = options_.rankSize; + transportOptions.localRankSize = options_.localRankSize; transportOptions.devId = options_.devId; transportOptions.exchange = options_.exchange; int ret = transport_->Init(transportOptions); diff --git a/src/comm/udma/tilexr_udma_context.h b/src/comm/udma/tilexr_udma_context.h index cb94bff8..eac9600b 100644 --- a/src/comm/udma/tilexr_udma_context.h +++ b/src/comm/udma/tilexr_udma_context.h @@ -30,6 +30,7 @@ using TileXRUDMACommArgsUpdateFn = int (*)(const TileXRUDMACommArgsState& state, struct TileXRUDMAContextOptions { int rank = 0; int rankSize = 0; + int localRankSize = 1; int devId = 0; bool threadMode = false; TileXRSockExchange* exchange = nullptr; diff --git a/src/comm/udma/tilexr_udma_transport.cpp b/src/comm/udma/tilexr_udma_transport.cpp index 7bafa5f1..b18035c7 100644 --- a/src/comm/udma/tilexr_udma_transport.cpp +++ b/src/comm/udma/tilexr_udma_transport.cpp @@ -324,7 +324,8 @@ int TileXRUDMATransport::Init(const TileXRUDMATransportOptions& options) if (options.rankSize <= 1) { return TILEXR_SUCCESS; } - if (options.rank < 0 || options.rank >= options.rankSize || options.exchange == nullptr) { + if (options.rank < 0 || options.rank >= options.rankSize || options.localRankSize <= 0 || + options.localRankSize > options.rankSize || options.exchange == nullptr) { return TILEXR_ERROR_PARA_CHECK_FAIL; } options_ = options; @@ -459,7 +460,7 @@ int TileXRUDMATransport::BuildRoutes() std::vector localRouteByPeer(options_.rankSize, -1); for (int peer = 0; peer < options_.rankSize; ++peer) { - if (peer == options_.rank) { + if (!UsesUDMAPeer(peer)) { continue; } uint32_t localEid = devEids[0].eidIndex; @@ -480,7 +481,7 @@ int TileXRUDMATransport::BuildRoutes() } for (int peer = 0; peer < options_.rankSize; ++peer) { - if (peer == options_.rank) { + if (!UsesUDMAPeer(peer)) { continue; } int32_t remoteEid = allRouteByPeer[peer * options_.rankSize + options_.rank]; @@ -720,6 +721,17 @@ uint32_t TileXRUDMATransport::FallbackLocalEid() const return 0; } +bool TileXRUDMATransport::UsesUDMAPeer(int peer) const +{ + if (peer < 0 || peer >= options_.rankSize || peer == options_.rank) { + return false; + } + if (options_.localRankSize >= options_.rankSize) { + return true; + } + return peer / options_.localRankSize != options_.rank / options_.localRankSize; +} + int TileXRUDMATransport::RefreshUDMAInfo() { if (eidCount_ == 0 || states_.empty()) { @@ -777,7 +789,8 @@ int TileXRUDMATransport::RefreshUDMAInfo() for (int rank = 0; rank < options_.rankSize; ++rank) { uint32_t localEid = fallbackEid; uint32_t remoteEid = fallbackEid; - if (rank != options_.rank) { + const bool usesUDMA = UsesUDMAPeer(rank); + if (usesUDMA) { localEid = peerLocalEid_[rank]; remoteEid = peerRemoteEid_[rank]; } @@ -795,7 +808,7 @@ int TileXRUDMATransport::RefreshUDMAInfo() if (localMemIt != localMemInfoByEid_.end()) { mem[rank] = localMemIt->second; } - } else { + } else if (usesUDMA) { mem[rank] = allMem[rank * eidCount_ + remoteEid]; mem[rank].tpn = state.tpnList[rank]; } @@ -832,14 +845,23 @@ int TileXRUDMATransport::RegisterMemory(GM_ADDR localPtr, size_t bytes) } int ret = RegisterMemoryOnContexts(localPtr, bytes); if (ret != TILEXR_SUCCESS) { + TILEXR_LOG(ERROR) << "TileXR UDMA local memory registration failed, rank=" << options_.rank + << ", bytes=" << bytes << ", ret=" << ret; return ret; } registeredPtr_ = localPtr; ret = ExchangeAndImportMemory(); if (ret != TILEXR_SUCCESS) { + TILEXR_LOG(ERROR) << "TileXR UDMA remote memory exchange/import failed, rank=" << options_.rank + << ", bytes=" << bytes << ", ret=" << ret; return ret; } - return RefreshUDMAInfo(); + ret = RefreshUDMAInfo(); + if (ret != TILEXR_SUCCESS) { + TILEXR_LOG(ERROR) << "TileXR UDMA info refresh failed after memory registration, rank=" << options_.rank + << ", bytes=" << bytes << ", ret=" << ret; + } + return ret; } int TileXRUDMATransport::RegisterMemoryOnContexts(GM_ADDR localPtr, size_t bytes) @@ -863,6 +885,9 @@ int TileXRUDMATransport::RegisterMemoryOnContexts(GM_ADDR localPtr, size_t bytes void* lmemHandle = nullptr; int ret = loader_.RaCtxLmemRegister(ctxEntry.second, &mrInfo, &lmemHandle); if (ret != 0 || lmemHandle == nullptr) { + TILEXR_LOG(ERROR) << "TileXR UDMA RaCtxLmemRegister failed, rank=" << options_.rank + << ", eid=" << eidIndex << ", bytes=" << bytes << ", ret=" << ret + << ", handle=" << lmemHandle; return TILEXR_ERROR_INTERNAL; } @@ -903,6 +928,8 @@ int TileXRUDMATransport::ExchangeAndImportMemory() std::vector allCounts(options_.rankSize); int ret = options_.exchange->AllGather(&localCount, 1, allCounts.data()); if (ret != TILEXR_SUCCESS) { + TILEXR_LOG(ERROR) << "TileXR UDMA memory-count AllGather failed, rank=" << options_.rank + << ", localCount=" << localCount << ", ret=" << ret; return ret; } const uint32_t maxCount = *std::max_element(allCounts.begin(), allCounts.end()); @@ -927,12 +954,14 @@ int TileXRUDMATransport::ExchangeAndImportMemory() std::vector all(options_.rankSize * maxCount); ret = options_.exchange->AllGather(local.data(), local.size(), all.data()); if (ret != TILEXR_SUCCESS) { + TILEXR_LOG(ERROR) << "TileXR UDMA memory-info AllGather failed, rank=" << options_.rank + << ", maxCount=" << maxCount << ", ret=" << ret; return ret; } remoteMemHandles_.assign(options_.rankSize, nullptr); for (int peer = 0; peer < options_.rankSize; ++peer) { - if (peer == options_.rank) { + if (!UsesUDMAPeer(peer)) { continue; } const uint32_t remoteEid = peerRemoteEid_[peer]; @@ -945,6 +974,9 @@ int TileXRUDMATransport::ExchangeAndImportMemory() } } if (remote == nullptr) { + TILEXR_LOG(ERROR) << "TileXR UDMA remote memory info missing, rank=" << options_.rank + << ", peer=" << peer << ", remoteEid=" << remoteEid + << ", peerCount=" << allCounts[peer]; return TILEXR_ERROR_INTERNAL; } const uint32_t localEid = peerLocalEid_[peer]; @@ -956,6 +988,10 @@ int TileXRUDMATransport::ExchangeAndImportMemory() void* remoteHandle = nullptr; ret = loader_.RaCtxRmemImport(ctxHandleByEid_[localEid], &importInfo, &remoteHandle); if (ret != 0 || remoteHandle == nullptr) { + TILEXR_LOG(ERROR) << "TileXR UDMA RaCtxRmemImport failed, rank=" << options_.rank + << ", peer=" << peer << ", localEid=" << localEid + << ", remoteEid=" << remoteEid << ", ret=" << ret + << ", handle=" << remoteHandle; return TILEXR_ERROR_INTERNAL; } remoteMemHandles_[peer] = remoteHandle; diff --git a/src/comm/udma/tilexr_udma_transport.h b/src/comm/udma/tilexr_udma_transport.h index 0a787d12..2a877573 100644 --- a/src/comm/udma/tilexr_udma_transport.h +++ b/src/comm/udma/tilexr_udma_transport.h @@ -27,6 +27,7 @@ class TileXRSockExchange; struct TileXRUDMATransportOptions { int rank = 0; int rankSize = 0; + int localRankSize = 1; int devId = 0; TileXRSockExchange* exchange = nullptr; }; @@ -63,6 +64,7 @@ class TileXRUDMATransport { void CleanupMemory(); void CleanupContexts(); uint32_t FallbackLocalEid() const; + bool UsesUDMAPeer(int peer) const; TileXRHccpLoader loader_; TileXRUDMATransportOptions options_ {}; diff --git a/src/ep/CMakeLists.txt b/src/ep/CMakeLists.txt index c2fc93d5..dbe4142d 100644 --- a/src/ep/CMakeLists.txt +++ b/src/ep/CMakeLists.txt @@ -138,6 +138,7 @@ add_custom_target(tilexr_ep_combine_kernel ALL DEPENDS "${TILEXR_EP_COMBINE_KERN add_library(tilexr-ep SHARED host/ep_layout.cpp host/ep_dispatch_host.cpp + host/ep_transport_route.cpp host/ep_launch_context.cpp host/ep_kernel_launch.cpp host/tilexr_ep_dispatch.cpp @@ -158,7 +159,6 @@ target_link_directories(tilexr-ep ${CMAKE_CURRENT_BINARY_DIR} ${ASCEND_DRIVER_PATH}/lib64/driver ${ASCEND_HOME_PATH}/${ARCH}-linux/lib64 - ${ASCEND_HOME_PATH}/${ARCH}-linux/devlib ) target_link_libraries(tilexr-ep diff --git a/src/ep/common/ep_window.h b/src/ep/common/ep_window.h index 4e230640..cd61ed37 100644 --- a/src/ep/common/ep_window.h +++ b/src/ep/common/ep_window.h @@ -6,6 +6,7 @@ namespace TileXREp { constexpr int64_t kEpWindowAlignmentBytes = 32; +constexpr int64_t kEpUdmaReadyStrideBytes = 64; constexpr int64_t kEpAssistTupleInts = 4; constexpr int64_t kEpWindowHeaderBytes = 64; constexpr int64_t kEpSrcSlotHeaderBytes = 64; @@ -18,6 +19,8 @@ constexpr int32_t kEpStepCombineGatewayReady = 76; constexpr int32_t kEpStepCombineRelayReady = 77; constexpr int64_t kEpStatusOk = 0; constexpr int64_t kEpStatusRemoteReadyTimeout = 1; +constexpr int64_t kEpStatusDispatchReadyTimeout = 2; +constexpr int64_t kEpStatusDispatchSlotTimeout = 3; constexpr uint32_t kEpWindowMagic = 0x54584550U; struct EpWindowHeader { diff --git a/src/ep/host/ep_dispatch_host.cpp b/src/ep/host/ep_dispatch_host.cpp index c5a708c7..8fdc5ec2 100644 --- a/src/ep/host/ep_dispatch_host.cpp +++ b/src/ep/host/ep_dispatch_host.cpp @@ -7,11 +7,6 @@ namespace TileXREp { namespace { -bool TileXREpUsesCrossNodeComm(const TileXR::CommArgs &commArgs) -{ - return commArgs.localRankSize > 0 && commArgs.localRankSize < commArgs.rankSize; -} - int64_t TileXREpEffectiveTpWorldSize(int64_t tpWorldSize) { return tpWorldSize == 0 ? 1 : tpWorldSize; @@ -118,18 +113,6 @@ int TileXREpValidateDispatchConfig(const EpDispatchParams ¶ms, const TileXR: return TileXR::TILEXR_ERROR_PARA_CHECK_FAIL; } - for (int rank = 0; rank < commArgs.rankSize; ++rank) { - if (commArgs.peerMems[rank] == nullptr) { - return TileXR::TILEXR_ERROR_NOT_INITIALIZED; - } - } - - if (TileXREpUsesCrossNodeComm(commArgs) && - (params.workspace == nullptr || (commArgs.extraFlag & TileXR::ExtraFlag::UDMA) == 0 || - commArgs.udmaInfoPtr == nullptr || commArgs.udmaRegistryPtr == nullptr)) { - return TileXR::TILEXR_ERROR_NOT_INITIALIZED; - } - int ret = TileXREpValidateDispatchV2Config(params, commArgs); if (ret != TileXR::TILEXR_SUCCESS) { return ret; @@ -169,18 +152,6 @@ int TileXREpValidateCombineConfig(const EpCombineParams ¶ms, const TileXR::C commArgs.rank < 0 || commArgs.rank >= commArgs.rankSize) { return TileXR::TILEXR_ERROR_PARA_CHECK_FAIL; } - if (TileXREpUsesCrossNodeComm(commArgs) && - (params.workspace == nullptr || (commArgs.extraFlag & TileXR::ExtraFlag::UDMA) == 0 || - commArgs.udmaInfoPtr == nullptr || commArgs.udmaRegistryPtr == nullptr)) { - return TileXR::TILEXR_ERROR_NOT_INITIALIZED; - } - - for (int rank = 0; rank < commArgs.rankSize; ++rank) { - if (commArgs.peerMems[rank] == nullptr) { - return TileXR::TILEXR_ERROR_NOT_INITIALIZED; - } - } - return TileXREpBuildWindowConfig(commArgs.rankSize, params.bs, params.h, params.topK, params.moeExpertNum, params.dtype, window); } diff --git a/src/ep/host/ep_dispatch_host.h b/src/ep/host/ep_dispatch_host.h index 903e1a79..294e3a11 100644 --- a/src/ep/host/ep_dispatch_host.h +++ b/src/ep/host/ep_dispatch_host.h @@ -5,6 +5,7 @@ #include "acl/acl_base.h" #include "ep_layout.h" +#include "ep_transport_route.h" #include "tilexr_api.h" namespace TileXREp { @@ -64,6 +65,7 @@ struct EpHostLaunchContext { TileXR::CommArgs *hostArgs = nullptr; GM_ADDR devArgs = nullptr; EpWindowConfig window {}; + TileXR::TileXRTransportKind transport = TileXR::TileXRTransportKind::MEMORY; }; int TileXREpValidateBasicDispatchParams(const EpDispatchParams ¶ms); diff --git a/src/ep/host/ep_kernel_launch.cpp b/src/ep/host/ep_kernel_launch.cpp index f7095626..81fcfc92 100644 --- a/src/ep/host/ep_kernel_launch.cpp +++ b/src/ep/host/ep_kernel_launch.cpp @@ -47,10 +47,10 @@ namespace TileXREp { namespace { -bool TileXREpUsesCrossNodeKernel(const EpHostLaunchContext &context) +bool TileXREpUsesDirectUdmaKernel(const EpHostLaunchContext &context) { - return context.hostArgs != nullptr && context.hostArgs->localRankSize > 0 && - context.hostArgs->localRankSize < context.hostArgs->rankSize; + return context.transport == TileXR::TileXRTransportKind::DIRECT_URMA && context.hostArgs != nullptr && + context.hostArgs->rankSize > 1 && context.window.slotBytes > 0; } int64_t TileXREpUdmaStatusOffset(int64_t totalBytes, int64_t rankSize, int64_t slotBytes) @@ -62,6 +62,29 @@ int64_t TileXREpUdmaStatusOffset(int64_t totalBytes, int64_t rankSize, int64_t s return operationBytes * 2; } +int TileXREpCheckUdmaStatus(aclrtStream stream, void *workspace, const EpWindowConfig &window) +{ + if (workspace == nullptr) { + return TileXR::TILEXR_ERROR_NOT_INITIALIZED; + } + aclError aclRet = aclrtSynchronizeStream(stream); + if (aclRet != ACL_SUCCESS) { + return TileXR::TILEXR_ERROR_INTERNAL; + } + + const int64_t statusOffset = TileXREpUdmaStatusOffset(window.totalBytes, window.rankSize, window.slotBytes); + if (statusOffset == TileXR::TILEXR_INVALID_VALUE) { + return TileXR::TILEXR_ERROR_PARA_CHECK_FAIL; + } + uint64_t status = TileXREp::kEpStatusOk; + aclRet = aclrtMemcpy(&status, sizeof(status), static_cast(workspace) + statusOffset, + sizeof(status), ACL_MEMCPY_DEVICE_TO_HOST); + if (aclRet != ACL_SUCCESS) { + return TileXR::TILEXR_ERROR_INTERNAL; + } + return status == TileXREp::kEpStatusOk ? TileXR::TILEXR_SUCCESS : TileXR::TILEXR_ERROR_TIMEOUT; +} + } // namespace int TileXREpLaunchDispatchKernel(const EpDispatchParams ¶ms, const EpHostLaunchContext &context) @@ -73,7 +96,7 @@ int TileXREpLaunchDispatchKernel(const EpDispatchParams ¶ms, const EpHostLau } constexpr uint32_t kMvpBlockDim = 1; - if (TileXREpUsesCrossNodeKernel(context)) { + if (TileXREpUsesDirectUdmaKernel(context)) { launch_tilexr_ep_dispatch_cross_node_kernel(kMvpBlockDim, params.stream, context.devArgs, static_cast(params.x), reinterpret_cast(params.expertIds), static_cast(params.scales), reinterpret_cast(params.xActiveMask), @@ -87,14 +110,16 @@ int TileXREpLaunchDispatchKernel(const EpDispatchParams ¶ms, const EpHostLau context.window.assistBytesPerSlot, context.window.slotBytes, context.window.totalBytes, params.expertTokenNumsType, params.sharedExpertNum, params.sharedExpertRankNum, params.quantMode, params.tpWorldSize, params.tpRankId, magic); + return TileXREpCheckUdmaStatus(params.stream, params.workspace, context.window); } else { + GM_ADDR memoryWorkspace = nullptr; launch_tilexr_ep_dispatch_kernel(kMvpBlockDim, params.stream, context.devArgs, static_cast(params.x), reinterpret_cast(params.expertIds), static_cast(params.scales), reinterpret_cast(params.xActiveMask), static_cast(params.expandXOut), static_cast(params.dynamicScalesOut), reinterpret_cast(params.expertTokenNumsOut), reinterpret_cast(params.epRecvCountsOut), reinterpret_cast(params.tpRecvCountsOut), reinterpret_cast(params.assistInfoForCombineOut), - static_cast(params.workspace), params.bs, params.h, params.topK, params.moeExpertNum, context.window.dtypeBytes, + memoryWorkspace, params.bs, params.h, params.topK, params.moeExpertNum, context.window.dtypeBytes, context.window.maxRoutesPerSrc, context.window.rowBytes, context.window.payloadRowBytes, context.window.payloadBytesPerSlot, context.window.assistBytesPerSlot, context.window.slotBytes, context.window.totalBytes, params.expertTokenNumsType, params.sharedExpertNum, params.sharedExpertRankNum, @@ -112,7 +137,7 @@ int TileXREpLaunchCombineKernel(const EpCombineParams ¶ms, const EpHostLaunc } constexpr uint32_t kMvpBlockDim = 1; - if (TileXREpUsesCrossNodeKernel(context)) { + if (TileXREpUsesDirectUdmaKernel(context)) { launch_tilexr_ep_combine_cross_node_kernel(kMvpBlockDim, params.stream, context.devArgs, static_cast(params.expertOut), reinterpret_cast(params.assistInfoForCombine), reinterpret_cast(params.epRecvCounts), static_cast(params.yOut), @@ -120,23 +145,9 @@ int TileXREpLaunchCombineKernel(const EpCombineParams ¶ms, const EpHostLaunc static_cast(params.dtype), context.window.dtypeBytes, context.window.maxRoutesPerSrc, context.window.rowBytes, context.window.payloadBytesPerSlot, context.window.assistBytesPerSlot, context.window.slotBytes, context.window.totalBytes, magic); - aclError aclRet = aclrtSynchronizeStream(params.stream); - if (aclRet != ACL_SUCCESS) { - return TileXR::TILEXR_ERROR_INTERNAL; - } - uint64_t status = TileXREp::kEpStatusOk; - const int64_t statusOffset = TileXREpUdmaStatusOffset( - context.window.totalBytes, context.window.rankSize, context.window.slotBytes); - if (statusOffset == TileXR::TILEXR_INVALID_VALUE) { - return TileXR::TILEXR_ERROR_PARA_CHECK_FAIL; - } - aclRet = aclrtMemcpy(&status, sizeof(status), static_cast(params.workspace) + statusOffset, - sizeof(status), ACL_MEMCPY_DEVICE_TO_HOST); - if (aclRet != ACL_SUCCESS) { - return TileXR::TILEXR_ERROR_INTERNAL; - } - if (status != TileXREp::kEpStatusOk) { - return TileXR::TILEXR_ERROR_TIMEOUT; + const int statusRet = TileXREpCheckUdmaStatus(params.stream, params.workspace, context.window); + if (statusRet != TileXR::TILEXR_SUCCESS) { + return statusRet; } launch_tilexr_ep_combine_cross_node_drain_kernel(kMvpBlockDim, params.stream, context.devArgs, static_cast(params.yOut), static_cast(params.workspace), params.bs, params.h, diff --git a/src/ep/host/ep_launch_context.cpp b/src/ep/host/ep_launch_context.cpp index a48a579d..d5974f71 100644 --- a/src/ep/host/ep_launch_context.cpp +++ b/src/ep/host/ep_launch_context.cpp @@ -2,6 +2,7 @@ #include +#include "ep_transport_route.h" #include "ep_window.h" #include "tilexr_udma_reg.h" #include "tilexr_types.h" @@ -9,32 +10,9 @@ namespace TileXREp { namespace { -bool TileXREpUsesCrossNodeComm(const TileXR::CommArgs &commArgs) -{ - return commArgs.localRankSize > 0 && commArgs.localRankSize < commArgs.rankSize; -} - -int64_t TileXREpDispatchWorkspaceBytes(const EpWindowConfig &window, int64_t tpWorldSize) -{ - const int64_t alignedTotal = TileXREpAlignUp(window.totalBytes, kEpWindowAlignmentBytes); - const int64_t effectiveTpWorldSize = tpWorldSize == 0 ? 1 : tpWorldSize; - if (alignedTotal == TileXR::TILEXR_INVALID_VALUE || effectiveTpWorldSize <= 0 || - effectiveTpWorldSize > INT64_MAX - 2) { - return TileXR::TILEXR_INVALID_VALUE; - } - const int64_t factor = effectiveTpWorldSize + 2; - if (alignedTotal != 0 && factor > INT64_MAX / alignedTotal) { - return TileXR::TILEXR_INVALID_VALUE; - } - return alignedTotal * factor; -} - int ValidateRegisteredWorkspace( TileXRCommPtr comm, const TileXR::CommArgs &commArgs, void *workspace, int64_t requiredBytes) { - if (!TileXREpUsesCrossNodeComm(commArgs)) { - return TileXR::TILEXR_SUCCESS; - } if (workspace == nullptr || requiredBytes <= 0) { return TileXR::TILEXR_ERROR_NOT_INITIALIZED; } @@ -56,6 +34,37 @@ int ValidateRegisteredWorkspace( return TileXR::TILEXR_SUCCESS; } +int ValidateMemoryWindow(int64_t requiredBytes) +{ + if (requiredBytes <= 0 || requiredBytes > TileXR::IPC_BUFF_MAX_SIZE) { + return TileXR::TILEXR_ERROR_PARA_CHECK_FAIL; + } + return TileXR::TILEXR_SUCCESS; +} + +int ValidatePeerMems( + const TileXR::CommArgs &commArgs, TileXR::TileXRTransportKind transport) +{ + int beginRank = 0; + int endRank = commArgs.rankSize; + if (transport == TileXR::TileXRTransportKind::DIRECT_URMA) { + if (commArgs.localRankSize <= 0) { + return TileXR::TILEXR_ERROR_NOT_INITIALIZED; + } + beginRank = (commArgs.rank / commArgs.localRankSize) * commArgs.localRankSize; + endRank = beginRank + commArgs.localRankSize; + if (endRank > commArgs.rankSize) { + endRank = commArgs.rankSize; + } + } + for (int rank = beginRank; rank < endRank; ++rank) { + if (commArgs.peerMems[rank] == nullptr) { + return TileXR::TILEXR_ERROR_NOT_INITIALIZED; + } + } + return TileXR::TILEXR_SUCCESS; +} + } // namespace int TileXREpPrepareLaunchContext(const EpDispatchParams ¶ms, EpHostLaunchContext *context) @@ -89,8 +98,25 @@ int TileXREpPrepareLaunchContext(const EpDispatchParams ¶ms, EpHostLaunchCon *context = EpHostLaunchContext {}; return ret; } - const int64_t dispatchWorkspaceBytes = TileXREpDispatchWorkspaceBytes(context->window, params.tpWorldSize); - ret = ValidateRegisteredWorkspace(params.comm, *context->hostArgs, params.workspace, dispatchWorkspaceBytes); + + ret = TileXREpResolveTransportForWorkspaceFromEnv(*context->hostArgs, + static_cast(context->window.totalBytes), params.workspace, &context->transport); + if (ret != TileXR::TILEXR_SUCCESS) { + *context = EpHostLaunchContext {}; + return ret; + } + + ret = ValidatePeerMems(*context->hostArgs, context->transport); + if (ret != TileXR::TILEXR_SUCCESS) { + *context = EpHostLaunchContext {}; + return ret; + } + + const int64_t dispatchWorkspaceBytes = TileXREpUdmaRequiredWorkspaceBytes( + context->window.totalBytes, context->window.rankSize, context->window.slotBytes); + ret = context->transport == TileXR::TileXRTransportKind::DIRECT_URMA ? + ValidateRegisteredWorkspace(params.comm, *context->hostArgs, params.workspace, dispatchWorkspaceBytes) : + ValidateMemoryWindow(context->window.totalBytes); if (ret != TileXR::TILEXR_SUCCESS) { *context = EpHostLaunchContext {}; return ret; @@ -129,9 +155,25 @@ int TileXREpPrepareCombineLaunchContext(const EpCombineParams ¶ms, EpHostLau *context = EpHostLaunchContext {}; return ret; } + + ret = TileXREpResolveTransportForWorkspaceFromEnv(*context->hostArgs, + static_cast(context->window.totalBytes), params.workspace, &context->transport); + if (ret != TileXR::TILEXR_SUCCESS) { + *context = EpHostLaunchContext {}; + return ret; + } + + ret = ValidatePeerMems(*context->hostArgs, context->transport); + if (ret != TileXR::TILEXR_SUCCESS) { + *context = EpHostLaunchContext {}; + return ret; + } + const int64_t combineWorkspaceBytes = TileXREpUdmaRequiredWorkspaceBytes( context->window.totalBytes, context->window.rankSize, context->window.slotBytes); - ret = ValidateRegisteredWorkspace(params.comm, *context->hostArgs, params.workspace, combineWorkspaceBytes); + ret = context->transport == TileXR::TileXRTransportKind::DIRECT_URMA ? + ValidateRegisteredWorkspace(params.comm, *context->hostArgs, params.workspace, combineWorkspaceBytes) : + ValidateMemoryWindow(context->window.totalBytes); if (ret != TileXR::TILEXR_SUCCESS) { *context = EpHostLaunchContext {}; return ret; diff --git a/src/ep/host/ep_layout.cpp b/src/ep/host/ep_layout.cpp index ea78e995..409207cd 100644 --- a/src/ep/host/ep_layout.cpp +++ b/src/ep/host/ep_layout.cpp @@ -9,6 +9,9 @@ namespace TileXREp { namespace { +static_assert(kEpUdmaReadyStrideBytes == TileXR::TILEXR_UDMA_CACHE_LINE_SIZE, + "EP UDMA ready slots must match the UDMA cache-line size"); + bool MulInt64(int64_t lhs, int64_t rhs, int64_t *out) { if (out == nullptr || lhs < 0 || rhs < 0) { @@ -84,7 +87,7 @@ int64_t TileXREpUdmaOperationBytes(int64_t totalBytes, int64_t rankSize, int64_t int64_t readyBytes = 0; int64_t doubleTotal = 0; int64_t readyOffset = 0; - if (!MulInt64(static_cast(rankSize), static_cast(sizeof(uint64_t)), &readyBytes) || + if (!MulInt64(static_cast(rankSize), kEpUdmaReadyStrideBytes, &readyBytes) || !MulInt64(alignedTotal, 2, &doubleTotal) || !AddInt64(doubleTotal, readyBytes, &readyOffset)) { return TileXR::TILEXR_INVALID_VALUE; diff --git a/src/ep/host/ep_transport_route.cpp b/src/ep/host/ep_transport_route.cpp new file mode 100644 index 00000000..e979de33 --- /dev/null +++ b/src/ep/host/ep_transport_route.cpp @@ -0,0 +1,86 @@ +#include "ep_transport_route.h" + +#include +#include + +#include "tilexr_types.h" + +namespace TileXREp { + +int TileXREpParseTransportMode(const char *value, EpTransportMode *mode) +{ + if (mode == nullptr) { + return TileXR::TILEXR_ERROR_PARA_CHECK_FAIL; + } + if (value == nullptr || value[0] == '\0' || std::strcmp(value, "auto") == 0) { + *mode = EpTransportMode::AUTO; + return TileXR::TILEXR_SUCCESS; + } + if (std::strcmp(value, "memory") == 0) { + *mode = EpTransportMode::MEMORY; + return TileXR::TILEXR_SUCCESS; + } + if (std::strcmp(value, "direct_urma") == 0) { + *mode = EpTransportMode::DIRECT_URMA; + return TileXR::TILEXR_SUCCESS; + } + return TileXR::TILEXR_ERROR_PARA_CHECK_FAIL; +} + +bool TileXREpShouldRegisterWorkspace(EpTransportMode mode, const TileXR::CommArgs &args) +{ + return mode != EpTransportMode::MEMORY && TileXR::TileXRDirectUrmaCapable(&args); +} + +int TileXREpResolveTransport(EpTransportMode mode, const TileXR::CommArgs &args, uint64_t bytes, + TileXR::TileXRTransportKind *transport) +{ + if (transport == nullptr) { + return TileXR::TILEXR_ERROR_PARA_CHECK_FAIL; + } + + switch (mode) { + case EpTransportMode::AUTO: { + *transport = TileXR::TileXRSelectAutoTransport(&args, bytes); + return TileXR::TILEXR_SUCCESS; + } + case EpTransportMode::MEMORY: + *transport = TileXR::TileXRTransportKind::MEMORY; + return TileXR::TILEXR_SUCCESS; + case EpTransportMode::DIRECT_URMA: + if (!TileXR::TileXRDirectUrmaAvailable(&args)) { + return TileXR::TILEXR_ERROR_NOT_INITIALIZED; + } + *transport = TileXR::TileXRTransportKind::DIRECT_URMA; + return TileXR::TILEXR_SUCCESS; + default: + return TileXR::TILEXR_ERROR_PARA_CHECK_FAIL; + } +} + +int TileXREpResolveTransportFromEnv(const TileXR::CommArgs &args, uint64_t bytes, + TileXR::TileXRTransportKind *transport) +{ + EpTransportMode mode = EpTransportMode::AUTO; + int ret = TileXREpParseTransportMode(std::getenv("TILEXR_TRANSPORT_MODE"), &mode); + if (ret != TileXR::TILEXR_SUCCESS) { + return ret; + } + return TileXREpResolveTransport(mode, args, bytes, transport); +} + +int TileXREpResolveTransportForWorkspaceFromEnv(const TileXR::CommArgs &args, uint64_t bytes, + const void *workspace, TileXR::TileXRTransportKind *transport) +{ + EpTransportMode mode = EpTransportMode::AUTO; + int ret = TileXREpParseTransportMode(std::getenv("TILEXR_TRANSPORT_MODE"), &mode); + if (ret != TileXR::TILEXR_SUCCESS) { + return ret; + } + if (mode == EpTransportMode::AUTO && workspace == nullptr) { + mode = EpTransportMode::MEMORY; + } + return TileXREpResolveTransport(mode, args, bytes, transport); +} + +} // namespace TileXREp diff --git a/src/ep/host/ep_transport_route.h b/src/ep/host/ep_transport_route.h new file mode 100644 index 00000000..57041ca2 --- /dev/null +++ b/src/ep/host/ep_transport_route.h @@ -0,0 +1,28 @@ +#ifndef TILEXR_EP_HOST_EP_TRANSPORT_ROUTE_H +#define TILEXR_EP_HOST_EP_TRANSPORT_ROUTE_H + +#include + +#include "comm_args.h" +#include "tilexr_transport.h" + +namespace TileXREp { + +enum class EpTransportMode : uint8_t { + AUTO = 0, + MEMORY = 1, + DIRECT_URMA = 2, +}; + +int TileXREpParseTransportMode(const char *value, EpTransportMode *mode); +bool TileXREpShouldRegisterWorkspace(EpTransportMode mode, const TileXR::CommArgs &args); +int TileXREpResolveTransport(EpTransportMode mode, const TileXR::CommArgs &args, uint64_t bytes, + TileXR::TileXRTransportKind *transport); +int TileXREpResolveTransportFromEnv(const TileXR::CommArgs &args, uint64_t bytes, + TileXR::TileXRTransportKind *transport); +int TileXREpResolveTransportForWorkspaceFromEnv(const TileXR::CommArgs &args, uint64_t bytes, + const void *workspace, TileXR::TileXRTransportKind *transport); + +} // namespace TileXREp + +#endif // TILEXR_EP_HOST_EP_TRANSPORT_ROUTE_H diff --git a/src/ep/kernels/tilexr_ep_combine_helpers.h b/src/ep/kernels/tilexr_ep_combine_helpers.h index 60ea8a61..63ad10fe 100644 --- a/src/ep/kernels/tilexr_ep_combine_helpers.h +++ b/src/ep/kernels/tilexr_ep_combine_helpers.h @@ -36,13 +36,20 @@ __aicore__ inline int32_t TileXREpGetCombineTopKId(const TileXREp::EpAssistTuple __aicore__ inline int32_t TileXREpWaitDispatchSlotReady( GM_ADDR slotGM, int32_t srcRank, int64_t magic, AscendC::TBuf &tBuf) { + int64_t retries = 0; while (true) { + TileXREpInvalidateLocalCacheLines(slotGM, TileXREp::kEpSrcSlotHeaderBytes); const uint64_t packed = LoadUint64FromGm(slotGM, tBuf); const int32_t slotSrcRank = static_cast((packed >> 32) & 0xffffffffULL); const uint64_t slotMagic = LoadUint64FromGm(slotGM + 3 * static_cast(sizeof(uint64_t)), tBuf); if (slotSrcRank == srcRank && slotMagic == static_cast(magic)) { return static_cast(packed & 0xffffffffULL); } + ++retries; + if (retries >= static_cast(TileXR::TILEXR_UDMA_MAX_RETRY_TIMES)) { + AscendC::printf("tilexr_ep_dispatch_slot_ready timeout src %d\n", srcRank); + return -1; + } } } @@ -85,10 +92,13 @@ __aicore__ inline int64_t TileXREpDrainSourceWindow(GM_ADDR sourceWindow, int32_ int64_t sharedExpertRankNum, GM_ADDR expandXOutGM, GM_ADDR dynamicScalesOutGM, GM_ADDR expertTokenNumsOutGM, int64_t expertTokenNumsType, int64_t magic, __gm__ int32_t *epRecvCountsOut, __gm__ int32_t *tpRecvCountsOut, __gm__ TileXREp::EpAssistTuple *localAssistBase, int64_t outRecord, GM_ADDR tpPublishWindow, - int32_t tpPublishRank, AscendC::TBuf &tBuf) + int32_t tpPublishRank, int32_t *sourceCountOut, AscendC::TBuf &tBuf) { const int64_t count = TileXREpWaitDispatchSlotReady(sourceWindow + SlotOffset(slotRank, slotBytes), srcRank, magic, tBuf); + if (sourceCountOut != nullptr) { + *sourceCountOut = static_cast(count); + } epRecvCountsOut[srcRank] = static_cast(count); if (tpRecvCountsOut != nullptr) { tpRecvCountsOut[srcRank] = static_cast(count); diff --git a/src/ep/kernels/tilexr_ep_combine_kernel.cpp b/src/ep/kernels/tilexr_ep_combine_kernel.cpp index d21b554f..3a7181dd 100644 --- a/src/ep/kernels/tilexr_ep_combine_kernel.cpp +++ b/src/ep/kernels/tilexr_ep_combine_kernel.cpp @@ -184,6 +184,8 @@ extern "C" __global__ __aicore__ void tilexr_ep_combine_cross_node_kernel(GM_ADD const int32_t rank = args->rank; const int32_t rankSize = args->rankSize; const int32_t localRankSize = args->localRankSize; + const bool useUdmaForAllPeers = localRankSize == rankSize; + const int32_t effectiveLocalRankSize = useUdmaForAllPeers ? 1 : localRankSize; if (rankSize <= 0 || rankSize > TileXR::TILEXR_MAX_RANK_SIZE || rank < 0 || rank >= rankSize || localRankSize <= 0 || localRankSize > rankSize || !TileXR::UDMARegistryEnabled(args) || !IsValidShape(bs, h, topK, moeExpertNum, dtypeBytes, maxRoutesPerSrc, rowBytes, payloadBytesPerSlot, @@ -192,13 +194,8 @@ extern "C" __global__ __aicore__ void tilexr_ep_combine_cross_node_kernel(GM_ADD } GM_ADDR shareAddrs[TileXR::TILEXR_MAX_RANK_SIZE]; - AscendC::GlobalTensor peerMems; - peerMems.SetGlobalBuffer(&(args->peerMems[0]), TileXR::TILEXR_MAX_RANK_SIZE); - for (int32_t peer = 0; peer < rankSize; ++peer) { - shareAddrs[peer] = peerMems.GetValue(peer); - if (shareAddrs[peer] == nullptr) { - return; - } + if (!TileXREpLoadLocalPeerMems(args, shareAddrs, rank, rankSize, localRankSize)) { + return; } AscendC::TPipe pipe; @@ -217,8 +214,8 @@ extern "C" __global__ __aicore__ void tilexr_ep_combine_cross_node_kernel(GM_ADD ClearLocalWindow(sendWindow, rankSize, maxRoutesPerSrc, rowBytes, slotBytes, totalBytes, tBuf); ClearLocalWindow(recvWindow, rankSize, maxRoutesPerSrc, rowBytes, slotBytes, totalBytes, tBuf); - const int32_t localNodeStart = TileXREpNodeStart(rank, localRankSize); - const int32_t localNodeEnd = TileXREpNodeEnd(rank, localRankSize, rankSize); + const int32_t localNodeStart = TileXREpNodeStart(rank, effectiveLocalRankSize); + const int32_t localNodeEnd = TileXREpNodeEnd(rank, effectiveLocalRankSize, rankSize); sync.SetInnerFlag(static_cast(magic), TileXREp::kEpStepCombineWindowCleared); for (int32_t peer = localNodeStart; peer < localNodeEnd; ++peer) { sync.WaitRankInnerFlag(static_cast(magic), TileXREp::kEpStepCombineWindowCleared, peer); @@ -227,7 +224,10 @@ extern "C" __global__ __aicore__ void tilexr_ep_combine_cross_node_kernel(GM_ADD ScatterCombineRows(sendWindow, expertOutGM, assistInfoForCombineGM, epRecvCountsGM, rankSize, maxRoutesPerSrc, rowBytes, slotBytes, payloadBytesPerSlot, tBuf); for (int32_t dstRank = 0; dstRank < rankSize; ++dstRank) { - if (TileXREpSameNode(dstRank, rank, localRankSize)) { + if (useUdmaForAllPeers && dstRank == rank) { + CopyBytesGmToGm(recvWindow + SlotOffset(rank, slotBytes), + sendWindow + SlotOffset(rank, slotBytes), tBuf, slotBytes); + } else if (TileXREpSameNode(dstRank, rank, effectiveLocalRankSize)) { CopyBytesGmToGm(shareAddrs[dstRank] + TileXR::IPC_DATA_OFFSET + SlotOffset(rank, slotBytes), sendWindow + SlotOffset(dstRank, slotBytes), tBuf, slotBytes); } @@ -241,14 +241,15 @@ extern "C" __global__ __aicore__ void tilexr_ep_combine_cross_node_kernel(GM_ADD sync.WaitRankInnerFlag(static_cast(magic), TileXREp::kEpStepCombineReady, peer); } - TileXREpNotifyRemoteUdmaReadySeparate(args, sendWindow, rank, rankSize, localRankSize, totalBytes, magic, - slotBytes, TileXREp::kEpStepCombineReady, combineWindowOffset + UDMARecvWindowOffset(totalBytes)); + TileXREpNotifyRemoteUdmaReadySeparate(args, sendWindow, rank, rankSize, effectiveLocalRankSize, totalBytes, + magic, slotBytes, TileXREp::kEpStepCombineReady, + combineWindowOffset + UDMARecvWindowOffset(totalBytes)); sync.SetInnerFlag(static_cast(magic), TileXREp::kEpStepCombineGatewayReady); for (int32_t peer = localNodeStart; peer < localNodeEnd; ++peer) { sync.WaitRankInnerFlag(static_cast(magic), TileXREp::kEpStepCombineGatewayReady, peer); } - const bool remoteReady = TileXREpWaitRemoteUdmaReady(sendWindow, rank, rankSize, localRankSize, totalBytes, - magic, TileXREp::kEpStepCombineReady, tBuf); + const bool remoteReady = TileXREpWaitRemoteUdmaReady(sendWindow, rank, rankSize, effectiveLocalRankSize, + totalBytes, magic, TileXREp::kEpStepCombineReady, tBuf); if (!remoteReady) { TileXREpStoreStatusValue(workspaceGM, totalBytes, rankSize, slotBytes, TileXREp::kEpStatusRemoteReadyTimeout); @@ -277,6 +278,7 @@ extern "C" __global__ __aicore__ void tilexr_ep_combine_cross_node_drain_kernel( const int32_t rank = args->rank; const int32_t rankSize = args->rankSize; const int32_t localRankSize = args->localRankSize; + const bool useUdmaForAllPeers = localRankSize == rankSize; if (rankSize <= 0 || rankSize > TileXR::TILEXR_MAX_RANK_SIZE || rank < 0 || rank >= rankSize || localRankSize <= 0 || localRankSize > rankSize || !IsValidShape(bs, h, topK, moeExpertNum, dtypeBytes, maxRoutesPerSrc, rowBytes, payloadBytesPerSlot, @@ -300,7 +302,7 @@ extern "C" __global__ __aicore__ void tilexr_ep_combine_cross_node_drain_kernel( ZeroRows(yOutGM, bs, h); for (int32_t expertRank = 0; expertRank < rankSize; ++expertRank) { GM_ADDR sourceWindow = recvWindow; - if (TileXREpSameNode(expertRank, rank, localRankSize)) { + if (!useUdmaForAllPeers && TileXREpSameNode(expertRank, rank, localRankSize)) { sourceWindow = localIpcWindow; } else { GM_ADDR slotBase = recvWindow + SlotOffset(expertRank, slotBytes); diff --git a/src/ep/kernels/tilexr_ep_dispatch_kernel.cpp b/src/ep/kernels/tilexr_ep_dispatch_kernel.cpp index 079c2c73..d854206c 100644 --- a/src/ep/kernels/tilexr_ep_dispatch_kernel.cpp +++ b/src/ep/kernels/tilexr_ep_dispatch_kernel.cpp @@ -17,6 +17,9 @@ constexpr int64_t kEpQuantModeNone = 0; constexpr int64_t kEpQuantModeStatic = 1; constexpr int64_t kEpQuantModePerTokenDynamic = 2; +static_assert(TileXREp::kEpUdmaReadyStrideBytes == TileXR::TILEXR_UDMA_CACHE_LINE_SIZE, + "EP UDMA ready slots must match the UDMA cache-line size"); + __aicore__ inline int64_t AlignUp(int64_t value, int64_t alignment) { if (alignment <= 0) { @@ -44,7 +47,7 @@ __aicore__ inline int64_t AssistOffset(int64_t srcRank, int64_t slotBytes, int64 __aicore__ inline int64_t UDMAReadyOffset(int64_t totalBytes, int32_t rank) { return AlignUp(totalBytes, TileXREp::kEpWindowAlignmentBytes) * 2 + - static_cast(rank) * static_cast(sizeof(uint64_t)); + static_cast(rank) * TileXREp::kEpUdmaReadyStrideBytes; } __aicore__ inline int64_t UDMARecvWindowOffset(int64_t totalBytes) @@ -52,6 +55,21 @@ __aicore__ inline int64_t UDMARecvWindowOffset(int64_t totalBytes) return AlignUp(totalBytes, TileXREp::kEpWindowAlignmentBytes); } +__aicore__ inline int64_t UDMAOperationBytes(int64_t totalBytes, int32_t rankSize, int64_t slotBytes) +{ + const int64_t readyEnd = UDMAReadyOffset(totalBytes, rankSize); + const int64_t relaySlotsOffset = AlignUp(readyEnd, TileXR::TILEXR_UDMA_CACHE_LINE_SIZE); + const int64_t relayBytes = static_cast(rankSize) * static_cast(rankSize) * slotBytes; + const int64_t relayReadyBase = AlignUp(relaySlotsOffset + relayBytes, TileXR::TILEXR_UDMA_CACHE_LINE_SIZE); + const int64_t relayReadyBytes = static_cast(rankSize) * TileXREp::kEpUdmaReadyStrideBytes; + return AlignUp(relayReadyBase + relayReadyBytes, TileXR::TILEXR_UDMA_CACHE_LINE_SIZE); +} + +__aicore__ inline int64_t UDMAStatusOffset(int64_t totalBytes, int32_t rankSize, int64_t slotBytes) +{ + return UDMAOperationBytes(totalBytes, rankSize, slotBytes) * 2; +} + __aicore__ inline int64_t TileXREpTpWindowOffset(int64_t totalBytes) { return AlignUp(totalBytes, TileXREp::kEpWindowAlignmentBytes) * 3; @@ -89,6 +107,40 @@ __aicore__ inline void CopyBytesGmToGm( AscendC::PipeBarrier(); } +__aicore__ inline void TileXREpInvalidateLocalCacheLines(GM_ADDR localGM, int64_t bytes) +{ + if (localGM == nullptr || bytes <= 0) { + return; + } + __gm__ uint8_t *start = reinterpret_cast<__gm__ uint8_t *>( + reinterpret_cast(localGM) / TileXR::TILEXR_UDMA_CACHE_LINE_SIZE * + TileXR::TILEXR_UDMA_CACHE_LINE_SIZE); + __gm__ uint8_t *end = reinterpret_cast<__gm__ uint8_t *>( + (reinterpret_cast(localGM) + static_cast(bytes) - 1) / + TileXR::TILEXR_UDMA_CACHE_LINE_SIZE * TileXR::TILEXR_UDMA_CACHE_LINE_SIZE); + AscendC::GlobalTensor global; + global.SetGlobalBuffer(start); + for (uint64_t offset = 0; offset <= static_cast(end - start); + offset += TileXR::TILEXR_UDMA_CACHE_LINE_SIZE) { + __asm__ __volatile__(""); + AscendC::DataCacheCleanAndInvalid(global[offset]); + __asm__ __volatile__(""); + } + AscendC::PipeBarrier(); +} + +__aicore__ inline void TileXREpStoreStatusValue(GM_ADDR workspaceGM, int64_t totalBytes, int32_t rankSize, + int64_t slotBytes, uint64_t status) +{ + GM_ADDR statusAddr = workspaceGM + UDMAStatusOffset(totalBytes, rankSize, slotBytes); + auto statusGM = reinterpret_cast<__gm__ uint64_t *>(statusAddr); + statusGM[0] = status; + AscendC::PipeBarrier(); + TileXR::UDMACleanCacheLines(reinterpret_cast<__gm__ uint8_t *>(statusAddr), sizeof(uint64_t)); + AscendC::PipeBarrier(); +} + __aicore__ inline int32_t LoadInt32FromGm(GM_ADDR srcGM, AscendC::TBuf &tBuf) { AscendC::LocalTensor local = @@ -276,10 +328,11 @@ __aicore__ inline void ClearLocalWindow( AscendC::PipeBarrier(); } -__aicore__ inline bool TileXREpUsesUdmaWindow(const __gm__ TileXR::CommArgs *args, GM_ADDR workspaceGM) +__aicore__ inline bool TileXREpUsesUdmaWindow( + const __gm__ TileXR::CommArgs *args, GM_ADDR workspaceGM, int64_t bytes) { - return workspaceGM != nullptr && args->localRankSize > 0 && args->localRankSize < args->rankSize && - TileXR::UDMARegistryEnabled(args); + return workspaceGM != nullptr && args != nullptr && args->localRankSize > 0 && + args->localRankSize < args->rankSize && bytes > 0 && TileXR::UDMARegistryEnabled(args); } __aicore__ inline GM_ADDR TileXREpWindowBase(GM_ADDR *shareAddrs, int32_t rank, bool useUdmaWindow, GM_ADDR workspaceGM) @@ -308,6 +361,31 @@ __aicore__ inline bool TileXREpUsesUdmaPeer(int32_t rank, int32_t peer, int32_t return TileXREpIsRemotePeer(rank, peer, localRankSize); } +__aicore__ inline bool TileXREpLoadLocalPeerMems(__gm__ TileXR::CommArgs *args, GM_ADDR *shareAddrs, + int32_t rank, int32_t rankSize, int32_t localRankSize) +{ + if (args == nullptr || shareAddrs == nullptr || localRankSize <= 0) { + return false; + } + for (int32_t peer = 0; peer < rankSize; ++peer) { + shareAddrs[peer] = nullptr; + } + const int32_t localNodeStart = (rank / localRankSize) * localRankSize; + int32_t localNodeEnd = localNodeStart + localRankSize; + if (localNodeEnd > rankSize) { + localNodeEnd = rankSize; + } + AscendC::GlobalTensor peerMems; + peerMems.SetGlobalBuffer(&(args->peerMems[0]), TileXR::TILEXR_MAX_RANK_SIZE); + for (int32_t peer = localNodeStart; peer < localNodeEnd; ++peer) { + shareAddrs[peer] = peerMems.GetValue(peer); + if (shareAddrs[peer] == nullptr) { + return false; + } + } + return true; +} + __aicore__ inline void TileXREpPublishLocalUdmaSlot(GM_ADDR localWindow, GM_ADDR sendWindow, int32_t dstRank, int64_t slotBytes, AscendC::TBuf &tBuf) { @@ -369,6 +447,20 @@ __aicore__ inline uint64_t TileXREpReadyValue(int64_t magic) return (static_cast(magic) << 32) | static_cast(TileXREp::kEpStepDispatchReady); } +__aicore__ inline void TileXREpStoreDispatchDebugMark(GM_ADDR localWindow, int64_t totalBytes, int32_t rank, + uint64_t mark) +{ + if (localWindow == nullptr || rank < 0) { + return; + } + GM_ADDR markAddr = localWindow + UDMAReadyOffset(totalBytes, rank); + auto markGM = reinterpret_cast<__gm__ uint64_t *>(markAddr); + markGM[0] = mark; + AscendC::PipeBarrier(); + TileXR::UDMACleanCacheLines(reinterpret_cast<__gm__ uint8_t *>(markAddr), sizeof(uint64_t)); + AscendC::PipeBarrier(); +} + __aicore__ inline int32_t TileXREpTpGroupStartRank(int32_t rank, int32_t tpWorldSize) { if (tpWorldSize <= 0) { @@ -431,7 +523,7 @@ __aicore__ inline void TileXREpNotifyAllUdmaReady( } } -__aicore__ inline void TileXREpWaitUdmaReady(GM_ADDR localWindow, int32_t rank, int32_t rankSize, +__aicore__ inline bool TileXREpWaitUdmaReady(GM_ADDR localWindow, int32_t rank, int32_t rankSize, int32_t localRankSize, int64_t totalBytes, int64_t magic, AscendC::TBuf &tBuf) { const uint64_t ready = TileXREpReadyValue(magic); @@ -439,12 +531,24 @@ __aicore__ inline void TileXREpWaitUdmaReady(GM_ADDR localWindow, int32_t rank, if (!TileXREpUsesUdmaPeer(rank, peer, localRankSize)) { continue; } - while (LoadUint64FromGm(localWindow + UDMAReadyOffset(totalBytes, peer), tBuf) != ready) { + GM_ADDR readyAddr = localWindow + UDMAReadyOffset(totalBytes, peer); + int64_t retries = 0; + while (true) { + TileXREpInvalidateLocalCacheLines(readyAddr, static_cast(sizeof(uint64_t))); + if (LoadUint64FromGm(readyAddr, tBuf) == ready) { + break; + } + ++retries; + if (retries >= static_cast(TileXR::TILEXR_UDMA_MAX_RETRY_TIMES)) { + AscendC::printf("tilexr_ep_udma_ready timeout rank %d peer %d\n", rank, peer); + return false; + } } } + return true; } -__aicore__ inline void TileXREpWaitAllUdmaReady(GM_ADDR localWindow, int32_t rank, int32_t rankSize, +__aicore__ inline bool TileXREpWaitAllUdmaReady(GM_ADDR localWindow, int32_t rank, int32_t rankSize, int64_t totalBytes, int64_t magic, AscendC::TBuf &tBuf) { const uint64_t ready = TileXREpReadyValue(magic); @@ -452,9 +556,21 @@ __aicore__ inline void TileXREpWaitAllUdmaReady(GM_ADDR localWindow, int32_t ran if (peer == rank) { continue; } - while (LoadUint64FromGm(localWindow + UDMAReadyOffset(totalBytes, peer), tBuf) != ready) { + GM_ADDR readyAddr = localWindow + UDMAReadyOffset(totalBytes, peer); + int64_t retries = 0; + while (true) { + TileXREpInvalidateLocalCacheLines(readyAddr, static_cast(sizeof(uint64_t))); + if (LoadUint64FromGm(readyAddr, tBuf) == ready) { + break; + } + ++retries; + if (retries >= static_cast(TileXR::TILEXR_UDMA_MAX_RETRY_TIMES)) { + AscendC::printf("tilexr_ep_udma_all_ready timeout rank %d peer %d\n", rank, peer); + return false; + } } } + return true; } __aicore__ inline void TileXREpFlushDispatchSlotHeaders( @@ -660,9 +776,12 @@ __aicore__ inline int64_t TileXREpAppendTpGroupRows(GM_ADDR *shareAddrs, GM_ADDR int64_t outRecord, int64_t maxRoutesPerSrc, int64_t rowBytes, int64_t payloadRowBytes, int64_t slotBytes, int64_t payloadBytesPerSlot, int64_t totalBytes, int64_t quantMode, int64_t localExpertNum, int64_t sharedExpertNum, int64_t sharedExpertRankNum, GM_ADDR dynamicScalesOutGM, GM_ADDR expertTokenNumsOutGM, - int64_t expertTokenNumsType, int64_t magic, __gm__ int32_t *tpRecvCountsOut, + int64_t expertTokenNumsType, int64_t magic, __gm__ int32_t *tpRecvCountsOut, bool *timedOut, AscendC::TBuf &tBuf) { + if (timedOut != nullptr) { + *timedOut = false; + } const int32_t tpGroupStartRank = TileXREpTpGroupStartRank(rank, tpWorldSize); if (tpGroupStartRank < 0 || expertRankSize <= 0) { return outRecord; @@ -682,6 +801,12 @@ __aicore__ inline int64_t TileXREpAppendTpGroupRows(GM_ADDR *shareAddrs, GM_ADDR outRecord, maxRoutesPerSrc, rowBytes, payloadRowBytes, slotBytes, payloadBytesPerSlot, quantMode, localExpertNum, sharedExpertNum, sharedExpertRankNum, dynamicScalesOutGM, expertTokenNumsOutGM, expertTokenNumsType, magic, &sourceCount, tBuf); + if (sourceCount < 0) { + if (timedOut != nullptr) { + *timedOut = true; + } + return outRecord; + } peerTpCount += sourceCount > 0 ? sourceCount : 0; } if (tpRecvCountsOut != nullptr) { @@ -747,7 +872,7 @@ extern "C" __global__ __aicore__ void tilexr_ep_dispatch_kernel(GM_ADDR commArgs SyncCollectives sync; sync.Init(rank, rankSize, shareAddrs, tBuf); - const bool useUdmaWindow = TileXREpUsesUdmaWindow(args, workspaceGM); + const bool useUdmaWindow = TileXREpUsesUdmaWindow(args, workspaceGM, slotBytes); GM_ADDR localWindow = TileXREpWindowBase(shareAddrs, rank, useUdmaWindow, workspaceGM); ClearLocalWindow(localWindow, rankSize, maxRoutesPerSrc, rowBytes, slotBytes, totalBytes, tBuf); sync.SetInnerFlag(static_cast(magic), TileXREp::kEpStepWindowCleared); @@ -806,7 +931,8 @@ extern "C" __global__ __aicore__ void tilexr_ep_dispatch_kernel(GM_ADDR commArgs outRecord = TileXREpDrainSourceWindow(sourceWindow, slotRank, srcRank, slotBytes, payloadBytesPerSlot, maxRoutesPerSrc, rowBytes, payloadRowBytes, quantMode, localExpertNum64, sharedExpertNum, sharedExpertRankNum, expandXOutGM, dynamicScalesOutGM, expertTokenNumsOutGM, expertTokenNumsType, - magic, epRecvCountsOut, tpRecvCountsOut, localAssistBase, outRecord, tpPublishWindow, rank, tBuf); + magic, epRecvCountsOut, tpRecvCountsOut, localAssistBase, outRecord, tpPublishWindow, rank, nullptr, + tBuf); } if (tpPublishWindow != nullptr) { const int32_t localTpCount = static_cast(outRecord); @@ -816,7 +942,7 @@ extern "C" __global__ __aicore__ void tilexr_ep_dispatch_kernel(GM_ADDR commArgs expertRankSize, effectiveTpWorldSize, outRecord, maxRoutesPerSrc, rowBytes, payloadRowBytes, slotBytes, payloadBytesPerSlot, totalBytes, quantMode, localExpertNum64, sharedExpertNum, sharedExpertRankNum, dynamicScalesOutGM, expertTokenNumsOutGM, expertTokenNumsType, magic, - tpRecvCountsOut, tBuf); + tpRecvCountsOut, nullptr, tBuf); } TileXREpFinalizeExpertTokenNums(expertTokenNumsOutGM, localExpertNum64, expertTokenNumsType); if (!useUdmaWindow) { @@ -854,11 +980,14 @@ extern "C" __global__ __aicore__ void tilexr_ep_dispatch_cross_node_kernel(GM_AD const int32_t rank = args->rank; const int32_t rankSize = args->rankSize; const int32_t localRankSize = args->localRankSize; + const bool useUdmaForAllPeers = localRankSize == rankSize; + const int32_t effectiveLocalRankSize = useUdmaForAllPeers ? 1 : localRankSize; + TileXREpStoreDispatchDebugMark(workspaceGM, totalBytes, rank, 0xd1500001ULL); const int32_t effectiveTpWorldSize = tpWorldSize == 0 ? 1 : static_cast(tpWorldSize); const int32_t effectiveTpRankId = static_cast(tpRankId); const int32_t expertRankSize = effectiveTpWorldSize > 0 ? rankSize / effectiveTpWorldSize : 0; if (rankSize <= 0 || rankSize > TileXR::TILEXR_MAX_RANK_SIZE || rank < 0 || rank >= rankSize || - localRankSize <= 0 || localRankSize >= rankSize || + localRankSize <= 0 || localRankSize > rankSize || effectiveTpWorldSize <= 0 || effectiveTpRankId < 0 || effectiveTpRankId >= effectiveTpWorldSize || expertRankSize <= 0 || expertRankSize * effectiveTpWorldSize != rankSize || (quantMode != kEpQuantModeNone && quantMode != kEpQuantModeStatic && @@ -873,28 +1002,28 @@ extern "C" __global__ __aicore__ void tilexr_ep_dispatch_cross_node_kernel(GM_AD } GM_ADDR shareAddrs[TileXR::TILEXR_MAX_RANK_SIZE]; - if (localRankSize > 1) { - AscendC::GlobalTensor peerMems; - peerMems.SetGlobalBuffer(&(args->peerMems[0]), TileXR::TILEXR_MAX_RANK_SIZE); - for (int32_t peer = 0; peer < rankSize; ++peer) { - shareAddrs[peer] = peerMems.GetValue(peer); - if (shareAddrs[peer] == nullptr) { - return; - } - } + if (localRankSize > 1 && + !TileXREpLoadLocalPeerMems(args, shareAddrs, rank, rankSize, localRankSize)) { + return; } AscendC::TPipe pipe; AscendC::TBuf tBuf; pipe.InitBuffer(tBuf, kEpUbBytes); + TileXREpStoreStatusValue(workspaceGM, totalBytes, rankSize, slotBytes, TileXREp::kEpStatusOk); + TileXREpStoreDispatchDebugMark(workspaceGM, totalBytes, rank, 0xd1500002ULL); GM_ADDR sendWindow = workspaceGM; GM_ADDR recvWindow = workspaceGM + UDMARecvWindowOffset(totalBytes); GM_ADDR localIpcWindow = localRankSize > 1 ? shareAddrs[rank] + TileXR::IPC_DATA_OFFSET : nullptr; + GM_ADDR dispatchIpcWindow = useUdmaForAllPeers ? nullptr : localIpcWindow; ClearLocalWindow(sendWindow, rankSize, maxRoutesPerSrc, rowBytes, slotBytes, totalBytes, tBuf); + TileXREpStoreDispatchDebugMark(workspaceGM, totalBytes, rank, 0xd1500003ULL); ClearLocalWindow(recvWindow, rankSize, maxRoutesPerSrc, rowBytes, slotBytes, totalBytes, tBuf); + TileXREpStoreDispatchDebugMark(workspaceGM, totalBytes, rank, 0xd1500004ULL); if (localRankSize > 1) { ClearLocalWindow(localIpcWindow, rankSize, maxRoutesPerSrc, rowBytes, slotBytes, totalBytes, tBuf); + TileXREpStoreDispatchDebugMark(workspaceGM, totalBytes, rank, 0xd1500005ULL); } int64_t dstCounts[TileXR::TILEXR_MAX_RANK_SIZE]; @@ -906,14 +1035,16 @@ extern "C" __global__ __aicore__ void tilexr_ep_dispatch_cross_node_kernel(GM_AD if (localExpertNum <= 0) { return; } + TileXREpStoreDispatchDebugMark(workspaceGM, totalBytes, rank, 0xd1500006ULL); const int64_t inputRowBytes = h * static_cast(sizeof(uint16_t)); - TileXREpRouteLocalTokens(sendWindow, localIpcWindow, rank, rankSize, expertRankSize, effectiveTpWorldSize, + TileXREpRouteLocalTokens(sendWindow, dispatchIpcWindow, rank, rankSize, expertRankSize, effectiveTpWorldSize, effectiveTpRankId, localRankSize, xGM, expertIdsGM, scalesGM, xActiveMaskGM, bs, h, topK, moeExpertNum, sharedExpertNum, sharedExpertRankNum, localExpertNum, maxRoutesPerSrc, inputRowBytes, rowBytes, payloadRowBytes, slotBytes, payloadBytesPerSlot, quantMode, dstCounts, tBuf); + TileXREpStoreDispatchDebugMark(workspaceGM, totalBytes, rank, 0xd1500007ULL); for (int32_t dstRank = 0; dstRank < rankSize; ++dstRank) { - GM_ADDR targetWindow = TileXREpDispatchWriteWindow(sendWindow, localIpcWindow, rank, dstRank, + GM_ADDR targetWindow = TileXREpDispatchWriteWindow(sendWindow, dispatchIpcWindow, rank, dstRank, localRankSize); const int64_t payloadBytes = dstCounts[dstRank] * payloadRowBytes + (quantMode == kEpQuantModePerTokenDynamic ? @@ -922,19 +1053,36 @@ extern "C" __global__ __aicore__ void tilexr_ep_dispatch_cross_node_kernel(GM_AD rank, payloadBytes, dstCounts[dstRank] * static_cast(sizeof(TileXREp::EpAssistTuple)), magic, tBuf); } + TileXREpStoreDispatchDebugMark(workspaceGM, totalBytes, rank, 0xd1500008ULL); CopyBytesGmToGm(recvWindow + SlotOffset(rank, slotBytes), sendWindow + SlotOffset(rank, slotBytes), tBuf, slotBytes); + TileXREpStoreDispatchDebugMark(workspaceGM, totalBytes, rank, 0xd1500009ULL); TileXREpFlushDispatchSlotHeaders(sendWindow, rankSize, slotBytes, tBuf); if (localRankSize > 1) { TileXREpFlushDispatchSlotHeaders(localIpcWindow, rankSize, slotBytes, tBuf); } + TileXREpStoreDispatchDebugMark(workspaceGM, totalBytes, rank, 0xd150000aULL); - if (localRankSize <= 1) { + if (effectiveLocalRankSize <= 1) { TileXREpNotifyAllUdmaReady(args, sendWindow, rank, rankSize, totalBytes, magic, slotBytes); - TileXREpWaitAllUdmaReady(workspaceGM, rank, rankSize, totalBytes, magic, tBuf); + TileXREpStoreDispatchDebugMark(workspaceGM, totalBytes, rank, 0xd150000bULL); + if (!TileXREpWaitAllUdmaReady(workspaceGM, rank, rankSize, totalBytes, magic, tBuf)) { + TileXREpStoreStatusValue(workspaceGM, totalBytes, rankSize, slotBytes, + TileXREp::kEpStatusDispatchReadyTimeout); + return; + } + TileXREpStoreDispatchDebugMark(workspaceGM, totalBytes, rank, 0xd150000cULL); } else { - TileXREpNotifyUdmaReady(args, sendWindow, rank, rankSize, localRankSize, totalBytes, magic, slotBytes, - tBuf); + TileXREpNotifyUdmaReady(args, sendWindow, rank, rankSize, effectiveLocalRankSize, totalBytes, magic, + slotBytes, tBuf); + TileXREpStoreDispatchDebugMark(workspaceGM, totalBytes, rank, 0xd150000dULL); + if (!TileXREpWaitUdmaReady( + workspaceGM, rank, rankSize, effectiveLocalRankSize, totalBytes, magic, tBuf)) { + TileXREpStoreStatusValue(workspaceGM, totalBytes, rankSize, slotBytes, + TileXREp::kEpStatusDispatchReadyTimeout); + return; + } + TileXREpStoreDispatchDebugMark(workspaceGM, totalBytes, rank, 0xd150000eULL); } auto epRecvCountsOut = reinterpret_cast<__gm__ int32_t *>(epRecvCountsOutGM); @@ -949,24 +1097,37 @@ extern "C" __global__ __aicore__ void tilexr_ep_dispatch_cross_node_kernel(GM_AD } for (int32_t expertSrcRank = 0; expertSrcRank < expertRankSize; ++expertSrcRank) { const int32_t srcRank = expertSrcRank * effectiveTpWorldSize + effectiveTpRankId; - const bool sameNodeSource = localRankSize > 1 && srcRank != rank && - TileXREpIsSameNodePeer(rank, srcRank, localRankSize); + const bool sameNodeSource = effectiveLocalRankSize > 1 && srcRank != rank && + TileXREpIsSameNodePeer(rank, srcRank, effectiveLocalRankSize); GM_ADDR sourceWindow = sameNodeSource ? shareAddrs[srcRank] + TileXR::IPC_DATA_OFFSET : recvWindow; const int32_t slotRank = sameNodeSource ? rank : srcRank; + int32_t sourceCount = 0; outRecord = TileXREpDrainSourceWindow(sourceWindow, slotRank, srcRank, slotBytes, payloadBytesPerSlot, maxRoutesPerSrc, rowBytes, payloadRowBytes, quantMode, localExpertNum, sharedExpertNum, sharedExpertRankNum, expandXOutGM, dynamicScalesOutGM, expertTokenNumsOutGM, expertTokenNumsType, - magic, epRecvCountsOut, tpRecvCountsOut, localAssistBase, outRecord, tpPublishWindow, rank, tBuf); + magic, epRecvCountsOut, tpRecvCountsOut, localAssistBase, outRecord, tpPublishWindow, rank, + &sourceCount, tBuf); + if (sourceCount < 0) { + TileXREpStoreStatusValue(workspaceGM, totalBytes, rankSize, slotBytes, + TileXREp::kEpStatusDispatchSlotTimeout); + return; + } } if (tpPublishWindow != nullptr) { const int32_t localTpCount = static_cast(outRecord); tpRecvCountsOut[tpRankId] = localTpCount; AscendC::PipeBarrier(); + bool tpSlotTimedOut = false; outRecord = TileXREpAppendTpGroupRows(shareAddrs, expandXOutGM, localAssistBase, rank, expertRankSize, effectiveTpWorldSize, outRecord, maxRoutesPerSrc, rowBytes, payloadRowBytes, slotBytes, payloadBytesPerSlot, totalBytes, quantMode, localExpertNum, sharedExpertNum, sharedExpertRankNum, dynamicScalesOutGM, expertTokenNumsOutGM, expertTokenNumsType, magic, - tpRecvCountsOut, tBuf); + tpRecvCountsOut, &tpSlotTimedOut, tBuf); + if (tpSlotTimedOut) { + TileXREpStoreStatusValue(workspaceGM, totalBytes, rankSize, slotBytes, + TileXREp::kEpStatusDispatchSlotTimeout); + return; + } } TileXREpFinalizeExpertTokenNums(expertTokenNumsOutGM, localExpertNum, expertTokenNumsType); } diff --git a/src/ep/kernels/tilexr_ep_kernel_common.h b/src/ep/kernels/tilexr_ep_kernel_common.h index bc057cf9..336956d6 100644 --- a/src/ep/kernels/tilexr_ep_kernel_common.h +++ b/src/ep/kernels/tilexr_ep_kernel_common.h @@ -15,6 +15,9 @@ constexpr uint32_t kEpScalarUbBytes = 64; constexpr uint32_t kEpScalarUbOffset = kEpSyncUbBytes - kEpScalarUbBytes; constexpr int64_t kTileXrDataTypeFp16 = 3; +static_assert(TileXREp::kEpUdmaReadyStrideBytes == TileXR::TILEXR_UDMA_CACHE_LINE_SIZE, + "EP UDMA ready slots must match the UDMA cache-line size"); + __aicore__ inline int64_t AlignUp(int64_t value, int64_t alignment) { if (alignment <= 0) { @@ -42,7 +45,7 @@ __aicore__ inline int64_t AssistOffset(int64_t srcRank, int64_t slotBytes, int64 __aicore__ inline int64_t UDMAReadyOffset(int64_t totalBytes, int32_t rank) { return AlignUp(totalBytes, TileXREp::kEpWindowAlignmentBytes) * 2 + - static_cast(rank) * static_cast(sizeof(uint64_t)); + static_cast(rank) * TileXREp::kEpUdmaReadyStrideBytes; } __aicore__ inline int64_t UDMARecvWindowOffset(int64_t totalBytes) @@ -73,12 +76,12 @@ __aicore__ inline int64_t UDMARelayReadyOffset( int64_t totalBytes, int32_t rankSize, int64_t slotBytes, int32_t srcRank) { return UDMARelayReadyBaseOffset(totalBytes, rankSize, slotBytes) + - static_cast(srcRank) * static_cast(sizeof(uint64_t)); + static_cast(srcRank) * TileXREp::kEpUdmaReadyStrideBytes; } __aicore__ inline int64_t UDMAOperationBytes(int64_t totalBytes, int32_t rankSize, int64_t slotBytes) { - const int64_t relayReadyBytes = static_cast(rankSize) * static_cast(sizeof(uint64_t)); + const int64_t relayReadyBytes = static_cast(rankSize) * TileXREp::kEpUdmaReadyStrideBytes; return AlignUp(UDMARelayReadyBaseOffset(totalBytes, rankSize, slotBytes) + relayReadyBytes, TileXR::TILEXR_UDMA_CACHE_LINE_SIZE); } @@ -347,6 +350,28 @@ __aicore__ inline int32_t TileXREpNodeEnd(int32_t rank, int32_t localRankSize, i return end > rankSize ? rankSize : end; } +__aicore__ inline bool TileXREpLoadLocalPeerMems(__gm__ TileXR::CommArgs *args, GM_ADDR *shareAddrs, + int32_t rank, int32_t rankSize, int32_t localRankSize) +{ + if (args == nullptr || shareAddrs == nullptr || localRankSize <= 0) { + return false; + } + for (int32_t peer = 0; peer < rankSize; ++peer) { + shareAddrs[peer] = nullptr; + } + AscendC::GlobalTensor peerMems; + peerMems.SetGlobalBuffer(&(args->peerMems[0]), TileXR::TILEXR_MAX_RANK_SIZE); + const int32_t localNodeStart = TileXREpNodeStart(rank, localRankSize); + const int32_t localNodeEnd = TileXREpNodeEnd(rank, localRankSize, rankSize); + for (int32_t peer = localNodeStart; peer < localNodeEnd; ++peer) { + shareAddrs[peer] = peerMems.GetValue(peer); + if (shareAddrs[peer] == nullptr) { + return false; + } + } + return true; +} + __aicore__ inline void TileXREpStoreLocalReadyValue(GM_ADDR localWindow, int64_t totalBytes, int32_t rank, uint64_t ready) { diff --git a/src/include/tilexr_transport.h b/src/include/tilexr_transport.h new file mode 100644 index 00000000..890dd582 --- /dev/null +++ b/src/include/tilexr_transport.h @@ -0,0 +1,87 @@ +/* + * Copyright (c) 2024-2026 TileXR Project + * Licensed under the Apache License, Version 2.0 + */ + +#ifndef TILEXR_TRANSPORT_H +#define TILEXR_TRANSPORT_H + +#include + +#include "comm_args.h" + +namespace TileXR { + +#if TILEXR_ASCENDC_AICORE_COMPILE +#define TILEXR_TRANSPORT_INLINE __aicore__ inline +#define TILEXR_TRANSPORT_GM __gm__ +#else +#define TILEXR_TRANSPORT_INLINE inline +#define TILEXR_TRANSPORT_GM +#endif + +enum class TileXRTransportKind : uint8_t { + MEMORY = 0, + DIRECT_URMA = 1, +}; + +constexpr uint64_t TILEXR_AUTO_SAME_NODE_DIRECT_URMA_THRESHOLD_BYTES = 4ULL * 1024ULL * 1024ULL; +constexpr uint64_t TILEXR_AUTO_CROSS_NODE_DIRECT_URMA_THRESHOLD_BYTES = 128ULL * 1024ULL; +constexpr uint64_t TILEXR_AUTO_DIRECT_URMA_THRESHOLD_BYTES = + TILEXR_AUTO_SAME_NODE_DIRECT_URMA_THRESHOLD_BYTES; + +TILEXR_TRANSPORT_INLINE bool TileXRDirectUrmaCapable(const TILEXR_TRANSPORT_GM CommArgs* args) +{ + return args != nullptr && + ((args->extraFlag & ExtraFlag::UDMA) != 0) && + args->udmaInfoPtr != nullptr; +} + +TILEXR_TRANSPORT_INLINE bool TileXRDirectUrmaAvailable(const TILEXR_TRANSPORT_GM CommArgs* args) +{ + return TileXRDirectUrmaCapable(args) && args->udmaRegistryPtr != nullptr; +} + +TILEXR_TRANSPORT_INLINE bool TileXRDirectUrmaPeerRoutable( + const TILEXR_TRANSPORT_GM CommArgs* args, int targetRank) +{ + if (!TileXRDirectUrmaCapable(args) || args->rankSize <= 1 || args->rank < 0 || + args->rank >= args->rankSize || targetRank < 0 || targetRank >= args->rankSize || + targetRank == args->rank || args->localRankSize <= 0) { + return false; + } + if (args->localRankSize >= args->rankSize) { + return true; + } + return targetRank / args->localRankSize != args->rank / args->localRankSize; +} + +TILEXR_TRANSPORT_INLINE bool TileXRCommSpansNodes(const TILEXR_TRANSPORT_GM CommArgs* args) +{ + return args != nullptr && args->localRankSize > 0 && args->localRankSize < args->rankSize; +} + +TILEXR_TRANSPORT_INLINE TileXRTransportKind TileXRSelectAutoTransport( + const TILEXR_TRANSPORT_GM CommArgs* args, uint64_t bytes) +{ + if (bytes == 0) { + return TileXRTransportKind::MEMORY; + } + if (!TileXRDirectUrmaAvailable(args)) { + return TileXRTransportKind::MEMORY; + } + const uint64_t threshold = TileXRCommSpansNodes(args) ? + TILEXR_AUTO_CROSS_NODE_DIRECT_URMA_THRESHOLD_BYTES : + TILEXR_AUTO_SAME_NODE_DIRECT_URMA_THRESHOLD_BYTES; + if (bytes >= threshold) { + return TileXRTransportKind::DIRECT_URMA; + } + return TileXRTransportKind::MEMORY; +} + +#undef TILEXR_TRANSPORT_INLINE +#undef TILEXR_TRANSPORT_GM + +} // namespace TileXR + +#endif // TILEXR_TRANSPORT_H diff --git a/src/include/tilexr_udma.h b/src/include/tilexr_udma.h index 388cea87..74cac2b4 100644 --- a/src/include/tilexr_udma.h +++ b/src/include/tilexr_udma.h @@ -8,6 +8,7 @@ #include "kernel_operator.h" #include "comm_args.h" +#include "tilexr_transport.h" #include "tilexr_udma_reg.h" #include "tilexr_udma_types.h" @@ -45,6 +46,24 @@ __aicore__ inline bool UDMARegistryEnabled(const __gm__ CommArgs* args) return UDMAEnabled(args) && args->udmaRegistryPtr != nullptr; } +__aicore__ inline bool UDMAPeerEnabled(const __gm__ CommArgs* args, int targetRank) +{ + return UDMARegistryEnabled(args) && TileXRDirectUrmaPeerRoutable(args, targetRank); +} + +__aicore__ inline bool UDMAAllPeersEnabled(const __gm__ CommArgs* args) +{ + if (!UDMARegistryEnabled(args) || args->rank < 0 || args->rank >= args->rankSize) { + return false; + } + for (int peer = 0; peer < args->rankSize; ++peer) { + if (peer != args->rank && !UDMAPeerEnabled(args, peer)) { + return false; + } + } + return true; +} + __aicore__ inline __gm__ UDMAInfo* GetUDMAInfo(const __gm__ CommArgs* args) { return reinterpret_cast<__gm__ UDMAInfo*>(args->udmaInfoPtr); @@ -296,7 +315,7 @@ template __aicore__ inline void UDMAPutNbi( const __gm__ CommArgs* args, int targetRank, const __gm__ T* localSrc, uint64_t byteOffset, uint32_t byteCount) { - if (!UDMARegistryEnabled(args)) return; + if (!UDMAPeerEnabled(args, targetRank)) return; auto registry = GetUDMARegistry(args); if (!UDMARegisteredRangeValid(registry, targetRank, byteOffset, byteCount)) return; @@ -317,7 +336,7 @@ template __aicore__ inline void UDMAGetNbi( const __gm__ CommArgs* args, int sourceRank, __gm__ T* localDst, uint64_t byteOffset, uint32_t byteCount) { - if (!UDMARegistryEnabled(args)) return; + if (!UDMAPeerEnabled(args, sourceRank)) return; auto registry = GetUDMARegistry(args); if (!UDMARegisteredRangeValid(registry, sourceRank, byteOffset, byteCount)) return; @@ -338,7 +357,7 @@ __aicore__ inline void UDMAPutSignalNbi( const __gm__ CommArgs* args, int targetRank, const __gm__ T* localSrc, uint64_t byteOffset, uint32_t byteCount, uint64_t signalByteOffset, uint64_t signal) { - if (!UDMARegistryEnabled(args)) return; + if (!UDMAPeerEnabled(args, targetRank)) return; auto registry = GetUDMARegistry(args); if (!UDMARegisteredRangeValid(registry, targetRank, byteOffset, byteCount) || @@ -365,7 +384,7 @@ __aicore__ inline void UDMAPutRegisteredSignalNbi( __aicore__ inline void UDMAQuiet(const __gm__ CommArgs* args, int targetRank) { - if (!UDMAEnabled(args)) return; + if (!TileXRDirectUrmaPeerRoutable(args, targetRank)) return; __gm__ UDMAInfo* udmaInfo = GetUDMAInfo(args); __gm__ UDMAWQCtx* qpCtxEntry = UDMAGetWQCtx(udmaInfo, targetRank, 0); uint32_t wqeCnt = ld_dev(reinterpret_cast<__gm__ uint32_t*>(qpCtxEntry->wqeCntAddr), 0); diff --git a/tests/ep/CMakeLists.txt b/tests/ep/CMakeLists.txt index e505f1d8..0fe80602 100644 --- a/tests/ep/CMakeLists.txt +++ b/tests/ep/CMakeLists.txt @@ -85,10 +85,17 @@ add_executable(test_tilexr_ep_api_sources add_executable(test_tilexr_ep_kernel_sources unit/test_tilexr_ep_kernel_sources.cpp) +add_executable(test_tilexr_ep_transport_route + unit/test_tilexr_ep_transport_route.cpp + ${TILEXR_ROOT}/src/ep/host/ep_transport_route.cpp +) + add_executable(test_tilexr_ep_host_validation unit/test_tilexr_ep_host_validation.cpp ${TILEXR_ROOT}/src/ep/host/ep_layout.cpp ${TILEXR_ROOT}/src/ep/host/ep_dispatch_host.cpp + ${TILEXR_ROOT}/src/ep/host/ep_transport_route.cpp + ${TILEXR_ROOT}/src/ep/host/ep_launch_context.cpp ) target_include_directories(test_tilexr_ep_layout PRIVATE @@ -99,18 +106,24 @@ target_include_directories(test_tilexr_ep_host_validation PRIVATE ${TILEXR_EP_TEST_INCLUDE_DIRS} ) +target_include_directories(test_tilexr_ep_transport_route PRIVATE + ${TILEXR_EP_TEST_INCLUDE_DIRS} +) + target_compile_definitions(test_tilexr_ep_api_sources PRIVATE TILEXR_SOURCE_ROOT="${TILEXR_ROOT}") target_compile_definitions(test_tilexr_ep_kernel_sources PRIVATE TILEXR_SOURCE_ROOT="${TILEXR_ROOT}") add_test(NAME test_tilexr_ep_layout COMMAND test_tilexr_ep_layout) add_test(NAME test_tilexr_ep_api_sources COMMAND test_tilexr_ep_api_sources) add_test(NAME test_tilexr_ep_kernel_sources COMMAND test_tilexr_ep_kernel_sources) +add_test(NAME test_tilexr_ep_transport_route COMMAND test_tilexr_ep_transport_route) add_test(NAME test_tilexr_ep_host_validation COMMAND test_tilexr_ep_host_validation) install(TARGETS test_tilexr_ep_layout test_tilexr_ep_api_sources test_tilexr_ep_kernel_sources + test_tilexr_ep_transport_route test_tilexr_ep_host_validation RUNTIME DESTINATION ${CMAKE_INSTALL_BINDIR} ) @@ -134,6 +147,7 @@ if(BUILD_TILEXR_EP_DEMO) target_include_directories(tilexr_ep_dispatch_demo PRIVATE "${TILEXR_INSTALL_INCLUDE_SEARCH_DIR}" + "${TILEXR_ROOT}/src/ep/host" "${ASCEND_HOME_PATH}/${ARCH}-linux/include" "${ASCEND_HOME_PATH}/${ARCH}-linux/pkg_inc" "${ASCEND_HOME_PATH}/${ARCH}-linux/pkg_inc/runtime" @@ -142,7 +156,6 @@ if(BUILD_TILEXR_EP_DEMO) target_link_directories(tilexr_ep_dispatch_demo PRIVATE "${TILEXR_INSTALL_LIB_SEARCH_DIR}" "${ASCEND_HOME_PATH}/${ARCH}-linux/lib64" - "${ASCEND_HOME_PATH}/${ARCH}-linux/devlib" "${ASCEND_DRIVER_PATH}/lib64/driver" ) diff --git a/tests/ep/demo/tilexr_ep_dispatch_demo.cpp b/tests/ep/demo/tilexr_ep_dispatch_demo.cpp index e3f34b77..5dfedfed 100644 --- a/tests/ep/demo/tilexr_ep_dispatch_demo.cpp +++ b/tests/ep/demo/tilexr_ep_dispatch_demo.cpp @@ -16,27 +16,29 @@ #include #include "acl/acl.h" +#include "ep_transport_route.h" #include "tilexr_api.h" #include "tilexr_ep.h" +#include "tilexr_transport.h" #include "tilexr_types.h" namespace { -constexpr int64_t kBs = 4; +constexpr int64_t kDefaultBs = 4; constexpr int64_t kH = 8; constexpr int64_t kTopK = 2; -constexpr int64_t kRoutes = kBs * kTopK; -constexpr int64_t kXElements = kBs * kH; constexpr int64_t kAssistInts = 4; constexpr uint16_t kFp16One = 0x3c00; constexpr uint16_t kFp16Two = 0x4000; constexpr std::size_t kUdmaCacheLineBytes = 64; +constexpr std::size_t kUdmaReadyStrideBytes = kUdmaCacheLineBytes; constexpr std::size_t kUdmaRegistrationAlignment = 2 * 1024 * 1024; TileXRUDMAMemHandle g_workspaceHandle = 0; bool g_workspaceRegistered = false; struct DemoConfig { + int64_t bs = kDefaultBs; int64_t moeExpertNum = 8; int64_t sharedExpertNum = 0; int64_t sharedExpertRankNum = 0; @@ -45,7 +47,7 @@ struct DemoConfig { int64_t maxRoutesPerRank() const { - return kBs * (kTopK + sharedExpertNum); + return bs * (kTopK + sharedExpertNum); } int64_t effectiveTpWorldSize() const @@ -332,12 +334,13 @@ std::size_t EpOperationBytes(int rankSize, const DemoConfig &config, std::size_t bool usePerTokenDynamicQuant) { const std::size_t windowBytes = EpWindowBytes(rankSize, config, payloadRowBytes, usePerTokenDynamicQuant); - const std::size_t readyOffset = windowBytes * 2 + static_cast(rankSize) * sizeof(uint64_t); + const std::size_t readyOffset = windowBytes * 2 + + static_cast(rankSize) * kUdmaReadyStrideBytes; const std::size_t relayOffset = AlignSize(readyOffset, kUdmaCacheLineBytes); const std::size_t relayBytes = static_cast(rankSize) * static_cast(rankSize) * EpSlotBytes(config, payloadRowBytes, usePerTokenDynamicQuant); const std::size_t relayReadyOffset = AlignSize(relayOffset + relayBytes, kUdmaCacheLineBytes); - const std::size_t relayReadyBytes = static_cast(rankSize) * sizeof(uint64_t); + const std::size_t relayReadyBytes = static_cast(rankSize) * kUdmaReadyStrideBytes; return AlignSize(relayReadyOffset + relayReadyBytes, kUdmaCacheLineBytes); } @@ -348,6 +351,51 @@ std::size_t EpRequiredWorkspaceBytes(int rankSize, const DemoConfig &config, std return AlignSize(operationBytes * 2 + sizeof(uint64_t), kUdmaCacheLineBytes); } +void DumpCrossNodeWindowSlots(int rank, int rankSize, void *workspaceDev, std::size_t baseOffset, + std::size_t slotBytes, const char *label) +{ + for (int slotRank = 0; slotRank < rankSize; ++slotRank) { + uint64_t slotHeader[8] = {}; + const GM_ADDR slotAddr = static_cast(workspaceDev) + baseOffset + 64 + + static_cast(slotRank) * slotBytes; + if (CheckAcl(aclrtMemcpy(slotHeader, sizeof(slotHeader), slotAddr, sizeof(slotHeader), + ACL_MEMCPY_DEVICE_TO_HOST), "dump dispatch slot header")) { + const uint32_t count = static_cast(slotHeader[0] & 0xffffffffULL); + const uint32_t slotSrc = static_cast((slotHeader[0] >> 32) & 0xffffffffULL); + std::cerr << "rank " << rank << " dump " << label << " slotRank " << slotRank + << " count " << count << " slotSrc " << slotSrc + << " payloadBytes " << slotHeader[1] << " assistBytes " << slotHeader[2] + << " magic " << slotHeader[3] << std::endl; + } + } +} + +void DumpCrossNodeReadyFlags(int rank, int rankSize, void *workspaceDev, std::size_t windowBytes) +{ + const std::size_t readyBase = windowBytes * 2; + for (int peer = 0; peer < rankSize; ++peer) { + uint64_t ready = 0; + const GM_ADDR readyAddr = static_cast(workspaceDev) + readyBase + + static_cast(peer) * kUdmaReadyStrideBytes; + if (CheckAcl(aclrtMemcpy(&ready, sizeof(ready), readyAddr, sizeof(ready), + ACL_MEMCPY_DEVICE_TO_HOST), "dump dispatch ready")) { + std::cerr << "rank " << rank << " dump ready peer " << peer << " value " << ready << std::endl; + } + } +} + +void DumpCrossNodeDispatchWindow(int rank, int rankSize, void *workspaceDev, std::size_t windowBytes, + std::size_t slotBytes, const char *reason) +{ + if (!EnvEnabled("TILEXR_EP_DEMO_DUMP_WINDOW") || workspaceDev == nullptr) { + return; + } + std::cerr << "rank " << rank << " dump dispatch window reason " << reason << std::endl; + DumpCrossNodeReadyFlags(rank, rankSize, workspaceDev, windowBytes); + DumpCrossNodeWindowSlots(rank, rankSize, workspaceDev, 0, slotBytes, "send window"); + DumpCrossNodeWindowSlots(rank, rankSize, workspaceDev, windowBytes, slotBytes, "recv window"); +} + uint16_t XValue(int rank, int64_t token, int64_t h) { return static_cast(0x3c00 + rank * 0x0400 + token * 0x0100 + h * 0x0010); @@ -412,18 +460,19 @@ int8_t DynamicQuantizedXValue(int rank, int64_t token, int64_t h) std::vector ExpertIds(const DemoConfig &config) { - std::vector expertIds(kRoutes); - for (int64_t route = 0; route < kRoutes; ++route) { + const int64_t routes = config.bs * kTopK; + std::vector expertIds(routes); + for (int64_t route = 0; route < routes; ++route) { expertIds[route] = static_cast(route % config.moeExpertNum); } return expertIds; } -std::vector ActiveMask(bool enabled) +std::vector ActiveMask(const DemoConfig &config, bool enabled) { - std::vector mask(kBs, 1); - if (enabled && kBs > 0) { - mask[kBs - 1] = 0; + std::vector mask(config.bs, 1); + if (enabled && config.bs > 0) { + mask[config.bs - 1] = 0; } return mask; } @@ -493,7 +542,7 @@ std::vector BuildExpectedRoutes( if (effectiveTpWorldSize > 1 && srcRank % effectiveTpWorldSize != targetTpRankId) { continue; } - for (int64_t token = 0; token < kBs; ++token) { + for (int64_t token = 0; token < config.bs; ++token) { if (!activeMask.empty() && activeMask[token] == 0) { continue; } @@ -624,9 +673,9 @@ bool ValidateOutputs(int rank, int rankSize, const DemoConfig &config, const std return true; } -bool ValidateCombineOutputs(int rank, const std::vector &yOut) +bool ValidateCombineOutputs(int rank, const DemoConfig &config, const std::vector &yOut) { - for (int64_t token = 0; token < kBs; ++token) { + for (int64_t token = 0; token < config.bs; ++token) { for (int64_t h = 0; h < kH; ++h) { const uint16_t actualValue = yOut[token * kH + h]; if (actualValue != kFp16Two) { @@ -719,6 +768,7 @@ int main(int argc, char **argv) const bool usePerTokenDynamicQuant = quantMode == 2; const float staticQuantScale = static_cast(GetEnvInt("TILEXR_EP_DEMO_STATIC_QUANT_SCALE", 1)); DemoConfig config {}; + config.bs = GetEnvInt("TILEXR_EP_DEMO_BS", static_cast(config.bs)); config.moeExpertNum = GetEnvInt("TILEXR_EP_DEMO_MOE_EXPERT_NUM", static_cast(config.moeExpertNum)); config.sharedExpertNum = GetEnvInt("TILEXR_EP_DEMO_SHARED_EXPERT_NUM", 0); config.sharedExpertRankNum = GetEnvInt("TILEXR_EP_DEMO_SHARED_EXPERT_RANK_NUM", 0); @@ -729,7 +779,8 @@ int main(int argc, char **argv) const int64_t expertRankSize = static_cast(rankSize) / config.effectiveTpWorldSize(); const int64_t moeRankNum = expertRankSize - config.sharedExpertRankNum; - if (rankSize <= 0 || rank < 0 || rank >= rankSize || config.effectiveTpWorldSize() <= 0 || + if (rankSize <= 0 || rank < 0 || rank >= rankSize || config.bs <= 0 || + config.effectiveTpWorldSize() <= 0 || rankSize % config.effectiveTpWorldSize() != 0 || moeRankNum <= 0 || config.moeExpertNum <= 0 || config.moeExpertNum % moeRankNum != 0 || config.sharedExpertNum < 0 || config.sharedExpertRankNum < 0 || @@ -739,7 +790,7 @@ int main(int argc, char **argv) (quantMode != 0 && quantMode != 1 && quantMode != 2) || ((useStaticQuant || usePerTokenDynamicQuant) && !dispatchOnly)) { std::cerr << "This demo expects a valid rank and moeExpertNum divisible by MoE rank num, got moeExpertNum=" - << config.moeExpertNum << " rankSize=" << rankSize + << config.moeExpertNum << " rankSize=" << rankSize << " bs=" << config.bs << " sharedExpertNum=" << config.sharedExpertNum << " sharedExpertRankNum=" << config.sharedExpertRankNum << " tpWorldSize=" << config.tpWorldSize @@ -777,14 +828,14 @@ int main(int argc, char **argv) return 1; } - std::vector hostX(kXElements); - for (int64_t token = 0; token < kBs; ++token) { + std::vector hostX(config.bs * kH); + for (int64_t token = 0; token < config.bs; ++token) { for (int64_t h = 0; h < kH; ++h) { hostX[token * kH + h] = XValue(rank, token, h); } } const std::vector hostExpertIds = ExpertIds(config); - const std::vector hostActiveMask = ActiveMask(useActiveMask); + const std::vector hostActiveMask = ActiveMask(config, useActiveMask); const std::size_t expectedRouteCount = BuildExpectedTpRoutes(rank, rankSize, config, hostActiveMask).size(); @@ -820,8 +871,9 @@ int main(int argc, char **argv) const std::size_t recvCountsBytes = rankSize * sizeof(int32_t); const std::size_t tpRecvCountsBytes = recvCountsBytes; const std::size_t assistBytes = (expandedElements / kH) * kAssistInts * sizeof(int32_t); - const std::size_t yOutBytes = kXElements * sizeof(uint16_t); + const std::size_t yOutBytes = static_cast(config.bs * kH) * sizeof(uint16_t); const std::size_t payloadRowBytes = kH * expandElementBytes; + const std::size_t dispatchSlotBytes = EpSlotBytes(config, payloadRowBytes, usePerTokenDynamicQuant); const std::size_t dispatchWindowBytes = EpWindowBytes(rankSize, config, payloadRowBytes, usePerTokenDynamicQuant); const std::size_t dispatchPayloadBytes = AlignSize(dispatchWindowBytes, 32) * static_cast(config.effectiveTpWorldSize() + 2); @@ -887,17 +939,44 @@ int main(int argc, char **argv) Cleanup(comm, stream, deviceId, deviceSet, aclReady, buffers); return 1; } - const bool crossNode = commArgsHost != nullptr && commArgsHost->localRankSize > 0 && + if (commArgsHost == nullptr) { + std::cerr << "TileXRGetCommArgsHost returned null communicator arguments" << std::endl; + Cleanup(comm, stream, deviceId, deviceSet, aclReady, buffers); + return 1; + } + + const bool crossNode = commArgsHost->localRankSize > 0 && commArgsHost->localRankSize < commArgsHost->rankSize; - if (crossNode && !CheckTileXR(TileXRUDMARegister(comm, static_cast(workspaceDev), workspaceBytes, + TileXREp::EpTransportMode requestedTransport = TileXREp::EpTransportMode::AUTO; + if (!CheckTileXR(TileXREp::TileXREpParseTransportMode( + std::getenv("TILEXR_TRANSPORT_MODE"), &requestedTransport), "TileXREpParseTransportMode")) { + Cleanup(comm, stream, deviceId, deviceSet, aclReady, buffers); + return 1; + } + const bool registerWorkspace = + TileXREp::TileXREpShouldRegisterWorkspace(requestedTransport, *commArgsHost); + if (registerWorkspace && + !CheckTileXR(TileXRUDMARegister(comm, static_cast(workspaceDev), workspaceBytes, &workspaceHandle), "TileXRUDMARegister workspace")) { Cleanup(comm, stream, deviceId, deviceSet, aclReady, buffers); return 1; } - if (crossNode) { + if (registerWorkspace) { g_workspaceHandle = workspaceHandle; g_workspaceRegistered = true; } + TileXR::TileXRTransportKind resolvedTransport = TileXR::TileXRTransportKind::MEMORY; + if (!CheckTileXR(TileXREp::TileXREpResolveTransport(requestedTransport, *commArgsHost, + static_cast(dispatchWindowBytes), &resolvedTransport), "TileXREpResolveTransport")) { + Cleanup(comm, stream, deviceId, deviceSet, aclReady, buffers); + return 1; + } + const bool useRegisteredWorkspace = + resolvedTransport == TileXR::TileXRTransportKind::DIRECT_URMA; + std::cout << "rank " << rank << " bs=" << config.bs + << " dispatchWindowBytes=" << dispatchWindowBytes + << " transport=" << (resolvedTransport == TileXR::TileXRTransportKind::DIRECT_URMA ? + "direct_urma" : "memory") << std::endl; const std::vector hostExpertOut(expandedElements, kFp16One); if (!CheckAcl(aclrtMemcpy(expertOutDev, expandXBytes, hostExpertOut.data(), expandXBytes, @@ -908,24 +987,31 @@ int main(int argc, char **argv) const bool useSharedExperts = config.sharedExpertNum != 0 || config.sharedExpertRankNum != 0; const bool useTp = config.effectiveTpWorldSize() != 1; - const bool useDispatchV2 = crossNode || useActiveMask || useTpRecvCounts || expertTokenNumsType != 1 || - useSharedExperts || useTp || useStaticQuant || usePerTokenDynamicQuant; + const bool useDispatchV2 = crossNode || useRegisteredWorkspace || useActiveMask || useTpRecvCounts || + expertTokenNumsType != 1 || useSharedExperts || useTp || useStaticQuant || usePerTokenDynamicQuant; const int dispatchRet = useDispatchV2 ? TileXRMoeEpDispatchV2(xDev, static_cast(expertIdsDev), scalesDev, - static_cast(xActiveMaskDev), nullptr, comm, kBs, kH, kTopK, config.moeExpertNum, + static_cast(xActiveMaskDev), nullptr, comm, config.bs, kH, kTopK, config.moeExpertNum, expertRankSize, ExpertRankForRank(rank, config), config.tpWorldSize, config.tpRankId, 0, - config.sharedExpertNum, config.sharedExpertRankNum, quantMode, kBs * rankSize, expertTokenNumsType, + config.sharedExpertNum, config.sharedExpertRankNum, quantMode, config.bs * rankSize, expertTokenNumsType, expandXDev, dynamicScalesDev, static_cast(assistDev), static_cast(expertTokenNumsDev), static_cast(recvCountsDev), static_cast(tpRecvCountsDev), nullptr, workspaceDev, (useStaticQuant || usePerTokenDynamicQuant) ? TileXR::TILEXR_DATA_TYPE_INT8 : TileXR::TILEXR_DATA_TYPE_FP16, stream) : - TileXRMoeEpDispatch(xDev, static_cast(expertIdsDev), comm, kBs, kH, kTopK, config.moeExpertNum, + TileXRMoeEpDispatch(xDev, static_cast(expertIdsDev), comm, config.bs, kH, kTopK, + config.moeExpertNum, expandXDev, static_cast(expertTokenNumsDev), static_cast(recvCountsDev), static_cast(assistDev), TileXR::TILEXR_DATA_TYPE_FP16, stream); - if (!CheckTileXR(dispatchRet, "TileXRMoeEpDispatch") || - !CheckAcl(aclrtSynchronizeStream(stream), "aclrtSynchronizeStream")) { + const bool dispatchLaunchOk = CheckTileXR(dispatchRet, "TileXRMoeEpDispatch"); + const bool dispatchSyncOk = dispatchLaunchOk && + CheckAcl(aclrtSynchronizeStream(stream), "aclrtSynchronizeStream"); + if (!dispatchLaunchOk || !dispatchSyncOk) { + if (crossNode) { + DumpCrossNodeDispatchWindow(rank, rankSize, workspaceDev, dispatchWindowBytes, dispatchSlotBytes, + "dispatch failure"); + } Cleanup(comm, stream, deviceId, deviceSet, aclReady, buffers); return 1; } @@ -1017,12 +1103,13 @@ int main(int argc, char **argv) return dispatchOk ? 0 : 1; } - const int combineRet = crossNode ? + const int combineRet = (crossNode || useRegisteredWorkspace) ? TileXRMoeEpCombineV2(expertOutDev, static_cast(assistDev), - static_cast(recvCountsDev), comm, kBs, kH, kTopK, config.moeExpertNum, yOutDev, workspaceDev, + static_cast(recvCountsDev), comm, config.bs, kH, kTopK, config.moeExpertNum, yOutDev, + workspaceDev, TileXR::TILEXR_DATA_TYPE_FP16, stream) : TileXRMoeEpCombine(expertOutDev, static_cast(assistDev), - static_cast(recvCountsDev), comm, kBs, kH, kTopK, config.moeExpertNum, yOutDev, + static_cast(recvCountsDev), comm, config.bs, kH, kTopK, config.moeExpertNum, yOutDev, TileXR::TILEXR_DATA_TYPE_FP16, stream); if (!CheckTileXR(combineRet, "TileXRMoeEpCombine") || !CheckAcl(aclrtSynchronizeStream(stream), "aclrtSynchronizeStream combine")) { @@ -1034,14 +1121,14 @@ int main(int argc, char **argv) return 1; } - std::vector hostYOut(kXElements); + std::vector hostYOut(config.bs * kH); if (!CheckAcl(aclrtMemcpy(hostYOut.data(), yOutBytes, yOutDev, yOutBytes, ACL_MEMCPY_DEVICE_TO_HOST), "copy yOut")) { Cleanup(comm, stream, deviceId, deviceSet, aclReady, buffers); return 1; } - const bool combineOk = ValidateCombineOutputs(rank, hostYOut); + const bool combineOk = ValidateCombineOutputs(rank, config, hostYOut); std::cout << "rank " << rank << " combine validation " << (combineOk ? "PASS" : "FAIL") << std::endl; Cleanup(comm, stream, deviceId, deviceSet, aclReady, buffers); return combineOk ? 0 : 1; diff --git a/tests/ep/unit/test_tilexr_ep_api_sources.cpp b/tests/ep/unit/test_tilexr_ep_api_sources.cpp index 7f942804..bbbb766b 100644 --- a/tests/ep/unit/test_tilexr_ep_api_sources.cpp +++ b/tests/ep/unit/test_tilexr_ep_api_sources.cpp @@ -54,6 +54,19 @@ void CheckNotContains(const std::string &label, const std::string &contents, con } } +void CheckBlockNotContains(const std::string &label, const std::string &contents, const std::string &start, + const std::string &end, const std::string &needle) +{ + const std::size_t startPos = contents.find(start); + const std::size_t endPos = startPos == std::string::npos ? std::string::npos : contents.find(end, startPos); + if (startPos == std::string::npos || endPos == std::string::npos) { + std::cerr << label << " missing block: " << start << std::endl; + ++g_failures; + return; + } + CheckNotContains(label, contents.substr(startPos, endPos - startPos), needle); +} + void TestPublicHeader() { std::string contents; @@ -104,6 +117,22 @@ void TestBuildPlacement() } } +void TestEpHostTargetsExcludeCannDevlib() +{ + std::string epCmake; + if (ReadFile("src/ep/CMakeLists.txt", &epCmake)) { + CheckBlockNotContains("src/ep/CMakeLists.txt tilexr-ep link directories", epCmake, + "target_link_directories(tilexr-ep", "target_link_libraries(tilexr-ep", "-linux/devlib"); + } + + std::string testCmake; + if (ReadFile("tests/ep/CMakeLists.txt", &testCmake)) { + CheckBlockNotContains("tests/ep/CMakeLists.txt demo link directories", testCmake, + "target_link_directories(tilexr_ep_dispatch_demo", "target_link_libraries(tilexr_ep_dispatch_demo", + "-linux/devlib"); + } +} + void TestEpHostChecksRegisteredWorkspace() { std::string launchContext; @@ -115,6 +144,8 @@ void TestEpHostChecksRegisteredWorkspace() CheckContains("src/ep/host/ep_launch_context.cpp", launchContext, "TileXRGetUDMARegistryHost"); CheckContains("src/ep/host/ep_launch_context.cpp", launchContext, "UDMARegionContains"); CheckContains("src/ep/host/ep_launch_context.cpp", launchContext, "TileXREpUdmaRequiredWorkspaceBytes"); + CheckNotContains("src/ep/host/ep_launch_context.cpp", launchContext, + "TileXREpDispatchWorkspaceBytes(context->window, params.tpWorldSize)"); } void TestEpSocDefaultFollowsEnvironment() @@ -242,6 +273,15 @@ void TestDispatchDemoRegistersAlignedUdmaWorkspace() "workspaceDev = reinterpret_cast(AlignAddress("); CheckContains("tests/ep/demo/tilexr_ep_dispatch_demo.cpp", demo, "TileXRUDMARegister(comm, static_cast(workspaceDev), workspaceBytes"); + CheckContains("tests/ep/demo/tilexr_ep_dispatch_demo.cpp", demo, "#include \"ep_transport_route.h\""); + CheckContains("tests/ep/demo/tilexr_ep_dispatch_demo.cpp", demo, + "TileXREp::TileXREpShouldRegisterWorkspace(requestedTransport, *commArgsHost)"); + CheckContains("tests/ep/demo/tilexr_ep_dispatch_demo.cpp", demo, + "TileXREp::TileXREpResolveTransport(requestedTransport, *commArgsHost,"); + CheckContains("tests/ep/demo/tilexr_ep_dispatch_demo.cpp", demo, + "resolvedTransport == TileXR::TileXRTransportKind::DIRECT_URMA"); + CheckNotContains("tests/ep/demo/tilexr_ep_dispatch_demo.cpp", demo, + "TransportNeedsUdmaRegistration"); CheckContains("tests/ep/demo/tilexr_ep_dispatch_demo.cpp", demo, "EpRequiredWorkspaceBytes"); } @@ -255,6 +295,8 @@ void TestDispatchDemoUsesHostBarrierBeforeValidation() CheckContains("tests/ep/demo/tilexr_ep_dispatch_demo.cpp", demo, "DemoBarrierAll"); CheckContains("tests/ep/demo/tilexr_ep_dispatch_demo.cpp", demo, "TILEXR_DEMO_BARRIER_ADDR"); CheckContains("tests/ep/demo/tilexr_ep_dispatch_demo.cpp", demo, "dispatch synchronized"); + CheckContains("tests/ep/demo/tilexr_ep_dispatch_demo.cpp", demo, "DumpCrossNodeDispatchWindow"); + CheckContains("tests/ep/demo/tilexr_ep_dispatch_demo.cpp", demo, "dispatch failure"); } void TestNoForbiddenDependencies() @@ -290,6 +332,7 @@ int main() { TestPublicHeader(); TestBuildPlacement(); + TestEpHostTargetsExcludeCannDevlib(); TestEpHostChecksRegisteredWorkspace(); TestEpSocDefaultFollowsEnvironment(); TestChipMapRecognizesAscend950Dt9582(); diff --git a/tests/ep/unit/test_tilexr_ep_host_validation.cpp b/tests/ep/unit/test_tilexr_ep_host_validation.cpp index d999bbcc..4d3bbd1c 100644 --- a/tests/ep/unit/test_tilexr_ep_host_validation.cpp +++ b/tests/ep/unit/test_tilexr_ep_host_validation.cpp @@ -1,14 +1,58 @@ +#include #include #include #include "comm_args.h" #include "ep_dispatch_host.h" +#include "ep_window.h" +#include "tilexr_udma_reg.h" #include "tilexr_types.h" namespace { +TileXR::CommArgs *g_commArgs = nullptr; +GM_ADDR g_devArgs = reinterpret_cast(0x70000000); +const TileXR::TileXRUDMARegistry *g_registry = nullptr; +int g_registryCalls = 0; + +} // namespace + +extern "C" int TileXRGetCommArgsHost(TileXRCommPtr, TileXR::CommArgs *&commArgsPtr) +{ + commArgsPtr = g_commArgs; + return TileXR::TILEXR_SUCCESS; +} + +extern "C" int TileXRGetCommArgsDev(TileXRCommPtr, GM_ADDR &commArgsPtr) +{ + commArgsPtr = g_devArgs; + return TileXR::TILEXR_SUCCESS; +} + +extern "C" int TileXRGetUDMARegistryHost(TileXRCommPtr, const TileXR::TileXRUDMARegistry **registry) +{ + ++g_registryCalls; + *registry = g_registry; + return TileXR::TILEXR_SUCCESS; +} + +namespace { + int g_failures = 0; +void SetTransportMode(const char *value) +{ +#ifdef _WIN32 + _putenv_s("TILEXR_TRANSPORT_MODE", value == nullptr ? "" : value); +#else + if (value == nullptr) { + unsetenv("TILEXR_TRANSPORT_MODE"); + } else { + setenv("TILEXR_TRANSPORT_MODE", value, 1); + } +#endif +} + void CheckInt(const char *label, int actual, int expected) { if (actual != expected) { @@ -127,7 +171,8 @@ TileXR::CommArgs CrossNodeCommArgs() args.rankSize = 8; args.localRankSize = 4; for (int rank = 0; rank < args.rankSize; ++rank) { - args.peerMems[rank] = reinterpret_cast(0x10000000 + rank * 0x10000000); + args.peerMems[rank] = reinterpret_cast( + static_cast(0x10000000ULL + static_cast(rank) * 0x10000000ULL)); } return args; } @@ -182,12 +227,12 @@ void TestCommValidation() commArgs = ValidCommArgs(); commArgs.peerMems[1] = nullptr; - CheckInt("missing peer mem", TileXREp::TileXREpValidateDispatchConfig(params, commArgs, &window), - TileXR::TILEXR_ERROR_NOT_INITIALIZED); + CheckInt("dispatch shape validation is transport independent", + TileXREp::TileXREpValidateDispatchConfig(params, commArgs, &window), TileXR::TILEXR_SUCCESS); commArgs = CrossNodeCommArgs(); - CheckInt("cross-node dispatch needs udma registry", TileXREp::TileXREpValidateDispatchConfig(params, commArgs, &window), - TileXR::TILEXR_ERROR_NOT_INITIALIZED); + CheckInt("cross-node dispatch supports peer memory without udma", + TileXREp::TileXREpValidateDispatchConfig(params, commArgs, &window), TileXR::TILEXR_SUCCESS); commArgs = CrossNodeUdmaCommArgs(); params = ValidParams(); @@ -237,22 +282,186 @@ void TestCombineValidation() commArgs = ValidCommArgs(); commArgs.peerMems[1] = nullptr; - CheckInt("combine missing peer mem", TileXREp::TileXREpValidateCombineConfig(params, commArgs, &window), - TileXR::TILEXR_ERROR_NOT_INITIALIZED); + CheckInt("combine shape validation is transport independent", + TileXREp::TileXREpValidateCombineConfig(params, commArgs, &window), TileXR::TILEXR_SUCCESS); commArgs = CrossNodeCommArgs(); - CheckInt("cross-node combine needs udma registry", TileXREp::TileXREpValidateCombineConfig(params, commArgs, - &window), TileXR::TILEXR_ERROR_NOT_INITIALIZED); + CheckInt("cross-node combine supports peer memory without udma", + TileXREp::TileXREpValidateCombineConfig(params, commArgs, &window), TileXR::TILEXR_SUCCESS); commArgs = CrossNodeUdmaCommArgs(); - CheckInt("cross-node combine needs workspace", TileXREp::TileXREpValidateCombineConfig(params, commArgs, &window), - TileXR::TILEXR_ERROR_NOT_INITIALIZED); + CheckInt("cross-node combine shape validation is transport independent", + TileXREp::TileXREpValidateCombineConfig(params, commArgs, &window), TileXR::TILEXR_SUCCESS); params.workspace = reinterpret_cast(0x50000000); CheckInt("cross-node combine config", TileXREp::TileXREpValidateCombineConfig(params, commArgs, &window), TileXR::TILEXR_SUCCESS); } +void FillRegistry(TileXR::TileXRUDMARegistry *registry, const TileXR::CommArgs &args, void *localWorkspace) +{ + *registry = TileXR::TileXRUDMARegistry {}; + registry->rankSize = static_cast(args.rankSize); + registry->regionCount = 1; + for (int rank = 0; rank < args.rankSize; ++rank) { + registry->regions[rank].base = rank == args.rank ? static_cast(localWorkspace) : + reinterpret_cast( + static_cast(0x80000000ULL + static_cast(rank) * 0x10000000ULL)); + registry->regions[rank].bytes = 1024ULL * 1024ULL * 1024ULL; + } +} + +void TestLaunchContextUsesResolvedTransport() +{ + TileXREp::EpDispatchParams params = ValidParams(); + TileXR::CommArgs args = CrossNodeUdmaCommArgs(); + g_commArgs = &args; + g_registry = nullptr; + g_registryCalls = 0; + SetTransportMode("memory"); + + TileXREp::EpHostLaunchContext context {}; + CheckInt("forced memory launch context", TileXREp::TileXREpPrepareLaunchContext(params, &context), + TileXR::TILEXR_SUCCESS); + CheckInt("forced memory saved route", static_cast(context.transport), + static_cast(TileXR::TileXRTransportKind::MEMORY)); + CheckInt("forced memory skips registry", g_registryCalls, 0); + + args.peerMems[7] = nullptr; + CheckInt("forced memory requires every peer mapping", TileXREp::TileXREpPrepareLaunchContext(params, &context), + TileXR::TILEXR_ERROR_NOT_INITIALIZED); + args.peerMems[7] = reinterpret_cast(0x80000000); + + SetTransportMode("direct_urma"); + args = CrossNodeCommArgs(); + g_commArgs = &args; + CheckInt("forced direct requires capability", TileXREp::TileXREpPrepareLaunchContext(params, &context), + TileXR::TILEXR_ERROR_NOT_INITIALIZED); + CheckInt("missing capability skips registry", g_registryCalls, 0); + + args = CrossNodeUdmaCommArgs(); + args.peerMems[4] = nullptr; + g_commArgs = &args; + static uint8_t workspace[1024] = {}; + params.workspace = workspace; + TileXR::TileXRUDMARegistry registry {}; + FillRegistry(®istry, args, workspace); + g_registry = ®istry; + TileXREp::EpWindowConfig directWindow {}; + CheckInt("build direct dispatch window", TileXREp::TileXREpValidateDispatchConfig(params, args, &directWindow), + TileXR::TILEXR_SUCCESS); + const uint64_t legacyDispatchBytes = static_cast( + TileXREp::TileXREpAlignUp(directWindow.totalBytes, TileXREp::kEpWindowAlignmentBytes) * 3); + for (int rank = 0; rank < args.rankSize; ++rank) { + registry.regions[rank].bytes = legacyDispatchBytes; + } + CheckInt("direct dispatch rejects workspace without relay and status area", + TileXREp::TileXREpPrepareLaunchContext(params, &context), TileXR::TILEXR_ERROR_PARA_CHECK_FAIL); + + FillRegistry(®istry, args, workspace); + g_registryCalls = 0; + CheckInt("forced direct launch context", TileXREp::TileXREpPrepareLaunchContext(params, &context), + TileXR::TILEXR_SUCCESS); + CheckInt("forced direct saved route", static_cast(context.transport), + static_cast(TileXR::TileXRTransportKind::DIRECT_URMA)); + CheckInt("forced direct checks registry", g_registryCalls, 1); + + args.peerMems[1] = nullptr; + g_registryCalls = 0; + CheckInt("forced direct requires local peer mapping", TileXREp::TileXREpPrepareLaunchContext(params, &context), + TileXR::TILEXR_ERROR_NOT_INITIALIZED); + CheckInt("missing local peer skips registry", g_registryCalls, 0); + + SetTransportMode("invalid"); + CheckInt("invalid transport mode", TileXREp::TileXREpPrepareLaunchContext(params, &context), + TileXR::TILEXR_ERROR_PARA_CHECK_FAIL); + SetTransportMode(nullptr); +} + +void TestCombineContextUsesRouteAwarePeerMappings() +{ + TileXREp::EpCombineParams params = ValidCombineParams(); + static uint8_t workspace[1024] = {}; + params.workspace = workspace; + TileXR::CommArgs args = CrossNodeUdmaCommArgs(); + args.peerMems[4] = nullptr; + g_commArgs = &args; + TileXR::TileXRUDMARegistry registry {}; + FillRegistry(®istry, args, workspace); + g_registry = ®istry; + g_registryCalls = 0; + SetTransportMode("direct_urma"); + + TileXREp::EpHostLaunchContext context {}; + CheckInt("direct combine ignores remote peer mapping", + TileXREp::TileXREpPrepareCombineLaunchContext(params, &context), TileXR::TILEXR_SUCCESS); + CheckInt("direct combine checks registry", g_registryCalls, 1); + + SetTransportMode("memory"); + g_registryCalls = 0; + CheckInt("memory combine requires remote peer mapping", + TileXREp::TileXREpPrepareCombineLaunchContext(params, &context), TileXR::TILEXR_ERROR_NOT_INITIALIZED); + CheckInt("memory combine skips registry", g_registryCalls, 0); + SetTransportMode(nullptr); +} + +void TestDirectTransportChecksRegistryOnSameNode() +{ + TileXREp::EpDispatchParams params = ValidParams(); + static uint8_t workspace[1024] = {}; + params.workspace = workspace; + TileXR::CommArgs args = ValidCommArgs(); + args.extraFlag |= TileXR::ExtraFlag::UDMA; + args.udmaInfoPtr = reinterpret_cast(0x30000000); + args.udmaRegistryPtr = reinterpret_cast(0x40000000); + g_commArgs = &args; + TileXR::TileXRUDMARegistry registry {}; + FillRegistry(®istry, args, workspace); + g_registry = ®istry; + g_registryCalls = 0; + SetTransportMode("direct_urma"); + + TileXREp::EpHostLaunchContext context {}; + CheckInt("same-node direct launch context", TileXREp::TileXREpPrepareLaunchContext(params, &context), + TileXR::TILEXR_SUCCESS); + CheckInt("same-node direct checks registry", g_registryCalls, 1); + CheckInt("same-node direct saved route", static_cast(context.transport), + static_cast(TileXR::TileXRTransportKind::DIRECT_URMA)); + SetTransportMode(nullptr); +} + +void TestAutoWithoutWorkspaceStaysOnMemory() +{ + TileXREp::EpDispatchParams params = ValidParams(); + params.bs = 64; + params.h = 128; + params.workspace = nullptr; + TileXR::CommArgs args = CrossNodeUdmaCommArgs(); + g_commArgs = &args; + g_registry = nullptr; + g_registryCalls = 0; + SetTransportMode("auto"); + + TileXREp::EpHostLaunchContext context {}; + CheckInt("auto without workspace uses memory", + TileXREp::TileXREpPrepareLaunchContext(params, &context), TileXR::TILEXR_SUCCESS); + CheckInt("auto without workspace saves memory route", static_cast(context.transport), + static_cast(TileXR::TileXRTransportKind::MEMORY)); + CheckInt("auto without workspace skips registry", g_registryCalls, 0); + + TileXREp::EpCombineParams combineParams = ValidCombineParams(); + combineParams.bs = 64; + combineParams.h = 128; + combineParams.workspace = nullptr; + g_registryCalls = 0; + CheckInt("auto combine without workspace uses memory", + TileXREp::TileXREpPrepareCombineLaunchContext(combineParams, &context), TileXR::TILEXR_SUCCESS); + CheckInt("auto combine without workspace saves memory route", static_cast(context.transport), + static_cast(TileXR::TileXRTransportKind::MEMORY)); + CheckInt("auto combine without workspace skips registry", g_registryCalls, 0); + SetTransportMode(nullptr); +} + void TestV2CapabilityValidation() { TileXREp::EpDispatchParams params = ValidV2Params(); @@ -377,7 +586,8 @@ void TestV2CapabilityValidation() commArgs.rankSize = 8; commArgs.localRankSize = 8; for (int rank = 0; rank < commArgs.rankSize; ++rank) { - commArgs.peerMems[rank] = reinterpret_cast(0x10000000 + rank * 0x10000000); + commArgs.peerMems[rank] = reinterpret_cast( + static_cast(0x10000000ULL + static_cast(rank) * 0x10000000ULL)); } CheckInt("v2 supports tp with shared expert", TileXREp::TileXREpValidateDispatchV2Config(params, commArgs), TileXR::TILEXR_SUCCESS); @@ -405,6 +615,10 @@ int main() TestBasicValidation(); TestCommValidation(); TestCombineValidation(); + TestLaunchContextUsesResolvedTransport(); + TestCombineContextUsesRouteAwarePeerMappings(); + TestDirectTransportChecksRegistryOnSameNode(); + TestAutoWithoutWorkspaceStaysOnMemory(); TestV2CapabilityValidation(); return g_failures == 0 ? 0 : 1; } diff --git a/tests/ep/unit/test_tilexr_ep_kernel_sources.cpp b/tests/ep/unit/test_tilexr_ep_kernel_sources.cpp index a8100561..112ff1e7 100644 --- a/tests/ep/unit/test_tilexr_ep_kernel_sources.cpp +++ b/tests/ep/unit/test_tilexr_ep_kernel_sources.cpp @@ -99,6 +99,7 @@ void TestCrossNodeDispatchUsesUDMARegistry() CheckContains(path, contents, "tilexr_udma.h"); CheckContains(path, contents, "TileXR::UDMARegistryEnabled(args)"); CheckContains(path, contents, "TileXREpUsesUdmaWindow"); + CheckContains(path, contents, "TileXREpLoadLocalPeerMems"); CheckContains(path, contents, "TileXR::UDMAPutNbi"); CheckContains(path, contents, "TileXR::UDMAQuiet(args, dstRank)"); CheckContains(path, contents, "TileXREpFlushDispatchSlotHeaders"); @@ -115,12 +116,22 @@ void TestCrossNodeDispatchPullsRemoteSlots() CheckContains(path, contents, "tilexr_ep_dispatch_cross_node_kernel"); CheckContains(path, contents, "launch_tilexr_ep_dispatch_cross_node_kernel"); + CheckNotContains(path, contents, "tilexr_ep_dispatch_cross_node_smoke_kernel"); + CheckNotContains(path, contents, "launch_tilexr_ep_dispatch_cross_node_smoke_kernel"); CheckContains(path, contents, "TileXREpPullUdmaSlots"); CheckContains(path, contents, "TileXR::UDMAGetNbi"); CheckContains(path, contents, "TileXREpNotifyUdmaReady"); - CheckContains(path, contents, "TileXREpWaitUdmaReady"); + CheckContains(path, contents, "bool TileXREpWaitUdmaReady"); CheckContains(path, contents, "TileXREpNotifyAllUdmaReady"); - CheckContains(path, contents, "TileXREpWaitAllUdmaReady"); + CheckContains(path, contents, "bool TileXREpWaitAllUdmaReady"); + CheckContains(path, contents, "TileXREpInvalidateLocalCacheLines"); + CheckContains(path, contents, "TileXR::TILEXR_UDMA_MAX_RETRY_TIMES"); + CheckContains(path, contents, "tilexr_ep_udma_ready timeout"); + CheckContains(path, contents, "tilexr_ep_udma_all_ready timeout"); + CheckContains(path, contents, "if (!TileXREpWaitAllUdmaReady"); + CheckContains(path, contents, "kEpStatusDispatchReadyTimeout"); + CheckContains(path, contents, "kEpStatusDispatchSlotTimeout"); + CheckContains(path, contents, "TileXREpStoreStatusValue(workspaceGM"); CheckContains(path, contents, "TileXR::UDMAPutSignalNbi"); } @@ -141,7 +152,7 @@ void TestCrossNodeDispatchSeparatesLocalAndRemotePeers() CheckContains(path, contents, "if (localRankSize > 1)"); CheckContains(path, contents, "dstRank != rank && TileXREpIsSameNodePeer(rank, dstRank, localRankSize)"); CheckContains(path, contents, "srcRank != rank &&"); - CheckContains(path, contents, "TileXREpIsSameNodePeer(rank, srcRank, localRankSize)"); + CheckContains(path, contents, "TileXREpIsSameNodePeer(rank, srcRank, effectiveLocalRankSize)"); CheckContains(path, contents, "sameNodeSource ? rank : srcRank"); CheckContains(path, contents, "!TileXREpIsSameNodePeer(rank, peer, localRankSize)"); } @@ -155,7 +166,45 @@ void TestHostDispatchSplitsCrossNodeKernel() } CheckContains(path, contents, "launch_tilexr_ep_dispatch_cross_node_kernel"); - CheckContains(path, contents, "TileXREpUsesCrossNodeKernel"); + CheckContains(path, contents, "TileXREpUsesDirectUdmaKernel"); + CheckContains(path, contents, "context.transport =="); + CheckContains(path, contents, "TileXR::TileXRTransportKind::DIRECT_URMA"); + CheckContains(path, contents, "GM_ADDR memoryWorkspace = nullptr;"); + CheckContains(path, contents, "TileXREpCheckUdmaStatus(params.stream, params.workspace, context.window)"); + CheckNotContains(path, contents, "TileXRSelectAutoTransport"); + CheckNotContains(path, contents, "TILEXR_EP_DISPATCH_CROSS_NODE_SMOKE"); + CheckNotContains(path, contents, "launch_tilexr_ep_dispatch_cross_node_smoke_kernel"); + CheckNotContains(path, contents, "return TileXR::TILEXR_SUCCESS;\n launch_tilexr_ep_dispatch_cross_node_kernel"); +} + +void TestSameNodeDirectUrmaUsesRegisteredWorkspace() +{ + std::string dispatch; + if (ReadFile("src/ep/kernels/tilexr_ep_dispatch_kernel.cpp", &dispatch)) { + CheckContains("src/ep/kernels/tilexr_ep_dispatch_kernel.cpp", dispatch, "useUdmaForAllPeers"); + CheckContains("src/ep/kernels/tilexr_ep_dispatch_kernel.cpp", dispatch, "effectiveLocalRankSize"); + CheckContains("src/ep/kernels/tilexr_ep_dispatch_kernel.cpp", dispatch, + "useUdmaForAllPeers ? nullptr : localIpcWindow"); + } + + std::string combine; + if (ReadFile("src/ep/kernels/tilexr_ep_combine_kernel.cpp", &combine)) { + CheckContains("src/ep/kernels/tilexr_ep_combine_kernel.cpp", combine, "useUdmaForAllPeers"); + CheckContains("src/ep/kernels/tilexr_ep_combine_kernel.cpp", combine, "effectiveLocalRankSize"); + CheckContains("src/ep/kernels/tilexr_ep_combine_kernel.cpp", combine, + "if (useUdmaForAllPeers && dstRank == rank)"); + } +} + +void TestCrossNodeDispatchHasNoTemporaryEarlyReturn() +{ + const std::string path = "src/ep/kernels/tilexr_ep_dispatch_kernel.cpp"; + std::string contents; + if (!ReadFile(path, &contents)) { + return; + } + + CheckNotContains(path, contents, "const int32_t localRankSize = args->localRankSize;\n return;"); } void TestDispatchHelpersLiveInDispatchHelperFile() @@ -188,6 +237,9 @@ void TestCombineHelpersLiveInCombineHelperFile() CheckContains(path, contents, "TileXREpGetCombineTokenId"); CheckContains(path, contents, "TileXREpGetCombineTopKId"); CheckContains(path, contents, "sourceWindow + SlotOffset(slotRank, slotBytes)"); + CheckContains(path, contents, "TileXREpInvalidateLocalCacheLines(slotGM, TileXREp::kEpSrcSlotHeaderBytes)"); + CheckContains(path, contents, "TileXR::TILEXR_UDMA_MAX_RETRY_TIMES"); + CheckContains(path, contents, "tilexr_ep_dispatch_slot_ready timeout"); } void TestCombineKernelUsesTileXRPeerMemory() @@ -209,6 +261,9 @@ void TestCombineKernelUsesTileXRPeerMemory() CheckContains(path, contents, "tilexr_ep_combine_cross_node_drain_kernel"); CheckContains(path, contents, "TileXREpNotifyRemoteUdmaReadySeparate"); CheckContains(path, contents, "TileXREpWaitRemoteUdmaReady"); + CheckContains(path, contents, "TileXREpLoadLocalPeerMems"); + CheckContains(path, contents, "TileXR::UDMARegistryEnabled(args)"); + CheckNotContains(path, contents, "TileXREpUsesDirectUrmaTransport"); CheckNotContains(path, contents, "tilexr_ep_dispatch_kernel"); std::string hostLaunch; @@ -233,6 +288,9 @@ void TestKernelCommonHasCombineHelpers() CheckContains(path, contents, "UDMASecondOperationOffset"); CheckContains(path, contents, "TileXREpNotifyRemoteUdmaReadySeparate"); CheckContains(path, contents, "TileXREpWaitRemoteUdmaReady"); + CheckContains(path, contents, "TileXR::UDMAPutNbi"); + CheckNotContains(path, contents, "TileXRPutAutoNbi"); + CheckNotContains(path, contents, "TileXRSelectAutoTransport"); CheckContains(path, contents, "TileXREpStoreStatusValue"); CheckContains(path, contents, "TileXREpFlushUdmaSourceWindow"); CheckContains(path, contents, "IsValidShape"); @@ -254,6 +312,7 @@ void TestDispatchDemoRunsCombine() CheckContains(path, contents, "TileXRMoeEpCombine"); CheckContains(path, contents, "ValidateCombineOutputs"); CheckContains(path, contents, "combine validation"); + CheckContains(path, contents, "TILEXR_EP_DEMO_BS"); } void TestKernelForwardsActiveMask() @@ -417,6 +476,34 @@ void TestClearLocalWindowDoesNotPreclearSlotHeaders() " }"); } +void TestUdmaReadyFlagsUseCacheLineStride() +{ + const std::vector kernelPaths = { + "src/ep/kernels/tilexr_ep_kernel_common.h", + "src/ep/kernels/tilexr_ep_dispatch_kernel.cpp", + }; + for (std::vector::const_iterator path = kernelPaths.begin(); path != kernelPaths.end(); ++path) { + std::string contents; + if (!ReadFile(*path, &contents)) { + continue; + } + CheckContains(*path, contents, + "static_cast(rank) * TileXREp::kEpUdmaReadyStrideBytes"); + } + + std::string common; + if (ReadFile("src/ep/kernels/tilexr_ep_kernel_common.h", &common)) { + CheckContains("src/ep/kernels/tilexr_ep_kernel_common.h", common, + "static_cast(srcRank) * TileXREp::kEpUdmaReadyStrideBytes"); + } + + std::string demo; + if (ReadFile("tests/ep/demo/tilexr_ep_dispatch_demo.cpp", &demo)) { + CheckContains("tests/ep/demo/tilexr_ep_dispatch_demo.cpp", demo, + "static_cast(rankSize) * kUdmaReadyStrideBytes"); + } +} + void TestNoForbiddenDependencies() { const std::vector paths = { @@ -455,6 +542,8 @@ int main() TestCrossNodeDispatchPullsRemoteSlots(); TestCrossNodeDispatchSeparatesLocalAndRemotePeers(); TestHostDispatchSplitsCrossNodeKernel(); + TestSameNodeDirectUrmaUsesRegisteredWorkspace(); + TestCrossNodeDispatchHasNoTemporaryEarlyReturn(); TestDispatchHelpersLiveInDispatchHelperFile(); TestCombineHelpersLiveInCombineHelperFile(); TestCombineKernelUsesTileXRPeerMemory(); @@ -468,6 +557,7 @@ int main() TestKernelForwardsStaticQuantConfig(); TestKernelForwardsPerTokenDynamicQuantConfig(); TestClearLocalWindowDoesNotPreclearSlotHeaders(); + TestUdmaReadyFlagsUseCacheLineStride(); TestNoForbiddenDependencies(); if (g_failures != 0) { std::cerr << g_failures << " TileXR EP kernel source checks failed" << std::endl; diff --git a/tests/ep/unit/test_tilexr_ep_layout.cpp b/tests/ep/unit/test_tilexr_ep_layout.cpp index 69c0e7b9..ec303931 100644 --- a/tests/ep/unit/test_tilexr_ep_layout.cpp +++ b/tests/ep/unit/test_tilexr_ep_layout.cpp @@ -79,6 +79,15 @@ void TestWindowConfig() CheckInt64("total bytes", config.totalBytes, 704); } +void TestUdmaWorkspaceUsesCacheLineReadySlots() +{ + CheckInt64("2-rank operation bytes", TileXREp::TileXREpUdmaOperationBytes(704, 2, 320), 2944); + CheckInt64("2-rank workspace bytes", TileXREp::TileXREpUdmaRequiredWorkspaceBytes(704, 2, 320), 5952); + CheckInt64("4-rank operation bytes", TileXREp::TileXREpUdmaOperationBytes(131392, 4, 32832), 788608); + CheckInt64("4-rank workspace bytes", + TileXREp::TileXREpUdmaRequiredWorkspaceBytes(131392, 4, 32832), 1577280); +} + void TestRejectsInvalidConfig() { TileXREp::EpWindowConfig config {}; @@ -107,6 +116,7 @@ int main() TestExpertMapping(); TestDataTypes(); TestWindowConfig(); + TestUdmaWorkspaceUsesCacheLineReadySlots(); TestRejectsInvalidConfig(); return g_failures == 0 ? 0 : 1; } diff --git a/tests/ep/unit/test_tilexr_ep_transport_route.cpp b/tests/ep/unit/test_tilexr_ep_transport_route.cpp new file mode 100644 index 00000000..0f23add8 --- /dev/null +++ b/tests/ep/unit/test_tilexr_ep_transport_route.cpp @@ -0,0 +1,174 @@ +#include +#include +#include + +#include "ep_transport_route.h" +#include "tilexr_types.h" + +namespace { + +int g_failures = 0; + +#define CHECK_EQ(lhs, rhs) \ + do { \ + auto lhsValue = (lhs); \ + auto rhsValue = (rhs); \ + if (lhsValue != rhsValue) { \ + std::cerr << "CHECK_EQ failed at line " << __LINE__ << ": " #lhs " != " #rhs \ + << " (" << static_cast(lhsValue) << " vs " << static_cast(rhsValue) << ")" \ + << std::endl; \ + ++g_failures; \ + } \ + } while (0) + +TileXR::CommArgs MakeArgs(bool directUrmaAvailable) +{ + TileXR::CommArgs args {}; + args.rankSize = 2; + args.localRankSize = 1; + if (directUrmaAvailable) { + args.extraFlag |= TileXR::ExtraFlag::UDMA; + args.udmaInfoPtr = reinterpret_cast(0x100000); + args.udmaRegistryPtr = reinterpret_cast(0x200000); + } + return args; +} + +TileXR::CommArgs MakeUdmaCapableArgs() +{ + TileXR::CommArgs args {}; + args.rankSize = 2; + args.localRankSize = 1; + args.extraFlag |= TileXR::ExtraFlag::UDMA; + args.udmaInfoPtr = reinterpret_cast(0x100000); + return args; +} + +void SetTransportMode(const char *value) +{ +#ifdef _WIN32 + _putenv_s("TILEXR_TRANSPORT_MODE", value == nullptr ? "" : value); +#else + if (value == nullptr) { + unsetenv("TILEXR_TRANSPORT_MODE"); + } else { + setenv("TILEXR_TRANSPORT_MODE", value, 1); + } +#endif +} + +void TestParseTransportMode() +{ + TileXREp::EpTransportMode mode = TileXREp::EpTransportMode::MEMORY; + CHECK_EQ(TileXREp::TileXREpParseTransportMode(nullptr, &mode), TileXR::TILEXR_SUCCESS); + CHECK_EQ(mode, TileXREp::EpTransportMode::AUTO); + + CHECK_EQ(TileXREp::TileXREpParseTransportMode("", &mode), TileXR::TILEXR_SUCCESS); + CHECK_EQ(mode, TileXREp::EpTransportMode::AUTO); + CHECK_EQ(TileXREp::TileXREpParseTransportMode("auto", &mode), TileXR::TILEXR_SUCCESS); + CHECK_EQ(mode, TileXREp::EpTransportMode::AUTO); + CHECK_EQ(TileXREp::TileXREpParseTransportMode("memory", &mode), TileXR::TILEXR_SUCCESS); + CHECK_EQ(mode, TileXREp::EpTransportMode::MEMORY); + CHECK_EQ(TileXREp::TileXREpParseTransportMode("direct_urma", &mode), TileXR::TILEXR_SUCCESS); + CHECK_EQ(mode, TileXREp::EpTransportMode::DIRECT_URMA); + CHECK_EQ(TileXREp::TileXREpParseTransportMode("invalid", &mode), TileXR::TILEXR_ERROR_PARA_CHECK_FAIL); + CHECK_EQ(TileXREp::TileXREpParseTransportMode("auto", nullptr), TileXR::TILEXR_ERROR_PARA_CHECK_FAIL); +} + +void TestResolveForcedTransport() +{ + const TileXR::CommArgs noUdmaArgs = MakeArgs(false); + const TileXR::CommArgs udmaArgs = MakeArgs(true); + TileXR::TileXRTransportKind transport = TileXR::TileXRTransportKind::DIRECT_URMA; + + CHECK_EQ(TileXREp::TileXREpResolveTransport(TileXREp::EpTransportMode::MEMORY, noUdmaArgs, 4096, + &transport), TileXR::TILEXR_SUCCESS); + CHECK_EQ(transport, TileXR::TileXRTransportKind::MEMORY); + + CHECK_EQ(TileXREp::TileXREpResolveTransport(TileXREp::EpTransportMode::DIRECT_URMA, noUdmaArgs, 4096, + &transport), TileXR::TILEXR_ERROR_NOT_INITIALIZED); + CHECK_EQ(TileXREp::TileXREpResolveTransport(TileXREp::EpTransportMode::DIRECT_URMA, udmaArgs, 4096, + &transport), TileXR::TILEXR_SUCCESS); + CHECK_EQ(transport, TileXR::TileXRTransportKind::DIRECT_URMA); + CHECK_EQ(TileXREp::TileXREpResolveTransport(TileXREp::EpTransportMode::AUTO, udmaArgs, 4096, + &transport), TileXR::TILEXR_SUCCESS); + CHECK_EQ(transport, TileXR::TileXRTransportKind::MEMORY); + CHECK_EQ(TileXREp::TileXREpResolveTransport(TileXREp::EpTransportMode::AUTO, udmaArgs, + TileXR::TILEXR_AUTO_CROSS_NODE_DIRECT_URMA_THRESHOLD_BYTES, &transport), TileXR::TILEXR_SUCCESS); + CHECK_EQ(transport, TileXR::TileXRTransportKind::DIRECT_URMA); + CHECK_EQ(TileXREp::TileXREpResolveTransport(TileXREp::EpTransportMode::AUTO, noUdmaArgs, 4096, + &transport), TileXR::TILEXR_SUCCESS); + CHECK_EQ(transport, TileXR::TileXRTransportKind::MEMORY); + CHECK_EQ(TileXREp::TileXREpResolveTransport(TileXREp::EpTransportMode::AUTO, udmaArgs, 4096, + nullptr), TileXR::TILEXR_ERROR_PARA_CHECK_FAIL); +} + +void TestSameNodeUsesFourMiBThreshold() +{ + TileXR::CommArgs args = MakeArgs(true); + args.localRankSize = args.rankSize; + TileXR::TileXRTransportKind transport = TileXR::TileXRTransportKind::DIRECT_URMA; + + CHECK_EQ(TileXREp::TileXREpResolveTransport(TileXREp::EpTransportMode::AUTO, args, + TileXR::TILEXR_AUTO_SAME_NODE_DIRECT_URMA_THRESHOLD_BYTES - 1, &transport), TileXR::TILEXR_SUCCESS); + CHECK_EQ(transport, TileXR::TileXRTransportKind::MEMORY); + CHECK_EQ(TileXREp::TileXREpResolveTransport(TileXREp::EpTransportMode::AUTO, args, + TileXR::TILEXR_AUTO_SAME_NODE_DIRECT_URMA_THRESHOLD_BYTES, &transport), TileXR::TILEXR_SUCCESS); + CHECK_EQ(transport, TileXR::TileXRTransportKind::DIRECT_URMA); + CHECK_EQ(TileXREp::TileXREpResolveTransport(TileXREp::EpTransportMode::DIRECT_URMA, args, + 8ULL * 1024ULL * 1024ULL, &transport), TileXR::TILEXR_SUCCESS); + CHECK_EQ(transport, TileXR::TileXRTransportKind::DIRECT_URMA); +} + +void TestResolveTransportFromEnvironment() +{ + const TileXR::CommArgs noUdmaArgs = MakeArgs(false); + TileXR::TileXRTransportKind transport = TileXR::TileXRTransportKind::DIRECT_URMA; + + SetTransportMode(nullptr); + CHECK_EQ(TileXREp::TileXREpResolveTransportFromEnv(noUdmaArgs, 8ULL * 1024ULL * 1024ULL, + &transport), TileXR::TILEXR_SUCCESS); + CHECK_EQ(transport, TileXR::TileXRTransportKind::MEMORY); + + SetTransportMode("auto"); + CHECK_EQ(TileXREp::TileXREpResolveTransportFromEnv(noUdmaArgs, 8ULL * 1024ULL * 1024ULL, + &transport), TileXR::TILEXR_SUCCESS); + CHECK_EQ(transport, TileXR::TileXRTransportKind::MEMORY); + + SetTransportMode("direct_urma"); + CHECK_EQ(TileXREp::TileXREpResolveTransportFromEnv(noUdmaArgs, 8ULL * 1024ULL * 1024ULL, + &transport), TileXR::TILEXR_ERROR_NOT_INITIALIZED); + + SetTransportMode("invalid"); + CHECK_EQ(TileXREp::TileXREpResolveTransportFromEnv(noUdmaArgs, 8ULL * 1024ULL * 1024ULL, + &transport), TileXR::TILEXR_ERROR_PARA_CHECK_FAIL); + SetTransportMode(nullptr); +} + +void TestDemoRegistrationDecisionUsesCapabilityBeforeRegistry() +{ + const TileXR::CommArgs noUdmaArgs = MakeArgs(false); + const TileXR::CommArgs capableArgs = MakeUdmaCapableArgs(); + + CHECK_EQ(TileXREp::TileXREpShouldRegisterWorkspace(TileXREp::EpTransportMode::AUTO, capableArgs), true); + CHECK_EQ(TileXREp::TileXREpShouldRegisterWorkspace(TileXREp::EpTransportMode::DIRECT_URMA, capableArgs), true); + CHECK_EQ(TileXREp::TileXREpShouldRegisterWorkspace(TileXREp::EpTransportMode::MEMORY, capableArgs), false); + CHECK_EQ(TileXREp::TileXREpShouldRegisterWorkspace(TileXREp::EpTransportMode::AUTO, noUdmaArgs), false); +} + +} // namespace + +int main() +{ + TestParseTransportMode(); + TestResolveForcedTransport(); + TestSameNodeUsesFourMiBThreshold(); + TestResolveTransportFromEnvironment(); + TestDemoRegistrationDecisionUsesCapabilityBeforeRegistry(); + if (g_failures != 0) { + std::cerr << g_failures << " TileXR EP transport route checks failed" << std::endl; + return 1; + } + std::cout << "TileXR EP transport route checks passed" << std::endl; + return 0; +} diff --git a/tests/memory/CMakeLists.txt b/tests/memory/CMakeLists.txt index bc49e6ad..8cd920cf 100644 --- a/tests/memory/CMakeLists.txt +++ b/tests/memory/CMakeLists.txt @@ -56,7 +56,6 @@ include_directories( link_directories( ${ASCEND_DRIVER_PATH}/lib64/driver ${ASCEND_HOME_PATH}/${ARCH}-linux/lib64 - ${ASCEND_HOME_PATH}/${ARCH}-linux/devlib ) add_executable(test_tilexr_memory_demo_sources diff --git a/tests/memory/README.md b/tests/memory/README.md index b9fecff2..d84198d6 100644 --- a/tests/memory/README.md +++ b/tests/memory/README.md @@ -2,7 +2,9 @@ This directory contains a small memory-semantics example modeled after the reference-only `ascend-transformer-boost/src/kernels/lcal/src/kernels/lcal_allgather.cce` source. -The example uses TileXR `CommArgs::peerMems[]` as shared peer-memory windows and moves data through UB with Ascend C `DataCopy`. It intentionally does not use `TileXRUDMARegister`, `tilexr_udma.h`, or UDMA put/get APIs. +The default path uses TileXR `CommArgs::peerMems[]` as shared peer-memory windows and moves data through UB with Ascend C `DataCopyPad`. It intentionally does not use `TileXRUDMARegister`, `tilexr_udma.h`, or UDMA put/get APIs. + +TCP sockets are used only for communicator rendezvous and demo barriers. They do not carry payload on the default path. ## Build @@ -42,3 +44,21 @@ run_tilexr_memory_demo.sh 0) { + host = value.substr(0, colon); + basePort = parsedPort; + } } } - int barrierPort = basePort + kDemoBarrierPortOffset; - if (barrierPort <= 0 || barrierPort > 65535) { - barrierPort = kDefaultCommPort + kDemoBarrierPortOffset; + int port = basePort + portOffset; + if (port <= 0 || port > 65535) { + port = kDefaultCommPort + portOffset; } - return BarrierEndpoint{static_cast(barrierPort)}; + return BarrierEndpoint{host, static_cast(port)}; +} + +BarrierEndpoint GetBarrierEndpoint() +{ + return GetEndpoint(kDemoBarrierPortOffset); +} + +BarrierEndpoint GetHostExchangeEndpoint() +{ + return GetEndpoint(kDemoHostExchangePortOffset); } bool SendAll(int fd, const void* data, size_t bytes) @@ -177,7 +204,7 @@ int CreateBarrierServer(uint16_t port) } sockaddr_in addr{}; addr.sin_family = AF_INET; - addr.sin_addr.s_addr = htonl(INADDR_LOOPBACK); + addr.sin_addr.s_addr = htonl(INADDR_ANY); addr.sin_port = htons(port); if (bind(fd, reinterpret_cast(&addr), sizeof(addr)) != 0 || listen(fd, SOMAXCONN) != 0) { @@ -187,12 +214,14 @@ int CreateBarrierServer(uint16_t port) return fd; } -int ConnectBarrierServer(uint16_t port) +int ConnectBarrierServer(const BarrierEndpoint& endpoint) { sockaddr_in addr{}; addr.sin_family = AF_INET; - addr.sin_addr.s_addr = htonl(INADDR_LOOPBACK); - addr.sin_port = htons(port); + addr.sin_port = htons(endpoint.port); + if (inet_pton(AF_INET, endpoint.host.c_str(), &addr.sin_addr) != 1) { + return -1; + } for (int attempt = 0; attempt < kConnectRetryCount; ++attempt) { int fd = socket(AF_INET, SOCK_STREAM, 0); @@ -215,14 +244,15 @@ bool DemoBarrierAll(int rank, int rankSize, const std::string& step) } BarrierEndpoint endpoint = GetBarrierEndpoint(); - PrintStatus(rank, "demo tcp barrier begin: " + step + " port=" + std::to_string(endpoint.port)); + PrintStatus(rank, "demo tcp barrier begin: " + step + " endpoint=" + endpoint.host + ":" + + std::to_string(endpoint.port)); constexpr uint8_t kArrive = 1; constexpr uint8_t kRelease = 2; if (rank == 0) { int serverFd = CreateBarrierServer(endpoint.port); if (serverFd < 0) { - std::cerr << "[rank " << rank << "] ERROR: failed to create demo barrier server on 127.0.0.1:" + std::cerr << "[rank " << rank << "] ERROR: failed to create demo barrier server on 0.0.0.0:" << endpoint.port << ", errno=" << errno << std::endl; return false; } @@ -253,10 +283,10 @@ bool DemoBarrierAll(int rank, int rankSize, const std::string& step) return false; } } else { - int fd = ConnectBarrierServer(endpoint.port); + int fd = ConnectBarrierServer(endpoint); if (fd < 0) { - std::cerr << "[rank " << rank << "] ERROR: failed to connect demo barrier on 127.0.0.1:" - << endpoint.port << std::endl; + std::cerr << "[rank " << rank << "] ERROR: failed to connect demo barrier on " << endpoint.host + << ":" << endpoint.port << std::endl; return false; } uint8_t release = 0; @@ -272,6 +302,83 @@ bool DemoBarrierAll(int rank, int rankSize, const std::string& step) return true; } +bool ExchangeInputSegmentsOnHost( + int rank, int rankSize, const std::vector& hostInput, std::vector& hostOutput, + int32_t elementsPerRank) +{ + if (rankSize <= 0 || elementsPerRank <= 0 || + hostInput.size() != static_cast(elementsPerRank) || + hostOutput.size() != static_cast(rankSize) * elementsPerRank) { + std::cerr << "[rank " << rank << "] ERROR: invalid diagnostic host staging shape" << std::endl; + return false; + } + + BarrierEndpoint endpoint = GetHostExchangeEndpoint(); + const size_t segmentBytes = static_cast(elementsPerRank) * sizeof(int32_t); + const size_t outputBytes = hostOutput.size() * sizeof(int32_t); + PrintStatus(rank, "diagnostic host staging begin endpoint=" + endpoint.host + ":" + + std::to_string(endpoint.port)); + + if (rank == 0) { + std::copy(hostInput.begin(), hostInput.end(), hostOutput.begin()); + int serverFd = CreateBarrierServer(endpoint.port); + if (serverFd < 0) { + std::cerr << "[rank " << rank << "] ERROR: failed to create host exchange server on 0.0.0.0:" + << endpoint.port << ", errno=" << errno << std::endl; + return false; + } + + std::vector clients; + clients.reserve(static_cast(rankSize - 1)); + bool ok = true; + for (int peer = 1; peer < rankSize; ++peer) { + int clientFd = accept(serverFd, nullptr, nullptr); + if (clientFd < 0) { + ok = false; + break; + } + int32_t peerRank = -1; + if (!RecvAll(clientFd, &peerRank, sizeof(peerRank)) || peerRank <= 0 || peerRank >= rankSize || + !RecvAll(clientFd, hostOutput.data() + static_cast(peerRank) * elementsPerRank, + segmentBytes)) { + close(clientFd); + ok = false; + break; + } + clients.push_back(clientFd); + } + + for (int clientFd : clients) { + ok = SendAll(clientFd, hostOutput.data(), outputBytes) && ok; + close(clientFd); + } + close(serverFd); + if (!ok) { + std::cerr << "[rank " << rank << "] ERROR: diagnostic host staging exchange failed" << std::endl; + return false; + } + } else { + int fd = ConnectBarrierServer(endpoint); + if (fd < 0) { + std::cerr << "[rank " << rank << "] ERROR: failed to connect host exchange on " << endpoint.host + << ":" << endpoint.port << std::endl; + return false; + } + int32_t rankValue = rank; + bool ok = SendAll(fd, &rankValue, sizeof(rankValue)) && + SendAll(fd, hostInput.data(), segmentBytes) && + RecvAll(fd, hostOutput.data(), outputBytes); + close(fd); + if (!ok) { + std::cerr << "[rank " << rank << "] ERROR: diagnostic host staging exchange failed" << std::endl; + return false; + } + } + + PrintStatus(rank, "diagnostic host staging end"); + return true; +} + void PrintCommArgs(int rank, const TileXR::CommArgs& args, GM_ADDR commArgsDev) { std::cout << "[rank " << rank << "] CommArgs host fields:" << std::endl; @@ -283,6 +390,37 @@ void PrintCommArgs(int rank, const TileXR::CommArgs& args, GM_ADDR commArgsDev) } } +bool PushInputToPeerWindowsOnHost( + int rank, int rankSize, const TileXR::CommArgs& args, const int32_t* input, int32_t elementsPerRank) +{ + const size_t segmentBytes = static_cast(elementsPerRank) * sizeof(int32_t); + for (int dstRank = 0; dstRank < rankSize; ++dstRank) { + if (args.peerMems[dstRank] == nullptr) { + std::cerr << "[rank " << rank << "] ERROR: peerMems[" << dstRank << "] is null" << std::endl; + return false; + } + GM_ADDR dst = args.peerMems[dstRank] + TileXR::IPC_DATA_OFFSET + + static_cast(rank) * segmentBytes; + if (!CopyDeviceToDevice(rank, dst, segmentBytes, input, segmentBytes, + "input to peer window dstRank=" + std::to_string(dstRank))) { + return false; + } + } + return true; +} + +bool CollectLocalWindowOnHost( + int rank, int rankSize, const TileXR::CommArgs& args, int32_t* output, int32_t elementsPerRank) +{ + if (args.peerMems[rank] == nullptr) { + std::cerr << "[rank " << rank << "] ERROR: local peerMems[" << rank << "] is null" << std::endl; + return false; + } + const size_t outputBytes = static_cast(rankSize) * elementsPerRank * sizeof(int32_t); + GM_ADDR src = args.peerMems[rank] + TileXR::IPC_DATA_OFFSET; + return CopyDeviceToDevice(rank, output, outputBytes, src, outputBytes, "local peer window to output"); +} + bool ValidateData(int rank, int rankSize, const std::vector& output, int32_t elementsPerRank) { bool ok = true; @@ -427,18 +565,72 @@ int main(int argc, char** argv) return 1; } - uint32_t blockDim = static_cast(std::max(kDefaultBlockDim, rankSize)); - PrintStatus(rank, "launch memory all-gather kernel blockDim=" + std::to_string(blockDim)); - launch_tilexr_memory_all_gather(blockDim, stream, commArgsDev, reinterpret_cast(input), - reinterpret_cast(output), reinterpret_cast(debug), elementsPerRank); - if (!CheckAcl(rank, "aclrtSynchronizeStream", aclrtSynchronizeStream(stream))) { - Cleanup(comm, stream, input, output, debug, rank, deviceId); - return 1; - } + const bool useHostStaging = GetEnvInt("TILEXR_MEMORY_DEMO_HOST_STAGING", 0) != 0; + const bool useHostPeerCopy = !useHostStaging && GetEnvInt("TILEXR_MEMORY_DEMO_HOST_COPY", 0) != 0; + if (useHostStaging) { + if (!ExchangeInputSegmentsOnHost(rank, rankSize, hostInput, hostOutput, elementsPerRank)) { + Cleanup(comm, stream, input, output, debug, rank, deviceId); + return 1; + } + if (!CopyHostToDevice(rank, output, outputCount * sizeof(int32_t), hostOutput.data(), + hostOutput.size() * sizeof(int32_t), "output from diagnostic host staging")) { + Cleanup(comm, stream, input, output, debug, rank, deviceId); + return 1; + } - if (!DemoBarrierAll(rank, rankSize, "all ranks completed memory kernels")) { - Cleanup(comm, stream, input, output, debug, rank, deviceId); - return 1; + if (!DemoBarrierAll(rank, rankSize, "all ranks completed diagnostic host staging")) { + Cleanup(comm, stream, input, output, debug, rank, deviceId); + return 1; + } + } else if (useHostPeerCopy) { + PrintStatus(rank, "diagnostic host peer-memory copy push begin"); + if (!PushInputToPeerWindowsOnHost(rank, rankSize, *commArgsHost, input, elementsPerRank)) { + Cleanup(comm, stream, input, output, debug, rank, deviceId); + return 1; + } + + if (!DemoBarrierAll(rank, rankSize, "all ranks completed diagnostic host peer-memory copy push")) { + Cleanup(comm, stream, input, output, debug, rank, deviceId); + return 1; + } + + PrintStatus(rank, "diagnostic host peer-memory copy collect begin"); + if (!CollectLocalWindowOnHost(rank, rankSize, *commArgsHost, output, elementsPerRank)) { + Cleanup(comm, stream, input, output, debug, rank, deviceId); + return 1; + } + + if (!DemoBarrierAll(rank, rankSize, "all ranks completed diagnostic host peer-memory copy collect")) { + Cleanup(comm, stream, input, output, debug, rank, deviceId); + return 1; + } + } else { + uint32_t blockDim = static_cast(std::max(kDefaultBlockDim, rankSize)); + PrintStatus(rank, "launch memory push kernel blockDim=" + std::to_string(blockDim)); + launch_tilexr_memory_push(blockDim, stream, commArgsDev, reinterpret_cast(input), + reinterpret_cast(output), reinterpret_cast(debug), elementsPerRank); + if (!CheckAcl(rank, "aclrtSynchronizeStream", aclrtSynchronizeStream(stream))) { + Cleanup(comm, stream, input, output, debug, rank, deviceId); + return 1; + } + + if (!DemoBarrierAll(rank, rankSize, "all ranks completed memory push kernels")) { + Cleanup(comm, stream, input, output, debug, rank, deviceId); + return 1; + } + + PrintStatus(rank, "launch memory collect kernel blockDim=" + std::to_string(blockDim)); + launch_tilexr_memory_collect(blockDim, stream, commArgsDev, reinterpret_cast(input), + reinterpret_cast(output), reinterpret_cast(debug), elementsPerRank); + if (!CheckAcl(rank, "aclrtSynchronizeStream", aclrtSynchronizeStream(stream))) { + Cleanup(comm, stream, input, output, debug, rank, deviceId); + return 1; + } + + if (!DemoBarrierAll(rank, rankSize, "all ranks completed memory collect kernels")) { + Cleanup(comm, stream, input, output, debug, rank, deviceId); + return 1; + } } if (!CopyDeviceToHost(rank, hostOutput.data(), hostOutput.size() * sizeof(int32_t), output, diff --git a/tests/memory/demo/tilexr_memory_demo_kernel.cpp b/tests/memory/demo/tilexr_memory_demo_kernel.cpp index e7188515..b0ccabb4 100644 --- a/tests/memory/demo/tilexr_memory_demo_kernel.cpp +++ b/tests/memory/demo/tilexr_memory_demo_kernel.cpp @@ -5,11 +5,11 @@ #include "comm_args.h" #include "kernel_operator.h" -#include "tilexr_sync.h" namespace { constexpr int32_t TILEXR_MEMORY_DEMO_MAGIC = 0x544d454d; // "TMEM" -constexpr int32_t TILEXR_MEMORY_DEMO_STEP_READY = 1; +constexpr int32_t TILEXR_MEMORY_DEMO_STEP_PUSH = 1; +constexpr int32_t TILEXR_MEMORY_DEMO_STEP_COLLECT = 2; constexpr uint32_t TILEXR_MEMORY_DEMO_UB_BYTES = 64 * 1024; constexpr uint32_t TILEXR_MEMORY_DEMO_SYNC_UB_BYTES = 4 * 1024; @@ -64,100 +64,145 @@ __aicore__ inline void CopyGmToGm( if (tile > kTileElements) { tile = kTileElements; } - AscendC::DataCopy(local, src[copied], static_cast(tile)); + const uint32_t tileBytes = static_cast(tile * static_cast(sizeof(T))); + AscendC::DataCopyParams copyParams {1, static_cast(tileBytes), 0, 0}; + AscendC::DataCopyPadParams padParams {false, 0, 0, 0}; + AscendC::DataCopyPad(local, src[copied], copyParams, padParams); AscendC::SetFlag(EVENT_ID0); AscendC::WaitFlag(EVENT_ID0); - AscendC::DataCopy(dst[copied], local, static_cast(tile)); + AscendC::DataCopyPad(dst[copied], local, copyParams); AscendC::SetFlag(EVENT_ID0); AscendC::WaitFlag(EVENT_ID0); } AscendC::PipeBarrier(); } + +__aicore__ inline bool LoadMemoryDemoArgs( + GM_ADDR commArgsGM, int32_t elementsPerRank, int32_t& rank, int32_t& rankSize, GM_ADDR* shareAddrs) +{ + auto args = reinterpret_cast<__gm__ TileXR::CommArgs*>(commArgsGM); + rank = args->rank; + rankSize = args->rankSize; + if (elementsPerRank <= 0 || rankSize <= 0 || rank < 0 || rank >= rankSize || + rankSize > TileXR::TILEXR_MAX_RANK_SIZE) { + return false; + } + + AscendC::GlobalTensor peerMems; + peerMems.SetGlobalBuffer(&(args->peerMems[0]), TileXR::TILEXR_MAX_RANK_SIZE); + for (int32_t peer = 0; peer < rankSize; ++peer) { + shareAddrs[peer] = peerMems.GetValue(peer); + if (shareAddrs[peer] == nullptr) { + return false; + } + } + return true; +} + +__aicore__ inline void WriteMemoryDemoDebug( + __gm__ int32_t* debug, int32_t rank, int32_t rankSize, int32_t elementsPerRank, int32_t blockNum, + int32_t step) +{ + if (debug == nullptr || AscendC::GetBlockIdx() != 0) { + return; + } + debug[0] = TILEXR_MEMORY_DEMO_MAGIC; + debug[1] = rank; + debug[2] = rankSize; + debug[3] = elementsPerRank; + debug[4] = blockNum; + debug[5] = step; +} } // namespace -extern "C" __global__ __aicore__ void tilexr_memory_all_gather_kernel( +extern "C" __global__ __aicore__ void tilexr_memory_push_kernel( GM_ADDR commArgsGM, GM_ADDR inputGM, GM_ADDR outputGM, GM_ADDR debugGM, int32_t elementsPerRank) { if constexpr (g_coreType == AscendC::AIV) { - auto args = reinterpret_cast<__gm__ TileXR::CommArgs*>(commArgsGM); auto input = reinterpret_cast<__gm__ int32_t*>(inputGM); - auto output = reinterpret_cast<__gm__ int32_t*>(outputGM); auto debug = reinterpret_cast<__gm__ int32_t*>(debugGM); - int32_t rank = args->rank; - int32_t rankSize = args->rankSize; int32_t blockIdx = AscendC::GetBlockIdx(); int32_t blockNum = AscendC::GetBlockNum(); - if (debug != nullptr && blockIdx == 0) { - debug[0] = TILEXR_MEMORY_DEMO_MAGIC; - debug[1] = rank; - debug[2] = rankSize; - debug[3] = elementsPerRank; - debug[4] = blockNum; - } - if (elementsPerRank <= 0 || rankSize <= 0 || rank < 0 || rank >= rankSize) { + int32_t rank = 0; + int32_t rankSize = 0; + GM_ADDR shareAddrs[TileXR::TILEXR_MAX_RANK_SIZE]; + if (!LoadMemoryDemoArgs(commArgsGM, elementsPerRank, rank, rankSize, shareAddrs)) { return; } + WriteMemoryDemoDebug(debug, rank, rankSize, elementsPerRank, blockNum, TILEXR_MEMORY_DEMO_STEP_PUSH); AscendC::TPipe pipe; AscendC::TBuf tBuf; pipe.InitBuffer(tBuf, TILEXR_MEMORY_DEMO_UB_BYTES); - GM_ADDR shareAddrs[TileXR::TILEXR_MAX_RANK_SIZE]; - AscendC::GlobalTensor peerMems; - peerMems.SetGlobalBuffer(&(args->peerMems[0]), TileXR::TILEXR_MAX_RANK_SIZE); - for (int32_t peer = 0; peer < rankSize; ++peer) { - shareAddrs[peer] = peerMems.GetValue(peer); - } - int64_t localOffset = 0; int64_t localCount = 0; GetBlockSlice(elementsPerRank, blockNum, blockIdx, localOffset, localCount); AscendC::GlobalTensor inputTensor; - AscendC::GlobalTensor shareTensor; inputTensor.SetGlobalBuffer(input + localOffset, localCount); - shareTensor.SetGlobalBuffer(reinterpret_cast<__gm__ int32_t*>(shareAddrs[rank] + TileXR::IPC_DATA_OFFSET) + - localOffset, - localCount); - CopyGmToGm(shareTensor, inputTensor, tBuf, localCount); - - SyncCollectives sync; - sync.Init(rank, rankSize, shareAddrs, tBuf); - sync.SetInnerFlag(TILEXR_MEMORY_DEMO_MAGIC, TILEXR_MEMORY_DEMO_STEP_READY); - - int32_t blocksPerRank = blockNum / rankSize; - if (blocksPerRank <= 0) { - blocksPerRank = 1; + for (int32_t dstRank = 0; dstRank < rankSize; ++dstRank) { + AscendC::GlobalTensor shareTensor; + shareTensor.SetGlobalBuffer( + reinterpret_cast<__gm__ int32_t*>(shareAddrs[dstRank] + TileXR::IPC_DATA_OFFSET) + + static_cast(rank) * elementsPerRank + localOffset, + localCount); + CopyGmToGm(shareTensor, inputTensor, tBuf, localCount); } - int32_t activeBlocks = blocksPerRank * rankSize; - if (blockIdx >= activeBlocks) { + } +} + +extern "C" __global__ __aicore__ void tilexr_memory_collect_kernel( + GM_ADDR commArgsGM, GM_ADDR inputGM, GM_ADDR outputGM, GM_ADDR debugGM, int32_t elementsPerRank) +{ + if constexpr (g_coreType == AscendC::AIV) { + auto output = reinterpret_cast<__gm__ int32_t*>(outputGM); + auto debug = reinterpret_cast<__gm__ int32_t*>(debugGM); + + int32_t blockIdx = AscendC::GetBlockIdx(); + int32_t blockNum = AscendC::GetBlockNum(); + int32_t rank = 0; + int32_t rankSize = 0; + GM_ADDR shareAddrs[TileXR::TILEXR_MAX_RANK_SIZE]; + if (!LoadMemoryDemoArgs(commArgsGM, elementsPerRank, rank, rankSize, shareAddrs)) { return; } + WriteMemoryDemoDebug(debug, rank, rankSize, elementsPerRank, blockNum, TILEXR_MEMORY_DEMO_STEP_COLLECT); - int32_t sourceRank = blockIdx / blocksPerRank; - int32_t sourceBlockIdx = blockIdx - sourceRank * blocksPerRank; - int64_t remoteOffset = 0; - int64_t remoteCount = 0; - GetBlockSlice(elementsPerRank, blocksPerRank, sourceBlockIdx, remoteOffset, remoteCount); + AscendC::TPipe pipe; + AscendC::TBuf tBuf; + pipe.InitBuffer(tBuf, TILEXR_MEMORY_DEMO_UB_BYTES); - sync.WaitRankInnerFlag(TILEXR_MEMORY_DEMO_MAGIC, TILEXR_MEMORY_DEMO_STEP_READY, sourceRank); + const int64_t totalElements = static_cast(rankSize) * elementsPerRank; + int64_t localOffset = 0; + int64_t localCount = 0; + GetBlockSlice(totalElements, blockNum, blockIdx, localOffset, localCount); + if (localCount <= 0) { + return; + } - AscendC::GlobalTensor remoteTensor; + AscendC::GlobalTensor shareTensor; AscendC::GlobalTensor outputTensor; - remoteTensor.SetGlobalBuffer(reinterpret_cast<__gm__ int32_t*>(shareAddrs[sourceRank] + TileXR::IPC_DATA_OFFSET) + - remoteOffset, - remoteCount); - outputTensor.SetGlobalBuffer(output + static_cast(sourceRank) * elementsPerRank + remoteOffset, - remoteCount); - CopyGmToGm(outputTensor, remoteTensor, tBuf, remoteCount); + shareTensor.SetGlobalBuffer( + reinterpret_cast<__gm__ int32_t*>(shareAddrs[rank] + TileXR::IPC_DATA_OFFSET) + localOffset, localCount); + outputTensor.SetGlobalBuffer(output + localOffset, localCount); + CopyGmToGm(outputTensor, shareTensor, tBuf, localCount); } } -void launch_tilexr_memory_all_gather( +void launch_tilexr_memory_push( + uint32_t blockDim, void* stream, GM_ADDR commArgs, GM_ADDR input, GM_ADDR output, GM_ADDR debug, + int32_t elementsPerRank) +{ + tilexr_memory_push_kernel<<>>( + commArgs, input, output, debug, elementsPerRank); +} + +void launch_tilexr_memory_collect( uint32_t blockDim, void* stream, GM_ADDR commArgs, GM_ADDR input, GM_ADDR output, GM_ADDR debug, int32_t elementsPerRank) { - tilexr_memory_all_gather_kernel<<>>( + tilexr_memory_collect_kernel<<>>( commArgs, input, output, debug, elementsPerRank); } diff --git a/tests/memory/unit/test_tilexr_memory_demo_sources.cpp b/tests/memory/unit/test_tilexr_memory_demo_sources.cpp index 52cee587..1fbf8364 100644 --- a/tests/memory/unit/test_tilexr_memory_demo_sources.cpp +++ b/tests/memory/unit/test_tilexr_memory_demo_sources.cpp @@ -66,16 +66,31 @@ void CheckAppearsBefore( } } +void CheckBlockNotContains(const std::string& path, const std::string& text, const std::string& start, + const std::string& end, const std::string& needle) +{ + const auto startPos = text.find(start); + const auto endPos = startPos == std::string::npos ? std::string::npos : text.find(end, startPos); + if (startPos == std::string::npos || endPos == std::string::npos) { + std::cerr << path << " missing block: " << start << std::endl; + ++g_failures; + return; + } + CheckNotContains(path, text.substr(startPos, endPos - startPos), needle); +} + void TestMemoryDemoKernelUsesPeerMemorySemantics() { const std::string path = "tests/memory/demo/tilexr_memory_demo_kernel.cpp"; const std::string text = ReadFile(path); CheckAppearsBefore(path, text, "#include \"comm_args.h\"", "#include \"kernel_operator.h\""); - CheckContains(path, text, "tilexr_memory_all_gather_kernel"); + CheckContains(path, text, "tilexr_memory_push_kernel"); + CheckContains(path, text, "tilexr_memory_collect_kernel"); CheckContains(path, text, "peerMems"); - CheckContains(path, text, "SyncCollectives"); - CheckContains(path, text, "AscendC::DataCopy"); + CheckContains(path, text, "AscendC::DataCopyPad"); CheckContains(path, text, "IPC_DATA_OFFSET"); + CheckNotContains(path, text, "WaitRankInnerFlag"); + CheckNotContains(path, text, "SyncCollectives"); CheckNotContains(path, text, "TileXRUDMARegister"); CheckNotContains(path, text, "UDMAPut"); CheckNotContains(path, text, "UDMAGet"); @@ -99,14 +114,38 @@ void TestMemoryDemoHostAndRunnerExist() const std::string hostText = ReadFile(hostPath); CheckContains(hostPath, hostText, "TileXRCommInitRankLocal"); CheckContains(hostPath, hostText, "TileXRGetCommArgsDev"); - CheckContains(hostPath, hostText, "launch_tilexr_memory_all_gather"); + CheckContains(hostPath, hostText, "PushInputToPeerWindowsOnHost"); + CheckContains(hostPath, hostText, "CollectLocalWindowOnHost"); + CheckContains(hostPath, hostText, "ExchangeInputSegmentsOnHost"); + CheckContains(hostPath, hostText, "TILEXR_MEMORY_DEMO_HOST_STAGING\", 0"); + CheckContains(hostPath, hostText, "TILEXR_MEMORY_DEMO_HOST_COPY\", 0"); + CheckContains(hostPath, hostText, "diagnostic host staging"); + CheckNotContains(hostPath, hostText, "host staging memory fallback"); + CheckContains(hostPath, hostText, "std::string host"); + CheckContains(hostPath, hostText, "inet_pton(AF_INET, endpoint.host.c_str()"); + CheckContains(hostPath, hostText, "htonl(INADDR_ANY)"); + CheckNotContains(hostPath, hostText, "htonl(INADDR_LOOPBACK)"); + CheckContains(hostPath, hostText, "ACL_MEMCPY_DEVICE_TO_DEVICE"); + CheckContains(hostPath, hostText, "diagnostic host peer-memory copy"); + CheckContains(hostPath, hostText, "launch memory push kernel"); + CheckContains(hostPath, hostText, "launch_tilexr_memory_push"); + CheckContains(hostPath, hostText, "launch_tilexr_memory_collect"); + CheckContains(hostPath, hostText, "all ranks completed memory push kernels"); + CheckContains(hostPath, hostText, "all ranks completed memory collect kernels"); CheckNotContains(hostPath, hostText, "TileXRUDMARegister"); const std::string runPath = "tests/memory/demo/run_tilexr_memory_demo.sh"; const std::string runText = ReadFile(runPath); CheckContains(runPath, runText, "tilexr_memory_demo"); CheckContains(runPath, runText, "${TILEXR_ROOT}/install/lib64"); - CheckNotContains(runPath, runText, "/usr/local/lib"); +} + +void TestMemoryHostTargetsExcludeCannDevlib() +{ + const std::string path = "tests/memory/CMakeLists.txt"; + const std::string text = ReadFile(path); + CheckBlockNotContains(path, text, "link_directories(", "add_executable(test_tilexr_memory_demo_sources", + "-linux/devlib"); } } // namespace @@ -116,6 +155,7 @@ int main() TestMemoryDemoKernelUsesPeerMemorySemantics(); TestCommArgsOwnsAicoreGmAddrCompatibility(); TestMemoryDemoHostAndRunnerExist(); + TestMemoryHostTargetsExcludeCannDevlib(); if (g_failures != 0) { std::cerr << g_failures << " TileXR memory demo source checks failed" << std::endl; return 1; diff --git a/tests/udma/CMakeLists.txt b/tests/udma/CMakeLists.txt index 653808c9..8ac51851 100644 --- a/tests/udma/CMakeLists.txt +++ b/tests/udma/CMakeLists.txt @@ -62,7 +62,6 @@ include_directories( link_directories( ${ASCEND_DRIVER_PATH}/lib64/driver ${ASCEND_HOME_PATH}/${ARCH}-linux/lib64 - ${ASCEND_HOME_PATH}/${ARCH}-linux/devlib ) add_executable(test_tilexr_udma_registry @@ -73,6 +72,14 @@ target_include_directories(test_tilexr_udma_registry PRIVATE ${TILEXR_ROOT}/src/include ) +add_executable(test_tilexr_transport_auto_route + unit/test_tilexr_transport_auto_route.cpp +) + +target_include_directories(test_tilexr_transport_auto_route PRIVATE + ${TILEXR_ROOT}/src/include +) + add_executable(test_tilexr_udma_transport_layout unit/test_tilexr_udma_transport_layout.cpp ${TILEXR_ROOT}/src/comm/udma/tilexr_udma_layout.cpp @@ -114,6 +121,7 @@ target_link_libraries(test_tilexr_udma set(INSTALL_TARGETS test_tilexr_udma test_tilexr_udma_registry + test_tilexr_transport_auto_route test_tilexr_udma_transport_layout test_tilexr_udma_demo_sources test_tilexr_udma_source_guard diff --git a/tests/udma/build.sh b/tests/udma/build.sh index 02177a27..19c168de 100755 --- a/tests/udma/build.sh +++ b/tests/udma/build.sh @@ -60,6 +60,7 @@ echo "==========================================" echo "Test binaries installed to: ${INSTALL_DIR}/bin" echo "" echo "Available tests:" +echo " - test_tilexr_transport_auto_route : auto transport route selector unit tests" echo " - test_tilexr_udma_transport_layout : UDMA info layout unit tests" echo " - test_tilexr_udma_registry : registered-memory metadata unit tests" echo " - test_tilexr_udma_source_guard : UDMA ownership/source boundary checks" diff --git a/tests/udma/demo/run_tilexr_udma_data_channel_probe_mpi.sh b/tests/udma/demo/run_tilexr_udma_data_channel_probe_mpi.sh index df34a9d3..3f141ebe 100644 --- a/tests/udma/demo/run_tilexr_udma_data_channel_probe_mpi.sh +++ b/tests/udma/demo/run_tilexr_udma_data_channel_probe_mpi.sh @@ -109,7 +109,22 @@ if [[ ! -x "${bin}" ]]; then fi export TILEXR_COMM_ID="${COMM_ID}" -export TILEXR_DEMO_BARRIER_ADDR="${COMM_ID}" +COMM_HOST="${COMM_ID%:*}" +COMM_PORT="${COMM_ID##*:}" +if [[ -z "${COMM_HOST}" || -z "${COMM_PORT}" || "${COMM_HOST}" == "${COMM_PORT}" || + ! "${COMM_PORT}" =~ ^[0-9]+$ ]]; then + echo "invalid --comm-id, expected host:port: ${COMM_ID}" >&2 + exit 2 +fi +BARRIER_PORT=$((COMM_PORT + 97)) +if (( BARRIER_PORT > 65535 )); then + BARRIER_PORT=$((COMM_PORT - 97)) +fi +if (( BARRIER_PORT <= 0 )); then + echo "cannot derive demo barrier port from --comm-id: ${COMM_ID}" >&2 + exit 2 +fi +export TILEXR_DEMO_BARRIER_ADDR="${COMM_HOST}:${BARRIER_PORT}" export TILEXR_DEMO_TEST_TYPE="${TEST_TYPE}" export TILEXR_DEMO_ELEMENTS_PER_RANK="${ELEMENTS}" export TILEXR_DEMO_NPUS="${NPU_COUNT}" diff --git a/tests/udma/demo/tilexr_udma_demo.cpp b/tests/udma/demo/tilexr_udma_demo.cpp index eabe662b..14d0fd49 100644 --- a/tests/udma/demo/tilexr_udma_demo.cpp +++ b/tests/udma/demo/tilexr_udma_demo.cpp @@ -39,9 +39,29 @@ constexpr int kConnectRetryCount = 500; constexpr int kConnectRetrySleepMs = 10; struct BarrierEndpoint { + std::string host = "127.0.0.1"; uint16_t port; }; +bool ParseHostPort(const char* value, BarrierEndpoint* endpoint) +{ + if (value == nullptr || value[0] == '\0' || endpoint == nullptr) { + return false; + } + const std::string text(value); + const size_t colon = text.rfind(':'); + if (colon == std::string::npos || colon == 0 || colon + 1 >= text.size()) { + return false; + } + const int port = std::atoi(text.c_str() + colon + 1); + if (port <= 0 || port > 65535) { + return false; + } + endpoint->host = text.substr(0, colon); + endpoint->port = static_cast(port); + return true; +} + int GetEnvInt(const char* name, int defaultValue) { const char* value = std::getenv(name); @@ -151,12 +171,19 @@ bool CopyDeviceToHost(int rank, void* dst, size_t dstSize, const void* src, size BarrierEndpoint GetBarrierEndpoint() { + BarrierEndpoint endpoint{}; + const char* barrier = std::getenv("TILEXR_DEMO_BARRIER_ADDR"); + if (ParseHostPort(barrier, &endpoint)) { + return endpoint; + } + int basePort = kDefaultCommPort; const char* commId = std::getenv("TILEXR_COMM_ID"); if (commId != nullptr) { std::string value(commId); size_t colon = value.rfind(':'); if (colon != std::string::npos && colon + 1 < value.size()) { + endpoint.host = value.substr(0, colon); basePort = std::atoi(value.c_str() + colon + 1); } } @@ -164,7 +191,8 @@ BarrierEndpoint GetBarrierEndpoint() if (barrierPort <= 0 || barrierPort > 65535) { barrierPort = kDefaultCommPort + kDemoBarrierPortOffset; } - return BarrierEndpoint{static_cast(barrierPort)}; + endpoint.port = static_cast(barrierPort); + return endpoint; } bool SendAll(int fd, const void* data, size_t bytes) @@ -220,7 +248,7 @@ int CreateBarrierServer(uint16_t port) } sockaddr_in addr{}; addr.sin_family = AF_INET; - addr.sin_addr.s_addr = htonl(INADDR_LOOPBACK); + addr.sin_addr.s_addr = htonl(INADDR_ANY); addr.sin_port = htons(port); if (bind(fd, reinterpret_cast(&addr), sizeof(addr)) != 0 || listen(fd, SOMAXCONN) != 0) { @@ -230,12 +258,14 @@ int CreateBarrierServer(uint16_t port) return fd; } -int ConnectBarrierServer(uint16_t port) +int ConnectBarrierServer(const BarrierEndpoint& endpoint) { sockaddr_in addr{}; addr.sin_family = AF_INET; - addr.sin_addr.s_addr = htonl(INADDR_LOOPBACK); - addr.sin_port = htons(port); + addr.sin_port = htons(endpoint.port); + if (inet_pton(AF_INET, endpoint.host.c_str(), &addr.sin_addr) != 1) { + return -1; + } for (int attempt = 0; attempt < kConnectRetryCount; ++attempt) { int fd = socket(AF_INET, SOCK_STREAM, 0); @@ -258,14 +288,15 @@ bool DemoBarrierAll(int rank, int rankSize, const std::string& step) } BarrierEndpoint endpoint = GetBarrierEndpoint(); - PrintStatus(rank, "demo tcp barrier begin: " + step + " port=" + std::to_string(endpoint.port)); + PrintStatus(rank, "demo tcp barrier begin: " + step + " endpoint=" + endpoint.host + ":" + + std::to_string(endpoint.port)); constexpr uint8_t kArrive = 1; constexpr uint8_t kRelease = 2; if (rank == 0) { int serverFd = CreateBarrierServer(endpoint.port); if (serverFd < 0) { - std::cerr << "[rank " << rank << "] ERROR: failed to create demo barrier server on 127.0.0.1:" + std::cerr << "[rank " << rank << "] ERROR: failed to create demo barrier server on 0.0.0.0:" << endpoint.port << ", errno=" << errno << std::endl; return false; } @@ -296,10 +327,10 @@ bool DemoBarrierAll(int rank, int rankSize, const std::string& step) return false; } } else { - int fd = ConnectBarrierServer(endpoint.port); + int fd = ConnectBarrierServer(endpoint); if (fd < 0) { - std::cerr << "[rank " << rank << "] ERROR: failed to connect demo barrier on 127.0.0.1:" - << endpoint.port << std::endl; + std::cerr << "[rank " << rank << "] ERROR: failed to connect demo barrier on " << endpoint.host + << ":" << endpoint.port << std::endl; return false; } uint8_t release = 0; @@ -386,8 +417,9 @@ int main(int argc, char** argv) int argIndex = 1; int rankSize = argc > argIndex ? std::atoi(argv[argIndex++]) : GetRankSizeFromEnv(); int rank = argc > argIndex ? std::atoi(argv[argIndex++]) : GetRankFromEnv(); - int testType = argc > argIndex ? std::atoi(argv[argIndex++]) : 0; - int32_t elementsPerRank = argc > argIndex ? std::atoi(argv[argIndex++]) : kDefaultElementsPerRank; + int testType = argc > argIndex ? std::atoi(argv[argIndex++]) : GetEnvInt("TILEXR_DEMO_TEST_TYPE", 0); + int32_t elementsPerRank = argc > argIndex ? std::atoi(argv[argIndex++]) : + GetEnvInt("TILEXR_DEMO_ELEMENTS_PER_RANK", kDefaultElementsPerRank); int npuCount = argc > argIndex ? std::atoi(argv[argIndex++]) : GetEnvInt("TILEXR_DEMO_NPUS", 8); int firstNpu = argc > argIndex ? std::atoi(argv[argIndex++]) : GetEnvInt("TILEXR_DEMO_FIRST_NPU", 0); int deviceId = GetDeviceIdFromEnv(rank, npuCount, firstNpu); diff --git a/tests/udma/demo/tilexr_udma_demo_kernel.cpp b/tests/udma/demo/tilexr_udma_demo_kernel.cpp index 6f42fab7..d72ac25b 100644 --- a/tests/udma/demo/tilexr_udma_demo_kernel.cpp +++ b/tests/udma/demo/tilexr_udma_demo_kernel.cpp @@ -17,7 +17,7 @@ extern "C" __global__ __aicore__ void tilexr_udma_all_gather_kernel( int32_t rank = args->rank; int32_t rankSize = args->rankSize; - bool enabled = TileXR::UDMARegistryEnabled(args); + bool enabled = TileXR::UDMAAllPeersEnabled(args); if (debug != nullptr) { debug[0] = TILEXR_UDMA_DEMO_MAGIC; @@ -53,7 +53,7 @@ extern "C" __global__ __aicore__ void tilexr_udma_put_signal_kernel( int32_t rank = args->rank; int32_t rankSize = args->rankSize; - bool enabled = TileXR::UDMARegistryEnabled(args); + bool enabled = TileXR::UDMAAllPeersEnabled(args); if (debug != nullptr) { debug[0] = TILEXR_UDMA_DEMO_MAGIC; @@ -91,7 +91,7 @@ extern "C" __global__ __aicore__ void tilexr_udma_slot_signal_get_probe_kernel( int32_t rank = args->rank; int32_t rankSize = args->rankSize; - bool enabled = TileXR::UDMARegistryEnabled(args); + bool enabled = TileXR::UDMAAllPeersEnabled(args); if (debug != nullptr) { debug[0] = TILEXR_UDMA_DEMO_MAGIC; @@ -136,7 +136,7 @@ extern "C" __global__ __aicore__ void tilexr_udma_registered_smoke_kernel( auto local = reinterpret_cast<__gm__ uint8_t*>(localGM); auto debug = reinterpret_cast<__gm__ int32_t*>(debugGM); - bool enabled = TileXR::UDMARegistryEnabled(args); + bool enabled = TileXR::UDMAAllPeersEnabled(args); if (debug != nullptr) { debug[0] = TILEXR_UDMA_DEMO_MAGIC; debug[1] = enabled ? 1 : 0; diff --git a/tests/udma/run_tests.sh b/tests/udma/run_tests.sh index e899d2de..251e1c16 100755 --- a/tests/udma/run_tests.sh +++ b/tests/udma/run_tests.sh @@ -13,7 +13,7 @@ INSTALL_DIR="${SCRIPT_DIR}/install" source "${TILEXR_ROOT}/scripts/common_env.sh" # 设置 LD_LIBRARY_PATH:优先使用当前仓库刚编译安装的库,避免被 /usr/local/lib 中的旧库覆盖 -export LD_LIBRARY_PATH="${INSTALL_DIR}/lib:${TILEXR_ROOT}/install/lib:/usr/local/lib:${LD_LIBRARY_PATH}" +export LD_LIBRARY_PATH="${INSTALL_DIR}/lib64:${INSTALL_DIR}/lib:${TILEXR_ROOT}/install/lib64:${TILEXR_ROOT}/install/lib:/usr/local/lib:${LD_LIBRARY_PATH}" echo "==========================================" echo " Running UDMA Tests" @@ -45,6 +45,7 @@ fi # 检查测试二进制是否存在 if [ ! -f "${INSTALL_DIR}/bin/test_tilexr_udma_transport_layout" ] || + [ ! -f "${INSTALL_DIR}/bin/test_tilexr_transport_auto_route" ] || [ ! -f "${INSTALL_DIR}/bin/test_tilexr_udma_registry" ] || [ ! -f "${INSTALL_DIR}/bin/test_tilexr_udma_source_guard" ] || [ ! -f "${INSTALL_DIR}/bin/test_tilexr_udma" ]; then @@ -54,41 +55,49 @@ fi # 测试 1: UDMA info layout 单元测试(host-only) echo "==========================================" -echo "Test 1: TileXR UDMA Transport Layout Unit Test" +echo "Test 1: TileXR Auto Route Selector Unit Test" echo "==========================================" -"${INSTALL_DIR}/bin/test_tilexr_udma_transport_layout" +"${INSTALL_DIR}/bin/test_tilexr_transport_auto_route" TEST1_RESULT=$? echo "" -# 测试 2: TileXR UDMA registry 单元测试(host-only) +# 测试 2: UDMA info layout 单元测试(host-only) echo "==========================================" -echo "Test 2: TileXR UDMA Registry Unit Test" +echo "Test 2: TileXR UDMA Transport Layout Unit Test" echo "==========================================" -"${INSTALL_DIR}/bin/test_tilexr_udma_registry" +"${INSTALL_DIR}/bin/test_tilexr_udma_transport_layout" TEST2_RESULT=$? echo "" -# 测试 3: UDMA source guard(host-only) +# 测试 3: TileXR UDMA registry 单元测试(host-only) echo "==========================================" -echo "Test 3: TileXR UDMA Source Guard Unit Test" +echo "Test 3: TileXR UDMA Registry Unit Test" echo "==========================================" -"${INSTALL_DIR}/bin/test_tilexr_udma_source_guard" +"${INSTALL_DIR}/bin/test_tilexr_udma_registry" TEST3_RESULT=$? echo "" -# 测试 4: TileXR 集成测试(单进程,单卡) +# 测试 4: UDMA source guard(host-only) +echo "==========================================" +echo "Test 4: TileXR UDMA Source Guard Unit Test" +echo "==========================================" +"${INSTALL_DIR}/bin/test_tilexr_udma_source_guard" +TEST4_RESULT=$? +echo "" + +# 测试 5: TileXR 集成测试(单进程,单卡) echo "==========================================" -echo "Test 4: TileXR Integration Tests (Single Process)" +echo "Test 5: TileXR Integration Tests (Single Process)" echo "==========================================" export RANK=0 export RANK_SIZE=1 "${INSTALL_DIR}/bin/test_tilexr_udma" -TEST4_RESULT=$? +TEST5_RESULT=$? echo "" -# 测试 5: TileXR 多进程测试(需要 mpirun) +# 测试 6: TileXR 多进程测试(需要 mpirun) echo "==========================================" -echo "Test 5: TileXR Multi-Process Tests (MPI)" +echo "Test 6: TileXR Multi-Process Tests (MPI)" echo "==========================================" # 检查是否有 mpirun @@ -107,14 +116,14 @@ if command -v mpirun &> /dev/null; then unset RANK unset RANK_SIZE mpirun -n 2 "${INSTALL_DIR}/bin/test_tilexr_udma" - TEST5_RESULT=$? + TEST6_RESULT=$? else echo "SKIP: Need at least 2 usable NPUs for multi-rank test" - TEST5_RESULT=0 + TEST6_RESULT=0 fi else echo "SKIP: mpirun not found, skipping multi-process tests" - TEST5_RESULT=0 + TEST6_RESULT=0 fi echo "" @@ -122,16 +131,17 @@ echo "" echo "==========================================" echo " Test Results Summary" echo "==========================================" -echo "Test 1 (UDMA Layout): $([ $TEST1_RESULT -eq 0 ] && echo 'PASS' || echo 'FAIL')" -echo "Test 2 (UDMA Registry): $([ $TEST2_RESULT -eq 0 ] && echo 'PASS' || echo 'FAIL')" -echo "Test 3 (Source Guard): $([ $TEST3_RESULT -eq 0 ] && echo 'PASS' || echo 'FAIL')" -echo "Test 4 (TileXR Single): $([ $TEST4_RESULT -eq 0 ] && echo 'PASS' || echo 'FAIL')" -echo "Test 5 (TileXR Multi): $([ $TEST5_RESULT -eq 0 ] && echo 'PASS' || echo 'SKIP/FAIL')" +echo "Test 1 (Auto Route): $([ $TEST1_RESULT -eq 0 ] && echo 'PASS' || echo 'FAIL')" +echo "Test 2 (UDMA Layout): $([ $TEST2_RESULT -eq 0 ] && echo 'PASS' || echo 'FAIL')" +echo "Test 3 (UDMA Registry): $([ $TEST3_RESULT -eq 0 ] && echo 'PASS' || echo 'FAIL')" +echo "Test 4 (Source Guard): $([ $TEST4_RESULT -eq 0 ] && echo 'PASS' || echo 'FAIL')" +echo "Test 5 (TileXR Single): $([ $TEST5_RESULT -eq 0 ] && echo 'PASS' || echo 'FAIL')" +echo "Test 6 (TileXR Multi): $([ $TEST6_RESULT -eq 0 ] && echo 'PASS' || echo 'SKIP/FAIL')" echo "==========================================" # 返回失败状态 if [ $TEST1_RESULT -ne 0 ] || [ $TEST2_RESULT -ne 0 ] || [ $TEST3_RESULT -ne 0 ] || - [ $TEST4_RESULT -ne 0 ] || [ $TEST5_RESULT -ne 0 ]; then + [ $TEST4_RESULT -ne 0 ] || [ $TEST5_RESULT -ne 0 ] || [ $TEST6_RESULT -ne 0 ]; then exit 1 fi diff --git a/tests/udma/unit/test_tilexr_transport_auto_route.cpp b/tests/udma/unit/test_tilexr_transport_auto_route.cpp new file mode 100644 index 00000000..b3d5f9ca --- /dev/null +++ b/tests/udma/unit/test_tilexr_transport_auto_route.cpp @@ -0,0 +1,182 @@ +#include +#include + +#include "tilexr_transport.h" + +namespace { + +int g_failures = 0; + +#define CHECK_EQ(lhs, rhs) \ + do { \ + auto lhsValue = (lhs); \ + auto rhsValue = (rhs); \ + if (lhsValue != rhsValue) { \ + std::cerr << "CHECK_EQ failed at line " << __LINE__ << ": " #lhs " != " #rhs \ + << " (" << static_cast(lhsValue) << " vs " << static_cast(rhsValue) << ")" \ + << std::endl; \ + ++g_failures; \ + } \ + } while (0) + +TileXR::CommArgs MakeArgs(bool directUrmaAvailable) +{ + TileXR::CommArgs args = {}; + if (directUrmaAvailable) { + args.extraFlag |= TileXR::ExtraFlag::UDMA; + args.udmaInfoPtr = reinterpret_cast(0x100000); + args.udmaRegistryPtr = reinterpret_cast(0x200000); + } + return args; +} + +TileXR::CommArgs MakeArgsMissingInfo() +{ + TileXR::CommArgs args = {}; + args.extraFlag |= TileXR::ExtraFlag::UDMA; + args.udmaRegistryPtr = reinterpret_cast(0x200000); + return args; +} + +TileXR::CommArgs MakeArgsMissingRegistry() +{ + TileXR::CommArgs args = {}; + args.extraFlag |= TileXR::ExtraFlag::UDMA; + args.udmaInfoPtr = reinterpret_cast(0x100000); + return args; +} + +TileXR::CommArgs MakeCrossNodeArgs(bool directUrmaAvailable) +{ + TileXR::CommArgs args = MakeArgs(directUrmaAvailable); + args.rankSize = 2; + args.localRankSize = 1; + return args; +} + +void TestAutoUsesMemoryForNullArgs() +{ + CHECK_EQ(TileXR::TileXRSelectAutoTransport(nullptr, 64ULL * 1024ULL * 1024ULL), + TileXR::TileXRTransportKind::MEMORY); +} + +void TestAutoUsesMemoryForZeroBytes() +{ + auto args = MakeArgs(true); + CHECK_EQ(TileXR::TileXRSelectAutoTransport(&args, 0), + TileXR::TileXRTransportKind::MEMORY); +} + +void TestAutoUsesMemoryBelowFourMiB() +{ + auto args = MakeArgs(true); + CHECK_EQ(TileXR::TileXRSelectAutoTransport(&args, TileXR::TILEXR_AUTO_DIRECT_URMA_THRESHOLD_BYTES - 1), + TileXR::TileXRTransportKind::MEMORY); +} + +void TestAutoUsesDirectUrmaAtFourMiB() +{ + auto args = MakeArgs(true); + CHECK_EQ(TileXR::TileXRSelectAutoTransport(&args, TileXR::TILEXR_AUTO_DIRECT_URMA_THRESHOLD_BYTES), + TileXR::TileXRTransportKind::DIRECT_URMA); +} + +void TestAutoUsesDirectUrmaAboveFourMiB() +{ + auto args = MakeArgs(true); + CHECK_EQ(TileXR::TileXRSelectAutoTransport(&args, 8ULL * 1024ULL * 1024ULL), + TileXR::TileXRTransportKind::DIRECT_URMA); +} + +void TestCrossNodeUsesMemoryBelow128KiB() +{ + auto args = MakeCrossNodeArgs(true); + CHECK_EQ(TileXR::TileXRSelectAutoTransport( + &args, TileXR::TILEXR_AUTO_CROSS_NODE_DIRECT_URMA_THRESHOLD_BYTES - 1), + TileXR::TileXRTransportKind::MEMORY); +} + +void TestCrossNodeUsesDirectUrmaAt128KiB() +{ + auto args = MakeCrossNodeArgs(true); + CHECK_EQ(TileXR::TileXRSelectAutoTransport( + &args, TileXR::TILEXR_AUTO_CROSS_NODE_DIRECT_URMA_THRESHOLD_BYTES), + TileXR::TileXRTransportKind::DIRECT_URMA); +} + +void TestAutoKeepsZeroBytesOnMemoryForCrossNode() +{ + auto args = MakeCrossNodeArgs(true); + CHECK_EQ(TileXR::TileXRSelectAutoTransport(&args, 0), + TileXR::TileXRTransportKind::MEMORY); +} + +void TestAutoFallsBackToMemoryWhenUrmaUnavailable() +{ + auto args = MakeArgs(false); + CHECK_EQ(TileXR::TileXRSelectAutoTransport(&args, 64ULL * 1024ULL * 1024ULL), + TileXR::TileXRTransportKind::MEMORY); +} + +void TestAutoFallsBackToMemoryWhenUrmaInfoMissing() +{ + auto args = MakeArgsMissingInfo(); + CHECK_EQ(TileXR::TileXRSelectAutoTransport(&args, 64ULL * 1024ULL * 1024ULL), + TileXR::TileXRTransportKind::MEMORY); +} + +void TestAutoFallsBackToMemoryWhenUrmaRegistryMissing() +{ + auto args = MakeArgsMissingRegistry(); + CHECK_EQ(TileXR::TileXRSelectAutoTransport(&args, 64ULL * 1024ULL * 1024ULL), + TileXR::TileXRTransportKind::MEMORY); +} + +void TestDirectUrmaPeerRoutabilityMatchesAllocatedResources() +{ + auto args = MakeArgs(true); + args.rank = 0; + args.rankSize = 4; + args.localRankSize = 2; + CHECK_EQ(TileXR::TileXRDirectUrmaPeerRoutable(&args, 0), false); + CHECK_EQ(TileXR::TileXRDirectUrmaPeerRoutable(&args, 1), false); + CHECK_EQ(TileXR::TileXRDirectUrmaPeerRoutable(&args, 2), true); + CHECK_EQ(TileXR::TileXRDirectUrmaPeerRoutable(&args, 3), true); + + args.rankSize = 2; + args.localRankSize = 2; + CHECK_EQ(TileXR::TileXRDirectUrmaPeerRoutable(&args, 1), true); + + args.localRankSize = 0; + CHECK_EQ(TileXR::TileXRDirectUrmaPeerRoutable(&args, 1), false); + + args = MakeArgs(false); + args.rank = 0; + args.rankSize = 2; + args.localRankSize = 1; + CHECK_EQ(TileXR::TileXRDirectUrmaPeerRoutable(&args, 1), false); +} + +} // namespace + +int main() +{ + TestAutoUsesMemoryForNullArgs(); + TestAutoUsesMemoryForZeroBytes(); + TestAutoUsesMemoryBelowFourMiB(); + TestAutoUsesDirectUrmaAtFourMiB(); + TestAutoUsesDirectUrmaAboveFourMiB(); + TestCrossNodeUsesMemoryBelow128KiB(); + TestCrossNodeUsesDirectUrmaAt128KiB(); + TestAutoKeepsZeroBytesOnMemoryForCrossNode(); + TestAutoFallsBackToMemoryWhenUrmaUnavailable(); + TestAutoFallsBackToMemoryWhenUrmaInfoMissing(); + TestAutoFallsBackToMemoryWhenUrmaRegistryMissing(); + TestDirectUrmaPeerRoutabilityMatchesAllocatedResources(); + if (g_failures != 0) { + std::cerr << g_failures << " TileXR transport auto route checks failed" << std::endl; + return 1; + } + std::cout << "TileXR transport auto route checks passed" << std::endl; + return 0; +} diff --git a/tests/udma/unit/test_tilexr_udma_demo_sources.cpp b/tests/udma/unit/test_tilexr_udma_demo_sources.cpp index c6ed7f0c..f5fe5133 100644 --- a/tests/udma/unit/test_tilexr_udma_demo_sources.cpp +++ b/tests/udma/unit/test_tilexr_udma_demo_sources.cpp @@ -60,12 +60,23 @@ int main() CheckContains(demoPath, demo, "launch_tilexr_udma_all_gather"); CheckContains(demoPath, demo, "launch_tilexr_udma_put_signal"); CheckContains(demoPath, demo, "DemoBarrierAll"); + CheckContains(demoPath, demo, "TILEXR_DEMO_BARRIER_ADDR"); + CheckContains(demoPath, demo, "GetEnvInt(\"TILEXR_DEMO_TEST_TYPE\""); + CheckContains(demoPath, demo, "GetEnvInt(\"TILEXR_DEMO_ELEMENTS_PER_RANK\""); + CheckContains(demoPath, demo, "INADDR_ANY"); + CheckContains(demoPath, demo, "inet_pton"); CheckNotContains(demoPath, demo, "aclshmem"); CheckNotContains(demoPath, demo, "shmem_"); + const std::string mpiRunnerPath = "tests/udma/demo/run_tilexr_udma_data_channel_probe_mpi.sh"; + const std::string mpiRunner = ReadFile(mpiRunnerPath); + CheckContains(mpiRunnerPath, mpiRunner, "BARRIER_PORT=$((COMM_PORT + 97))"); + CheckContains(mpiRunnerPath, mpiRunner, "TILEXR_DEMO_BARRIER_ADDR=\"${COMM_HOST}:${BARRIER_PORT}\""); + CheckNotContains(mpiRunnerPath, mpiRunner, "TILEXR_DEMO_BARRIER_ADDR=\"${COMM_ID}\""); + const std::string kernelPath = "tests/udma/demo/tilexr_udma_demo_kernel.cpp"; const std::string kernel = ReadFile(kernelPath); - CheckContains(kernelPath, kernel, "UDMARegistryEnabled"); + CheckContains(kernelPath, kernel, "UDMAAllPeersEnabled"); CheckContains(kernelPath, kernel, "UDMAPutNbi"); CheckContains(kernelPath, kernel, "UDMAPutSignalNbi"); CheckContains(kernelPath, kernel, "UDMAQuiet"); diff --git a/tests/udma/unit/test_tilexr_udma_source_guard.cpp b/tests/udma/unit/test_tilexr_udma_source_guard.cpp index fb9f768b..59996e37 100644 --- a/tests/udma/unit/test_tilexr_udma_source_guard.cpp +++ b/tests/udma/unit/test_tilexr_udma_source_guard.cpp @@ -56,6 +56,19 @@ void CheckNoNeedles(const std::string& path, const std::vector& nee } } +void CheckBlockNotContains(const std::string& path, const std::string& text, const std::string& start, + const std::string& end, const std::string& needle) +{ + const auto startPos = text.find(start); + const auto endPos = startPos == std::string::npos ? std::string::npos : text.find(end, startPos); + if (startPos == std::string::npos || endPos == std::string::npos) { + std::cerr << path << " missing block: " << start << std::endl; + ++g_failures; + return; + } + CheckNotContains(path, text.substr(startPos, endPos - startPos), needle); +} + void TestTileXRCommUsesUDMAContextBoundary() { const std::string headerPath = "src/comm/tilexr_comm.h"; @@ -179,6 +192,92 @@ void TestCommSourcesDoNotUseShmem() } } +void TestEpKernelsUseExplicitUdmaPrimitives() +{ + const std::vector paths = { + "src/ep/kernels/tilexr_ep_dispatch_kernel.cpp", + "src/ep/kernels/tilexr_ep_kernel_common.h", + }; + for (const auto& path : paths) { + const auto text = ReadFile(path); + CheckContains(path, text, "TileXR::UDMAPutNbi"); + CheckNotContains(path, text, "TileXRPutAutoNbi"); + CheckNotContains(path, text, "TileXRGetAutoNbi"); + } + + const std::string transportPath = "src/include/tilexr_transport.h"; + const auto transportText = ReadFile(transportPath); + CheckNotContains(transportPath, transportText, "TileXRPutAutoNbi"); + CheckNotContains(transportPath, transportText, "TileXRGetAutoNbi"); +} + +void TestEpHostLaunchUsesResolvedTransportGate() +{ + const std::string path = "src/ep/host/ep_kernel_launch.cpp"; + const auto text = ReadFile(path); + CheckContains(path, text, "context.transport =="); + CheckContains(path, text, "TileXR::TileXRTransportKind::DIRECT_URMA"); + CheckContains(path, text, "GM_ADDR memoryWorkspace = nullptr;"); + CheckNotContains(path, text, "TileXRSelectAutoTransport"); +} + +void TestUdmaHostTargetsExcludeCannDevlib() +{ + const std::string path = "tests/udma/CMakeLists.txt"; + const auto text = ReadFile(path); + CheckBlockNotContains(path, text, "link_directories(", "add_executable(test_tilexr_udma_registry", + "-linux/devlib"); +} + +void TestSocketExchangeHandlesPartialIo() +{ + const std::string path = "src/comm/tools/socket/tilexr_sock_exchange.h"; + const auto text = ReadFile(path); + CheckContains(path, text, "while (sentBytes < sendSize)"); + CheckContains(path, text, "sendBytes + sentBytes"); + CheckContains(path, text, "sendSize - sentBytes"); + CheckContains(path, text, "while (receivedBytes < recvSize)"); + CheckContains(path, text, "recvBytes + receivedBytes"); + CheckContains(path, text, "recvSize - receivedBytes"); +} + +void TestHybridUdmaOnlyBuildsRemotePeerResources() +{ + const std::string contextHeaderPath = "src/comm/udma/tilexr_udma_context.h"; + const auto contextHeader = ReadFile(contextHeaderPath); + CheckContains(contextHeaderPath, contextHeader, "int localRankSize = 1;"); + + const std::string transportHeaderPath = "src/comm/udma/tilexr_udma_transport.h"; + const auto transportHeader = ReadFile(transportHeaderPath); + CheckContains(transportHeaderPath, transportHeader, "bool UsesUDMAPeer(int peer) const;"); + + const std::string contextPath = "src/comm/udma/tilexr_udma_context.cpp"; + const auto context = ReadFile(contextPath); + CheckContains(contextPath, context, "transportOptions.localRankSize = options_.localRankSize;"); + + const std::string commPath = "src/comm/tilexr_comm.cpp"; + const auto comm = ReadFile(commPath); + CheckContains(commPath, comm, "options.localRankSize = static_cast(localRankSize_);"); + + const std::string transportPath = "src/comm/udma/tilexr_udma_transport.cpp"; + const auto transport = ReadFile(transportPath); + CheckContains(transportPath, transport, "if (!UsesUDMAPeer(peer))"); + CheckContains(transportPath, transport, + "peer / options_.localRankSize != options_.rank / options_.localRankSize"); +} + +void TestPublicUdmaWrappersRejectUnroutablePeers() +{ + const std::string path = "src/include/tilexr_udma.h"; + const auto text = ReadFile(path); + CheckContains(path, text, "UDMAPeerEnabled(args, targetRank)"); + CheckContains(path, text, "UDMAPeerEnabled(args, sourceRank)"); + + const std::string demoPath = "tests/udma/demo/tilexr_udma_demo_kernel.cpp"; + const auto demo = ReadFile(demoPath); + CheckContains(demoPath, demo, "bool enabled = TileXR::UDMAAllPeersEnabled(args);"); +} + } // namespace int main() @@ -189,6 +288,12 @@ int main() TestUDMAReviewFeedbackGuards(); TestPublicHeadersDoNotExposeUDMAContext(); TestCommSourcesDoNotUseShmem(); + TestEpKernelsUseExplicitUdmaPrimitives(); + TestEpHostLaunchUsesResolvedTransportGate(); + TestUdmaHostTargetsExcludeCannDevlib(); + TestSocketExchangeHandlesPartialIo(); + TestHybridUdmaOnlyBuildsRemotePeerResources(); + TestPublicUdmaWrappersRejectUnroutablePeers(); if (g_failures != 0) { std::cerr << g_failures << " UDMA source guard checks failed" << std::endl; return 1;