Repository navigation
Conversation
* 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>
| 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)); |
There was a problem hiding this comment.
🔴 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
| for (const auto& [other_stream, other_workspace] : workspaces) { | ||
| if (other_workspace.defined() and other_workspace.device() == options.device()) | ||
| return other_workspace; |
There was a problem hiding this comment.
🔴 critical: 避免跨流复用仍可能被使用的稀疏工作区: 当多个捕获图在不同流并发重放,或图重放与其他流的 eager 调用并发时,这里会让独立执行共享同一个工作区。按设备选取 unordered_map 中首个张量并不能保证它属于重放流。sm100_sparse_mqa_logits_metadata 会修改并重置其中的 num_kv_splits、next_q_offset、num_finished_ctas 和 Q-block 信息,共享会破坏任务分配及调度元数据。应保留独立工作区,或建立明确的串行使用保证。
🤖 v6
| 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)); |
There was a problem hiding this comment.
🔴 critical: 为连续稀疏 KV 的尾块复制增加边界保护: 连续 KV 长度不足一个完整稀疏块时,这里仍无条件复制整块。例如接口允许 MXFP8 kv[4,128] 配合 sparse_block_kv=8,但实际只有 512 字节,首块复制却读取 1024 字节;SF 复制也有同样的尾部越界。忽略额外 logits 无法使这些读取安全,禁用缓存分配器或使用精确分配的外部缓冲区时会产生非法读。应屏蔽并补零越界部分,或显式保证足够的尾部 padding。
🤖 v6
| return jit->device.get_arch_major() == 9 or | ||
| (a.scalar_type() == kPackedFP4 and b.scalar_type() == kPackedFP4); |
There was a problem hiding this comment.
🟡 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
| 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. |
There was a problem hiding this comment.
🔴 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
| // 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) { |
There was a problem hiding this comment.
🟡 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
| 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 |
There was a problem hiding this comment.
🟡 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
| return {topk_idx, topk_weights}; | ||
| } | ||
|
|
||
| static std::map<std::string, int> get_bf16_mega_gate_config(const int& num_tokens, const int& hidden, |
There was a problem hiding this comment.
🔵 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; |
There was a problem hiding this comment.
🔵 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
| @@ -67,7 +67,7 @@ def sample_mqa_cases(name: str, cases: List[tuple]) -> List[tuple]: | |||
| if num_cases is None: | |||
There was a problem hiding this comment.
🔵 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 Code Reviewv6变更引入了 SM120 功能回归,以及稀疏 MQA 的跨流工作区竞争和尾块越界读取风险。当前环境缺少 PyTorch、CUDA 编译器及 GPU,结论基于 diff、调用链和内核边界检查,未复跑 GPU 测试。 v5Reviewed 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. v4This 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 📍 未定位到 diff 的评论🔵 suggestion 🔴 critical |
|
Hi Fanrong, did you see #447? |
|
Duplicated with #447. Closing. |
Summary
Bring
nv_devonto the DeepJIT-based runtime of the 26/09 public release (main78b6900 = #432 + #441) while keeping everynv_devfeature, plus three follow-up fixes found while integrating the sparse MQA-logits kernels into TensorRT-LLM.Merge of
main(26/09) intonv_devPolicy of the merge:
main's structure is the base and everynv_devfeature is re-applied on top of it.Now on this branch from
main:csrc/jit/*andfmt(jit->compile/jit->launch, C++20,std::format); CUDA >= 12.9fp8_fp4_sparse_mqa_logits,fp8_fp4_paged_sparse_mqa_logitsand their metadata builders)MegaMoESignalsworkspace, L2 readiness mask, GEMM alpha / deterministic paths, scheduled MQA metadatanv_devfeatures re-applied on the new runtime:MmaKind::NVFP4kept next to upstream'sMmaKind::MXF4; bit-based buffer layout; legacyl2_full_countsignal for the count protocol)smxx_clean_logitskept for the kernels that do not fuse cleaning (SM120 and FP16-weights paths)DG_JIT_BARRIER_TIMEOUT_SECONDSFollow-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: returnstd::map<std::string, int>instead ofpy::dictso generated Python stubs aredict[str, int](strict type checkers reject a baredict); the Python-side result is unchanged.cp.asyncchunks, 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_mhcneedtilelang, 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_jitstays at upstream e5bdee2. Itsutils/exception.hppincludes<elfutils/libdwfl.h>for declarations only (libdw isdlopened at runtime), so build hosts need the elfutils development headers; making that include optional (__has_includewith a symbol-only backtrace fallback) is proposed separately to DeepJIT.