Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
67 commits
Select commit Hold shift + click to select a range
b7ffb3a
feat(ws2): add TP-aware logprob contract and dispatch metadata
ryankert01 Aug 2, 2026
cdc11ba
fix(ws2): address CodeRabbit review on logprob contract PR
ryankert01 Aug 2, 2026
6455715
fix(ws2): address second CodeRabbit round on logprob contract
ryankert01 Aug 2, 2026
3b4eaef
feat(ws2): make determinism scope and invocation surface part of the …
ryankert01 Aug 2, 2026
e6dbeef
docs(ws2): drop standalone design doc per review
ryankert01 Aug 2, 2026
878ba88
style(ws2): align comment density with sibling kernel modules
ryankert01 Aug 2, 2026
934bc5b
feat: add single-gpu logprob comparison harness
hihaluemen Aug 4, 2026
b69426d
fix: keep logprob CLI stdout machine readable
hihaluemen Aug 4, 2026
0efcfe1
refactor: simplify logprob comparison harness
hihaluemen Aug 4, 2026
115d86c
docs: document SM90 logprob validation
hihaluemen Aug 4, 2026
c028b5b
fix: address logprob harness lint and provenance
hihaluemen Aug 4, 2026
7ba09b5
fix: type heterogeneous logprob backends
hihaluemen Aug 4, 2026
a63bea2
init vocab parallel logp
KJLdefeated Aug 5, 2026
4fcdc30
init vocab parallel logp
KJLdefeated Aug 5, 2026
6ffade4
adding cross tp testing
KJLdefeated Aug 5, 2026
4231625
fix comment
KJLdefeated Aug 6, 2026
2360a71
test(ws2): align dispatch tests with the auto+TP>1 unsafe-dispatch guard
KJLdefeated Aug 7, 2026
99e59f8
Merge branch 'RL-Align:main' into feat/ws2-logprob-single-gpu-harness…
hihaluemen Aug 8, 2026
4eebb3b
refactor: colocate logprob harness tooling and docs
hihaluemen Aug 8, 2026
19488cc
test: resolve logprob CLI path reliably
hihaluemen Aug 8, 2026
6e2a79e
Merge PR1 logprob contract into PR2 integration base
hihaluemen Aug 8, 2026
b7d9d89
init vocab parallel logp
KJLdefeated Aug 5, 2026
3866d3c
init vocab parallel logp
KJLdefeated Aug 5, 2026
65f3c6f
adding cross tp testing
KJLdefeated Aug 5, 2026
05d19eb
test: align PR3 dispatch with latest PR1 guard
hihaluemen Aug 8, 2026
a46891a
Merge branch 'RL-Align:main' into feat/ws2-logprob-single-gpu-harness…
hihaluemen Aug 10, 2026
57b04b1
Merge latest PR2 into PR1-PR3 integration base
hihaluemen Aug 11, 2026
f36a63d
feat(ws2): add distributed logprob drift runner
hihaluemen Aug 11, 2026
1d9bac1
ci(ws2): run logprob comparison tests
hihaluemen Aug 11, 2026
f6b5a07
fix(ws2): harden distributed drift reporting
hihaluemen Aug 11, 2026
8a4f4be
fix(ws2): clean up process groups on setup failure
hihaluemen Aug 11, 2026
3627fcd
ws2 deterministic grpo loss PR5
KJLdefeated Aug 12, 2026
e668339
Merge branch 'main' into feat/ws2-logprob-tp-contract-pr1
KJLdefeated Aug 12, 2026
d407561
Merge issue 241 PR1
hihaluemen Aug 19, 2026
e3be4f9
Merge issue 241 PR2
hihaluemen Aug 19, 2026
628a9f6
Merge issue 241 PR3
hihaluemen Aug 19, 2026
28dde33
Merge issue 241 PR4
hihaluemen Aug 19, 2026
b52fa19
Merge issue 241 PR5
hihaluemen Aug 19, 2026
c2c99c1
feat(ws2): record cross-topology logprob fingerprints
hihaluemen Aug 19, 2026
9aca059
Merge PR4 cross-topology fingerprints
hihaluemen Aug 19, 2026
8428676
feat(ws2): add Vime logprob provider for CP metadata
inaniloquentee Aug 21, 2026
189c222
fix(ws2): keep TP entropy merge order explicit
inaniloquentee Aug 21, 2026
5f8d656
feat(vime): add Qwen3 TP2 CP2 validation example
inaniloquentee Aug 21, 2026
3836f66
Merge branch 'RL-Align:main' into feat/ws2-logprob-distributed-report…
hihaluemen Aug 21, 2026
b811abc
Merge branch 'feat/ws2-logprob-distributed-report-pr4' into work/ws2-…
hihaluemen Aug 21, 2026
f79602c
feat(rocm): add deterministic vocab parallel logprob backend
hihaluemen Aug 21, 2026
1fcacbc
Merge upstream test into ROCm logprob PR
hihaluemen Aug 21, 2026
b776227
style: apply isort ordering to setup imports
hihaluemen Aug 21, 2026
9a05c1e
test: isolate ROCm vocab logprob dispatch cases
hihaluemen Aug 22, 2026
1c73a42
style: format ROCm dispatch test with black
hihaluemen Aug 22, 2026
8e793c9
logp benchmark
KJLdefeated Aug 23, 2026
e9f1d2a
modify report
KJLdefeated Aug 23, 2026
dd5fe05
feat(logprob): add Triton WS2 vocab-parallel backend, shared fused ke…
KJLdefeated Aug 23, 2026
46ec8d7
Update CPU and Cuda golden baseline
frank-2077 Aug 24, 2026
0831aa8
Merge pull request #361 from RL-Align/test
Flink-ddd Aug 30, 2026
c57a8ae
Optimize deterministic rollout tensor-parallel all-reduce
inaniloquentee Aug 30, 2026
9d5732b
Optimize deterministic rollout tensor-parallel all-reduce (#365)
inaniloquentee Aug 30, 2026
8e32313
Enable exact-batch CUDA graphs for strict rollout
inaniloquentee Aug 30, 2026
9b95e34
Precompile strict FA4 training kernels
inaniloquentee Aug 31, 2026
0aa1d63
Add strict FA4 training precompile
inaniloquentee Aug 31, 2026
afccecc
Optimize small deterministic all-reduce path
inaniloquentee Aug 31, 2026
01b4ae4
Merge pull request #367 from RL-Align/codex/deterministic-full-decode…
Flink-ddd Aug 31, 2026
eedb788
Merge upstream test into ROCm logprob PR
hihaluemen Aug 31, 2026
6ce841b
Merge collaborator updates into upstream-synced ROCm logprob PR
hihaluemen Aug 31, 2026
c5431ce
Merge branch 'RL-Align:main' into work/ws2-logprob-rocm
hihaluemen Aug 31, 2026
b4dc9fb
style: format synced runtime updates
hihaluemen Aug 31, 2026
3156686
fix(logprob): restore deterministic wrapper interface
hihaluemen Aug 31, 2026
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 7 additions & 0 deletions .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -216,3 +216,10 @@ _dev_notes/

# Local C8 execute dumps; default output is under TMPDIR.
ws1-c8-ci.json

# hipify outputs generated during ROCm builds (sources live in csrc/*.cu, csrc/cuda/, csrc/hip/hip_*.hip)
csrc/*.hip
csrc/*_hip.cpp
csrc/hip/**
!csrc/hip/hip_*.hip
csrc/hip/*_hip.hip
174 changes: 174 additions & 0 deletions _dev_notes/rocm_logprob_implementation_summary.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,174 @@
# ROCm Logprob 实现说明

日期:2026-08-22

这次工作是在临时分支 `work/ws2-logprob-rocm` 上完成的。它基于 issue #241 的 PR1-PR5 integration,再合入当前 PR4 的最新更新。

## 一句话概括

新增了一个 ROCm 版 WS2 vocab-parallel logprob 后端:每个 GPU 用 HIP/CUDA 可移植 kernel 计算本地 vocab tile 的 FP32 `(max, sumexp)`,TP 之间仍使用 issue #241 已确定的 all-gather + 固定顺序 merge。这样不会为了追求 ROCm 速度而改变原来的数值契约。

## 改了什么

### 1. 增加本地 tile partial kernel

文件:`csrc/deterministic_logp_kernel.cu`

新增 `deterministic_logp_tile_stats`:

- 输入一个 TP rank 的 local vocab logits;
- 每个 vocab tile 输出 FP32 `max` 和 `sumexp`;
- 过滤真实 vocab 之外的 padding 列;
- 使用固定 block reduction,不使用 atomic 或不确定的全局归约;
- 同一份 `.cu` 源码可由 CUDA 或 HIP 编译。

这个 kernel 只负责 rank-local 计算,不负责 RCCL/NCCL,也不负责全局 LSE 合并。跨 TP 的合并顺序仍由已有 WS2 Python 实现控制。

### 2. 增加 ROCm backend 注册

文件:

- `rl_engine/kernels/ops/rocm/loss/vocab_parallel_logp.py`
- `rl_engine/kernels/ops/rocm/loss/__init__.py`
- `rl_engine/kernels/registry.py`

新增 backend:

```text
rocm-vocab-parallel-logp-ws2
```

它只在 ROCm platform 注册,放在 PyTorch reference 前面。backend 继承现有 `VocabParallelLogprobOp`,所以以下逻辑仍然共用一份实现:

- TP contract 和 preflight;
- tile-aligned shard 检查;
- all-gather transport;
- global tile order merge;
- selected-token owner gather;
- entropy 的显式 rank-order merge;
- backward 和 active-mask 语义。

只有在 ROCm native extension 已经加载、且提供新符号时,registry 才会注册这个 production backend。native 不可用时,registry 不会把它显示成可用的 ROCm production 实现,而是明确使用已有的 PyTorch reference backend。这样 backend provenance 是准确的,不会出现“界面显示用了 ROCm,实际却悄悄跑了另一条路径”的情况。

如果调用方显式要求 native ROCm backend,但扩展或符号缺失,会直接抛出清晰的 `RuntimeError`,不会在 wrapper 内部静默 fallback。PyTorch reference 始终使用纯 PyTorch tile 统计;ROCm 子类只有在 native 能力已由 registry 检查通过时才启用 HIP tile kernel。

### 3. 增加 native Python binding 和类型声明

文件:

- `csrc/ops.cpp`
- `rl_engine/_C.pyi`

注册了 `deterministic_logp_tile_stats` 的 Python binding 和类型签名。

同时修复了一个 ROCm 构建隐患:`deterministic_collective_*` 使用 CUDA IPC,但 ROCm 构建不编译对应源文件。现在这些符号的声明和 pybind 注册在 HIP 编译时都会被排除,避免 ROCm 链接阶段出现 unresolved symbol。

Prefix-Shared Attention 的 NVIDIA PTX 注册也改成 HIP 编译时排除,与 PR4 的 source gating 保持一致。

### 4. 增加测试

文件:`tests/test_rocm_logprob_backend.py`

覆盖:

- ROCm backend 继承并保持 WS2 operator surface;
- backend 只在 ROCm registry 中注册;
- native extension 可用时才注册 ROCm production backend;不可用时首选 PyTorch reference;
- native tile kernel 和 binding 存在;
- CPU-only 环境可以导入 ROCm wrapper,不要求 native extension。
- reference/native 两条执行路径不会互相静默切换。

另外补上了和 PR #319/#325 思路一致的验证入口:

- ROCm native 路径复用 TP2/TP4 的 TP1 对照、重复执行、forward/backward 和 bitwise 检查;
- ROCm 多卡 `TP2 x CP2` CLI 测试要求所有 rank 的实际 backend 都是
`rocm-vocab-parallel-logp-ws2`,且 provenance 不得标记 fallback;
- 显式请求 native backend 但扩展缺失时,测试要求直接 fail fast。

这些测试在非 ROCm 或 GPU 数量不足的环境会跳过;真正执行需要带 ROCm native extension 的多卡机器。

## 没有改什么

- 没有改 issue #241 的 TP/CP 语义;CP 仍不是 logprob 的 merge axis。
- 没有用 RCCL `all_reduce` 取代固定顺序的 all-gather + local merge。
- 没有把 AITER/Composable Kernel 强行接进 strict path。
- 没有把 SM90/TMA/WGMMA 源文件加入 ROCm build。
- 没有修改 Vime provider 的 contract 或 entropy 语义。
- 没有删除任何仓库文件;只是在 ROCm 编译时排除了不适用的 CUDA IPC/PTX 注册。

## 验证结果

在当前 Windows/CUDA 开发环境中:

```text
70 passed, 11 skipped
```

通过的测试包括:

- ROCm backend dispatch 测试;
- issue #241 logprob contract 测试;
- vocab-parallel logprob reference 测试;
- Vime selected-logprob provider 测试。

CUDA extension build 没有完成,原因是当前机器缺少 Microsoft Visual C++ `cl.exe`;同时本地 `nvcc` 是 CUDA 12.6,而 PyTorch 是 CUDA 12.8。当前机器也没有 `hipcc` 和 ROCm runtime,因此以下项目尚未在本地验证:

- gfx942/gfx950 HIP 编译;
- MI300X RCCL TP2/TP4;
- ROCm native tile kernel 与 PyTorch reference 的真实数值/bitwise 对比。

建议在 ROCm 机器上运行:

```bash
PYTORCH_ROCM_ARCH=gfx942 RL_KERNEL_REQUIRE_EXT=1 MAX_JOBS=16 \
python setup.py build_ext --inplace

PYTHONHASHSEED=0 pytest -q -ra \
tests/test_rocm_logprob_backend.py \
tests/test_logprob_contract.py \
tests/test_vocab_parallel_logp.py \
tests/test_distributed_logprob_comparison.py
```

真实多 GPU 验证还需要补充 RCCL TP2/TP4 的运行命令和结果;在拿到 gfx942 结果前,不应把这个 backend 宣称为已完成性能优化,只能称为 strict semantics-preserving ROCm implementation。

## 当前提交

实现尚未推送远程;代码位于当前工作分支 `work/ws2-logprob-rocm`。合并基线提交是 `b811abc`,本次实现文件仍在工作区,待 ROCm 环境验证后再拆分成正式 PR commits。

## 2026-08-23 补充:ROCm 性能调优

在 MI300X 上做完基准测试(`benchmarks/benchmark_rocm_logp.py`,结果见
`benchmarks/results/pr328_rocm_mi300x/report.md`)后,对 ROCm backend 做了两处调整,
TP contract、tile 顺序 merge、selected-target 传输和 active-mask 语义都没有变:

1. 新增 ROCm 专用文件 `csrc/hip/hip_deterministic_logp_kernel.hip`(只在 ROCm 构建时编译,
共享的 `csrc/deterministic_logp_kernel.cu` 保持 SM90 调优版本不动)。其中
`hip_deterministic_logp_tile_stats` 直接读取 BF16/FP16/FP32 shard(kernel 内部逐元素精确
转成 FP32,并自行过滤 padding 列),不再先做一份 FP32 拷贝;每个线程固定处理 8 个连续
元素(向量化 load),累加顺序只由 `(BlockSize, Vec)` 决定,与 rank、shard 偏移和存储
dtype 无关,所以 TP=n 与 TP=1 仍然 bitwise 一致,BF16 直读与 FP32 上转的 partial 也
bitwise 一致。全 padding 的 tile 现在返回 `(-inf, 0)` identity partial(原来是 `-FLT_MAX`)。
2. 新增 `hip_deterministic_logp_backward`:从保存的输入 shard 一次 fused pass 生成
`grad_logits`(`g_logp * (onehot - p) + g_lse * p`,padding 列为 0,非有限行 `p = 0`),
替代共享 Python autograd 里约 9 次 `[tokens, vocab]` FP32 elementwise pass。
`RocmVocabParallelLogprobOp.apply` 走这条 HIP autograd 路径;`apply_with_entropy`
仍沿用共享路径(带 HIP tile kernel),因为 entropy 梯度本来就需要完整的概率张量。

构建期可调参数:`DETERMINISTIC_LOGP_TILE_BLOCK_SIZE`(默认 128)、
`DETERMINISTIC_LOGP_TILE_VECTOR_ELEMENTS`(默认 8)、`DETERMINISTIC_LOGP_BACKWARD_BLOCK_SIZE`
(默认 256),通过 `setup.py` 的环境变量注入。

## 2026-08-23 补充:Triton vocab-parallel backend

`apply` 所用的 fused autograd 路径被提成共享实现
(`rl_engine/kernels/ops/pytorch/loss/vocab_parallel_logp.py` 里的
`apply_with_kernels` / `VocabParallelLogprobKernels`):backend 只需提供
`tile_stats` 和 `backward` 两个 kernel,TP 传输、固定 tile 顺序 merge、
target ownership、mask 语义全部复用。基于这个接口新增了
`rl_engine/kernels/ops/triton/loss/vocab_parallel_logp.py`
(`triton-vocab-parallel-logp-ws2`):两个 Triton kernel,按 `BLOCK_V=1024`
从 tile 起点分块归约,masked lane 贡献 identity,所以归约顺序只由 `BLOCK_V`
决定,TP=n 与 TP=1 仍 bitwise 一致;同一份源码可在 CUDA 和 ROCm 上运行。
registry 在 `cuda`/`rocm` 平台都注册它,排在 PyTorch reference 之前;ROCm 上
HIP backend 仍然排第一。
Loading
Loading