Skip to content

Merge main (Public Release 26/09) into nv_dev, with sparse MQA-logits integration fixes - #458

Closed
lfr-0531 wants to merge 6 commits into
deepseek-ai:nv_devfrom
lfr-0531:pr/nv-dev-merge-main-2609
Closed

lfr-0531 wants to merge 6 commits into
deepseek-ai:nv_devfrom
lfr-0531:pr/nv-dev-merge-main-2609

Conversation

@lfr-0531

Copy link
Copy Markdown

Summary

Bring nv_dev onto the DeepJIT-based runtime of the 26/09 public release (main 78b6900 = #432 + #441) while keeping every nv_dev feature, plus three follow-up fixes found while integrating the sparse MQA-logits kernels into TensorRT-LLM.

Merge of main (26/09) into nv_dev

Policy of the merge: main's structure is the base and every nv_dev feature is re-applied on top of it.

Now on this branch from main:

  • DeepJIT submodule replaces csrc/jit/* and fmt (jit->compile / jit->launch, C++20, std::format); CUDA >= 12.9
  • Sparse MQA logits for the DeepSeek-V4.1 hierarchical sparse indexer (fp8_fp4_sparse_mqa_logits, fp8_fp4_paged_sparse_mqa_logits and their metadata builders)
  • MegaGate, Mega mHC, fp8xfp8 MegaMoE weights, MegaMoESignals workspace, L2 readiness mask, GEMM alpha / deterministic paths, scheduled MQA metadata
  • Fix Mega MoE task info slot release ordering #441 task-info slot release ordering fence in the MegaMoE scheduler

nv_dev features re-applied on the new runtime:

  • SM120/SM121: FP8/FP4 GEMM (AB-swap, split-K), BF16 GEMM, bmk/bnk einsum, TF32 HC prenorm GEMM, MQA logits (contiguous + paged), per-arch contiguous M/K alignment, packed UE8M0 SF on SM120
  • NVFP4 (fp4xfp4) MegaMoE with BF16 shared experts (MmaKind::NVFP4 kept next to upstream's MmaKind::MXF4; bit-based buffer layout; legacy l2_full_count signal for the count protocol)
  • SiTU activation for the FP8xFP4 MegaMoE
  • SM90 FP8 MegaMoE
  • FP16-weights SM100 MQA logits (two-CTA FP16-accumulator kernel)
  • SM90 paged MQA logits: block_kv 32, varlen metadata, next_n=4 multicast clusters
  • Standalone smxx_clean_logits kept for the kernels that do not fuse cleaning (SM120 and FP16-weights paths)
  • Barrier timeout policy (Diagnostic / TrapOnly) and the configurable DG_JIT_BARRIER_TIMEOUT_SECONDS
  • Tests for all of the above merged into the upstream test files

Follow-up fixes

  • get_sparse_mqa_logits_workspace: frameworks capture CUDA graphs on a dedicated side stream but replay on the launching stream, so the first sparse-logits call on the capture stream used to trip the pre-capture assertion. While capturing, reuse a workspace already warmed up on the device.
  • get_bf16_mega_gate_config: return std::map<std::string, int> instead of py::dict so generated Python stubs are dict[str, int] (strict type checkers reject a bare dict); the Python-side result is unchanged.
  • Paged sparse MQA logits: accept any 16-byte aligned KV page stride. The kernel addresses pages directly and copies them with 16-byte cp.async chunks, so page-footer caches that are not padded to 512 bytes (for example 64 x (64 + 4) = 4352-byte pages) can be read in place.

Validation

GB300 (NGC PyTorch 26.08, CUDA 13.4), standalone test suites: test_attention (contiguous MQA logits full, paged sampled cases, sparse MQA logits 161/161 bitwise vs dense), test_fp8_fp4, test_bf16, test_einsum, test_layout, test_hyperconnection, test_mega_moe (4 ranks), test_nvfp4_mega_moe, test_mega_moe_situ. SM90 and SM120 paths compile but were not run. test_mega_gate / test_mega_mhc need tilelang, which was not available in that environment.

The branch is consumed by TensorRT-LLM's DeepSeek-V4.1 bring-up (CSA2 indexer using the sparse MQA-logits kernels for the candidate-consuming layers), where the DeepGEMM-side changes above were exercised end to end (unit suite, GSM8K, long-context throughput on GB300).

Note on DeepJIT

third-party/deep_jit stays at upstream e5bdee2. Its utils/exception.hpp includes <elfutils/libdwfl.h> for declarations only (libdw is dlopened at runtime), so build hosts need the elfutils development headers; making that include optional (__has_include with a symbol-only backtrace fallback) is proposed separately to DeepJIT.

zheanxu and others added 6 commits September 10, 2026 14:31
* Public release 26/09

* Update News
…i#441) into nv_dev

Bring nv_dev onto the DeepJIT-based runtime introduced by the 26/09 public
release while keeping every nv_dev feature:

Upstream (main) additions now on this branch
- DeepJIT submodule replaces csrc/jit/* and fmt (jit->compile / jit->launch,
  C++20, std::format); CUDA >= 12.9 required
- Sparse MQA logits for the DeepSeek-V4.1 hierarchical sparse indexer
  (fp8_fp4_sparse_mqa_logits, fp8_fp4_paged_sparse_mqa_logits + metadata)
- MegaGate, Mega mHC, fp8xfp8 MegaMoE weights, MegaMoESignals workspace,
  L2 readiness mask, GEMM alpha / deterministic paths, scheduled MQA metadata
- deepseek-ai#441: task-info slot release ordering fence in the MegaMoE scheduler

nv_dev features re-applied on top of the new runtime
- SM120/SM121: FP8/FP4 GEMM (AB-swap, split-K), BF16 GEMM, bmk/bnk einsum,
  TF32 HC prenorm GEMM, MQA logits (contiguous + paged), per-arch contiguous
  M/K alignment, packed UE8M0 SF on SM120
- NVFP4 (fp4xfp4) MegaMoE with BF16 shared experts (MmaKind::NVFP4 kept next
  to upstream's MXF4; bit-based buffer layout; legacy l2_full_count signal)
- SiTU activation for the FP8xFP4 MegaMoE
- SM90 FP8 MegaMoE
- FP16-weights SM100 MQA logits (two-CTA FP16 accumulator kernel)
- SM90 paged MQA logits: block_kv 32, varlen metadata, next_n=4 multicast
  clusters
- Standalone smxx_clean_logits kept for the kernels that do not fuse cleaning
  (SM120 and FP16-weights paths)
- Barrier timeout policy (Diagnostic / TrapOnly) and the configurable
  DG_JIT_BARRIER_TIMEOUT_SECONDS (previously a TensorRT-LLM patch)
- Test coverage for all of the above merged into the upstream test files

Validated on GB300 (NGC PyTorch 26.08, CUDA 13.4): test_attention (contiguous
+ sampled paged + full sparse), test_fp8_fp4, test_bf16, test_einsum,
test_layout, test_hyperconnection, test_mega_moe (4 ranks),
test_nvfp4_mega_moe, test_mega_moe_situ. SM90 and SM120 paths are compiled but
not run here.
Frameworks such as TensorRT-LLM capture CUDA graphs on a dedicated side
stream after warming up eagerly on the launching stream. The sparse MQA
logits workspace is keyed by stream, so the first call on the capture
stream used to trip the pre-capture assertion. Captured work replays on
the launching stream, so fall back to that stream's warmed-up workspace
while capturing instead of recording a new allocation into the graph.

Signed-off-by: Fanrong Li <23290157+lfr-0531@users.noreply.github.com>
pybind11 converts std::map<std::string, int> to a Python dict exactly like
py::dict, but generated stubs become dict[str, int] instead of a bare dict,
which strict type checkers reject.

Signed-off-by: Fanrong Li <23290157+lfr-0531@users.noreply.github.com>
The paged sparse kernel addresses pages directly and copies them with
16-byte cp.async chunks, so a page stride only has to be a multiple of 16
bytes. Caches whose page-footer layout is not padded to 512 bytes (for
example 64 x (64 + 4) = 4352-byte pages) can now be read in place.

Signed-off-by: Fanrong Li <23290157+lfr-0531@users.noreply.github.com>
Comment thread csrc/apis/layout.hpp
Comment on lines +100 to +101
DG_HOST_ASSERT((arch_major == 9 and gran_k == 128 and k_alignment == 128) or
(arch_major == 10 and (gran_k == 32 or gran_k == 128) and k_alignment % 128 == 0));

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🔴 critical: 在 K-grouped SF 校验中保留 SM120 支持: 在 SM120/SM121 上,新增断言恒为假,导致非空的 k_grouped_fp8_gemm_nt_contiguous 和 k_grouped_fp8_gemm_tn_contiguous 在启动内核前直接失败,后面保留的 arch_major == 12 转换分支不可达。下游 get_k_grouped_mn_major_tma_aligned_packed_ue8m0_tensor 也新增了仅允许 SM100 的断言,预打包校验还排除了 SM120 原本支持的 gran_k=128;这些校验需要一起按架构保留原有支持。

🤖 v6

Comment thread csrc/apis/attention.hpp
Comment on lines +270 to +272
for (const auto& [other_stream, other_workspace] : workspaces) {
if (other_workspace.defined() and other_workspace.device() == options.device())
return other_workspace;

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🔴 critical: 避免跨流复用仍可能被使用的稀疏工作区: 当多个捕获图在不同流并发重放,或图重放与其他流的 eager 调用并发时,这里会让独立执行共享同一个工作区。按设备选取 unordered_map 中首个张量并不能保证它属于重放流。sm100_sparse_mqa_logits_metadata 会修改并重置其中的 num_kv_splits、next_q_offset、num_finished_ctas 和 Q-block 信息,共享会破坏任务分配及调度元数据。应保留独立工作区,或建立明确的串行使用保证。

🤖 v6

Comment on lines +477 to +482
for (uint32_t chunk_base = 0; chunk_base < kNumChunksPerKVBlock; chunk_base += 16) {
const uint32_t chunk_idx = chunk_base + lane_idx % 16;
ptx::cp_async_cg<256>(reinterpret_cast<const uint4*>(kv_block) + chunk_idx,
reinterpret_cast<uint4*>(smem.kv[kv_stage_idx][0])
+ (kv_block_base_in_split + kv_block_offset) * kNumChunksPerKVBlock +
get_swizzled_kv_chunk_idx<kHeadDim / kPackFactor>(chunk_idx));

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🔴 critical: 为连续稀疏 KV 的尾块复制增加边界保护: 连续 KV 长度不足一个完整稀疏块时,这里仍无条件复制整块。例如接口允许 MXFP8 kv[4,128] 配合 sparse_block_kv=8,但实际只有 512 字节,首块复制却读取 1024 字节;SF 复制也有同样的尾部越界。忽略额外 logits 无法使这些读取安全,禁用缓存分配器或使用精确分配的外部缓冲区时会产生非法读。应屏蔽并补零越界部分,或显式保证足够的尾部 padding。

🤖 v6

Comment thread csrc/utils/layout.hpp
Comment on lines +38 to +39
return jit->device.get_arch_major() == 9 or
(a.scalar_type() == kPackedFP4 and b.scalar_type() == kPackedFP4);

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🟡 warning: 将 FP4xFP4 的新增 K-major 限制限定到 SM100: 在 SM120 上,只要两个输入都是 FP4 且存在 MN-major 输入,这个新条件就会让 GEMM 入口断言失败。原有 fp8_fp4_gemm_nt_sm120 会通过 sm120_to_k_major 和 fp4_repack_to_k_major 转换这些输入,M-grouped 路径也保留了相同转换;现在它们被提前校验阻断。应将 SM100 的 FP4xFP4 布局限制按架构应用,避免删除 SM120 已支持的调用方式。

🤖 v6

Comment thread csrc/apis/attention.hpp
if (c10::cuda::currentStreamCaptureStatusMayInitCtx() != c10::cuda::CaptureStatus::None) {
// Frameworks capture on a dedicated side stream but replay on the stream that
// launched the eager warm-up, so reuse that stream's workspace instead of
// recording a fresh allocation and its zeroing into the graph.

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🔴 critical: The capture-time fallback returns the first defined workspace in unordered_map iteration order. The workspace holds live counters (num_kv_splits, next_q_offset, num_finished_ctas) that the metadata kernel resets in its last CTA, which is why it is keyed per stream. Borrowing another stream's workspace is only safe if the graph replays on (or is stream-ordered with) the owning stream and no eager metadata kernel runs on that stream concurrently with a replay. With more than one warmed stream (e.g. compute + comm side stream) the pick is implementation-defined and a mismatch produces wrong metadata silently rather than a crash. Please make the selection deterministic (e.g. prefer the device's current/default stream at warm-up, or record an explicit per-device primary workspace) and document the replay-stream contract on get_sparse_mqa_logits_metadata / get_paged_sparse_mqa_logits_metadata.

🤖 v5

Comment thread csrc/apis/attention.hpp
// Frameworks capture on a dedicated side stream but replay on the stream that
// launched the eager warm-up, so reuse that stream's workspace instead of
// recording a fresh allocation and its zeroing into the graph.
for (const auto& [other_stream, other_workspace] : workspaces) {

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🟡 warning: The capture-time fallback in get_sparse_mqa_logits_workspace picks an arbitrary already-warmed workspace on the same device by iterating an unordered_map. The selection order is unspecified and is not tied to the stream the captured graph will actually replay on, so when multiple streams have warmed a workspace the recorded scratch address can differ between capture invocations. It works today only because every workspace is a statically-owned, fixed-size buffer and the last metadata CTA resets the state, but the aliasing contract is implicit and fragile. Please select the replay/launching stream's workspace deterministically (e.g. key the lookup by the capture stream or the single warmed stream), and assert dtype/numel on the reused tensor. Also note other_stream is unused and can trip -Wunused-variable on stricter compilers.

🤖 v4

Comment thread csrc/apis/attention.hpp
head_dim_with_sf == (is_fp4 ? head_dim / 2 : head_dim) + static_cast<int>(sizeof(int)));
DG_HOST_ASSERT(fused_kv_cache.scalar_type() == torch::kUInt8 and fused_kv_cache.stride(1) == head_dim_with_sf and
fused_kv_cache.stride(3) == 1 and fused_kv_cache.stride(0) <= std::numeric_limits<int>::max() and
fused_kv_cache.stride(0) % 16 == 0); // pages are read with 16-byte cp.async chunks

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🟡 warning: Relaxing the page-stride check to stride(0) % 16 == 0 is consistent with the 16-byte cp.async chunks in PagedSparseKVAccessor, but the 16-byte requirement applies to the actual global addresses, not only the stride. Since only stride(1)/stride(3) are constrained, a fused_kv_cache view with a non-16-byte-aligned storage offset can pass this assert and then produce misaligned cp.async addresses. Add a base-pointer check such as reinterpret_cast<uintptr_t>(fused_kv_cache.data_ptr()) % 16 == 0.

🤖 v4

Comment thread csrc/apis/mega_gate.hpp
return {topk_idx, topk_weights};
}

static std::map<std::string, int> get_bf16_mega_gate_config(const int& num_tokens, const int& hidden,

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🔵 suggestion: Changing the return type from pybind11::dict to std::map<std::string, int> does give precise dict[str, int] stubs, but it also changes the resulting Python dict's key order from insertion order to lexicographic order (block_tokens, num_expert_groups, num_gate_warpgroups, num_mma_ctas, num_sms, num_split_k). Values are unchanged, yet any consumer that snapshots key order (or compares against the previous py::dict ordering) will observe a behavioral difference. If order is part of the contract, keep an insertion-ordered container or document the new ordering; otherwise it is only a minor compatibility note.

🤖 v4

#endif

// Timeout in cycles, at 2 GHz
constexpr int64_t kNumTimeoutCycles = static_cast<int64_t>(DG_BARRIER_TIMEOUT_SECONDS) * 2000000000ll;

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🔵 suggestion: kNumTimeoutCycles converts DG_BARRIER_TIMEOUT_SECONDS to cycles assuming a fixed 2 GHz clock (* 2000000000ll). On parts or clock domains running at a different rate the effective hang-detection timeout will be wrong, which can turn a slow-but-healthy rendezvous into a spurious trap or delay diagnostics. Consider deriving the rate at runtime or documenting that the timeout is approximate.

🤖 v4

Comment thread tests/test_attention.py
@@ -67,7 +67,7 @@ def sample_mqa_cases(name: str, cases: List[tuple]) -> List[tuple]:
if num_cases is None:

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🔵 suggestion: Validation coverage for this merge is uneven: the PR description states SM90 and SM120 paths were only compiled, and test_mega_gate/test_mega_mhc were skipped because tilelang was unavailable. That leaves the SM90 next_n=4 multicast metadata, per-arch contiguous M/K alignment, and the MegaGate/MHC machinery without runtime coverage. Please add a CI job or an explicit xfail/skip record so regressions in those paths are visible rather than silently untested.

🤖 v4

@ds-review-bot

Copy link
Copy Markdown
Collaborator

🤖 ds-review-bot Code Review

v6

变更引入了 SM120 功能回归,以及稀疏 MQA 的跨流工作区竞争和尾块越界读取风险。当前环境缺少 PyTorch、CUDA 编译器及 GPU,结论基于 diff、调用链和内核边界检查,未复跑 GPU 测试。

v5

Reviewed the merge commit (0274ad7) and the three follow-up commits against main@78b6900 by reading the diff (no build/run possible in this environment). Confirmed OK: the merge follows the stated policy (main's DeepJIT runtime as base, nv_dev features re-applied); the #441 fence lands in the shared MegaMoEScheduler::release_task_info() so the sm100 fp8xfp4, bf16 and the new fp4xfp4 kernels all get it (the SM90 adapter has no task-info slots, nothing to port); MmaKind::NVFP4 is handled in every switch (heuristics/utils.hpp, heuristics/sm100.hpp, heuristics/mega_moe.hpp); deep_jit and cutlass submodule pins match main exactly; the 16-byte page-stride relaxation is sound given the accessor's 4-byte/16-byte cp.async reads; std::map conversion for get_bf16_mega_gate_config compiles as-is since pybind11/stl.h arrives via torch/python.h (only key order changes from insertion to sorted, equality unaffected). Issues: the graph-capture workspace fallback picks an arbitrary stream's workspace (unordered_map order) which can silently produce wrong metadata if it is not the replay stream, and neither it nor the 16-byte page-stride path is exercised by any in-tree test although both are the motivation for the TensorRT-LLM integration fixes. Remaining items are dead code left over from the re-apply, README gaps (SM120/SM121 support, DG_JIT_BARRIER_TIMEOUT_SECONDS), and a duplicate include. Recommend making the workspace selection deterministic and adding test coverage for both fixes before merging; the rest is cleanup.

v4

This MR rebases nv_dev onto the 26/09 public-release runtime (DeepJIT compile/launch replacing csrc/jit + fmt, C++20/std::format, CUDA >= 12.9) and re-applies the full nv_dev feature surface on top of main, then adds three sparse MQA-logits integration fixes. The merge structure looks correct: the merge commit has the expected two parents (nv_dev base + main #432/#441), the DeepJIT submodule is pinned at e5bdee2, csrc/jit and third-party/fmt are removed, and the three follow-up fixes (capture-safe sparse-logits workspace, typed get_bf16_mega_gate_config, 16-byte KV page stride) are all present. Most of the ~130 files are mechanical jit->compile/launch and includes migrations, which makes targeted review of the handful of behavioral changes important. The main risks I see are build portability of the DeepJIT dependency, the implicit/nondeterministic scratch-buffer reuse during CUDA graph capture, and the fact that the SM90/SM120 and MegaGate/MHC paths are only compile-checked. Details and smaller findings below.

Files reviewed: 130
Issues found: 🔴 5 critical | 🟡 5 warning | 🔵 8 suggestion
Inline comments posted: 16
General comments (无法定位到 diff): 2


📍 未定位到 diff 的评论

🔵 suggestion deep_gemm/include/deep_gemm/comm/barrier_with_timeout_policy.cuh:L12: handle_grid_sync_timeout and handle_nvlink_barrier_timeout were the nv_dev call sites, but after re-basing on main's wait_until the policy branch is inlined in barrier.cuh and nothing references these two functions anymore (git grep finds no callers). They also print the old message without the (%ds) timeout that barrier.cuh now includes. Suggest deleting them and keeping only the BarrierTimeoutPolicy enum in this header. 🤖 v5

🔴 critical third-party/deep_jit/include/deep_jit/utils/exception.hpp:L3: The unconditional #include <elfutils/libdwfl.h> makes the whole DeepGEMM build depend on the elfutils development headers even though libdw is only dlopen()ed at runtime and every dwfl symbol is resolved through dlsym(). Any wheel/build host without elfutils-dev now fails to compile, independent of this MR's GPU features. Please guard the include with #if __has_include(<elfutils/libdwfl.h>) and provide the symbol-only backtrace fallback (declare the small set of needed types locally, or compile out the dwfl lookup path) as the PR note proposes; this is a portability blocker for the change set even if the DeepJIT patch lands separately. 🤖 v4

@lucifer1004

Copy link
Copy Markdown
Collaborator

Hi Fanrong, did you see #447?

@lfr-0531

Copy link
Copy Markdown
Author

Duplicated with #447. Closing.

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

4 participants