Skip to content

[Public release 26/04] Introducing Mega MoE, FP4 Indexer and other features/fixes - #304

Merged
LyricZhao merged 7 commits into
mainfrom
public-release-260416
Apr 17, 2026
Merged

LyricZhao merged 7 commits into
mainfrom
public-release-260416

Conversation

@LyricZhao

@LyricZhao LyricZhao commented Apr 16, 2026 •

Copy link
Copy Markdown
Collaborator

New features

  • Mega MoE, fusing & overlapping dispatch/linear 1/SwiGLU/linear 2/combine into a single mega-kernel, overlapping NVLink communication and tensor core computation
  • FP4 Indexer (MQA logits) with larger MTP support
  • FP8 x FP4 GEMM
  • PDL
  • Refactors on GEMM heuristics
  • Faster JIT compilation
  • GEMM optimizations (dynamic swap A/B, much faster MoE GEMM)
  • DeepEPv2 MoE GEMM layout

Bug fixes

  • JIT may crash on distributed FS
  • Some kernel hangs and IMA

Contributors

Additional notes

Mega MoE is still under development and optimizations, stay tuned and optimization ideas are welcome!
Disclaimer: this release is only related to DeepGEMM's development, has nothing to do with internal model release.

@LyricZhao LyricZhao changed the title [Public release 26/04] Introducing Mega MoE, FP4 Indexer and other bug fixes [Public release 26/04] Introducing Mega MoE, FP4 Indexer and other features/fixes Apr 16, 2026
@LyricZhao
LyricZhao force-pushed the public-release-260416 branch from a050d09 to 7dbe077 Compare April 16, 2026 08:23
@LyricZhao
LyricZhao requested a review from zheanxu April 16, 2026 08:48
@yuho8818

Copy link
Copy Markdown

LGTM

Barry-Delaney added a commit to Barry-Delaney/DeepGEMM that referenced this pull request Apr 21, 2026
Brings in upstream commits d30fc36 (TMEM alloc/dealloc sync fix, deepseek-ai#292)
and 7f2a703 (public release 26/04: Mega MoE, FP4 indexer, and related
refactors, deepseek-ai#304), plus the two NV-side feature commits:

  - Enable bf16 bias output for FP8 GEMM on SM100
  - SM90 paged MQA logits: support kv_block=32 and next_n=4

The upstream deepseek-ai#304 refactor reshapes the attention/MQA layer; conflicts
with the pre-merge nv_dev code are resolved by taking the post-deepseek-ai#304
version. Also drops the redundant __syncthreads() in
sm100_bmk_bnk_mn.cuh -- deepseek-ai#292 (d30fc36) already covers the same hazard
via NamedBarrier.sync.
Barry-Delaney pushed a commit to Barry-Delaney/DeepGEMM that referenced this pull request Apr 21, 2026
Extend the SM90 FP8 paged MQA logits kernel and host wrapper to cover
block_kv=32 and next_n=4 on post-deepseek-ai#304 code.

Kernel (deep_gemm/include/deep_gemm/impls/sm90_fp8_paged_mqa_logits.cuh):
  - Keep MMA tile at kComputeBlockKV=64 and let block_kv vary (32 or
    64). Each WGMMA iteration covers kNumBlocksPerMMA = 64 / block_kv
    physical blocks with one TMA copy each.
  - Load block_table entries with a loop of scalar __ldg's so
    block_table_stride no longer has to be a multiple of
    kNumBlocksPerMMA.
  - Add kNumKVMulticast template parameter for next_n=4; compute
    kNextNPerCTA = kNextN / kNumKVMulticast and split Q rows across
    two CTAs in a cluster via cta_rank_in_cluster, keeping smem and
    register budgets within SM90 limits (WGMMA N=256 otherwise
    exceeds kNumMathRegisters=104).
  - Use cute::cluster_id_in_grid().x for the scheduler index.

Host wrapper (csrc/jit_kernels/impls/smxx_fp8_fp4_paged_mqa_logits.hpp):
  - Select num_kv_multicast = 2 when SM90+next_n=4 else 1. Build
    Q/weights TMA descriptors with next_n_per_cta * num_heads box,
    size smem from next_n_per_cta, launch with
    cluster_dim=num_kv_multicast.
  - Emit the extra kNumKVMulticast template arg only for SM90; SM100
    instantiation is unchanged.

Metadata (deep_gemm/include/deep_gemm/scheduler/paged_mqa_logits.cuh,
csrc/apis/attention.hpp):
  - Pass num_next_n_atoms in as a runtime parameter instead of
    hard-coding it from next_n. SM90 multicast passes 1 (no
    atomization, one entry per cluster); SM100 keeps 2 for next_n=4.
  - Relax the schedule_meta size assertion to num_sms /
    num_kv_multicast + 1 to match the cluster grid.

Other:
  - Add utils::shfl_sync helper in common/utils.cuh.
  - tests/test_attention.py: extend enumerate_paged_mqa_logits to
    cover block_kv=32 and next_n=4 on SM90; call
    get_paged_mqa_logits_metadata with num_clusters = num_sms /
    num_kv_multicast; initialize block_table with torch.zeros so
    unused slots hold a valid block id.
Barry-Delaney added a commit to Barry-Delaney/DeepGEMM that referenced this pull request Apr 21, 2026
Brings in upstream commits d30fc36 (TMEM alloc/dealloc sync fix, deepseek-ai#292)
and 7f2a703 (public release 26/04: Mega MoE, FP4 indexer, and related
refactors, deepseek-ai#304), plus the two NV-side feature commits:

  - Enable bf16 bias output for FP8 GEMM on SM100
  - SM90 paged MQA logits: support kv_block=32 and next_n=4

The upstream deepseek-ai#304 refactor reshapes the attention/MQA layer; conflicts
with the pre-merge nv_dev code are resolved by taking the post-deepseek-ai#304
version. Also drops the redundant __syncthreads() in
sm100_bmk_bnk_mn.cuh -- deepseek-ai#292 (d30fc36) already covers the same hazard
via NamedBarrier.sync.
Barry-Delaney pushed a commit to Barry-Delaney/DeepGEMM that referenced this pull request Apr 21, 2026
Extend the SM90 FP8 paged MQA logits kernel and host wrapper to cover
block_kv=32 and next_n=4 on post-deepseek-ai#304 code.

Kernel (deep_gemm/include/deep_gemm/impls/sm90_fp8_paged_mqa_logits.cuh):
  - Keep MMA tile at kComputeBlockKV=64 and let block_kv vary (32 or
    64). Each WGMMA iteration covers kNumBlocksPerMMA = 64 / block_kv
    physical blocks with one TMA copy each.
  - Load block_table entries with a loop of scalar __ldg's so
    block_table_stride no longer has to be a multiple of
    kNumBlocksPerMMA.
  - Add kNumKVMulticast template parameter for next_n=4; compute
    kNextNPerCTA = kNextN / kNumKVMulticast and split Q rows across
    two CTAs in a cluster via cta_rank_in_cluster, keeping smem and
    register budgets within SM90 limits (WGMMA N=256 otherwise
    exceeds kNumMathRegisters=104).
  - Use cute::cluster_id_in_grid().x for the scheduler index.

Host wrapper (csrc/jit_kernels/impls/smxx_fp8_fp4_paged_mqa_logits.hpp):
  - Select num_kv_multicast = 2 when SM90+next_n=4 else 1. Build
    Q/weights TMA descriptors with next_n_per_cta * num_heads box,
    size smem from next_n_per_cta, launch with
    cluster_dim=num_kv_multicast.
  - Emit the extra kNumKVMulticast template arg only for SM90; SM100
    instantiation is unchanged.

Metadata (deep_gemm/include/deep_gemm/scheduler/paged_mqa_logits.cuh,
csrc/apis/attention.hpp):
  - Pass num_next_n_atoms in as a runtime parameter instead of
    hard-coding it from next_n. SM90 multicast passes 1 (no
    atomization, one entry per cluster); SM100 keeps 2 for next_n=4.
  - Relax the schedule_meta size assertion to num_sms /
    num_kv_multicast + 1 to match the cluster grid.

Other:
  - Add utils::shfl_sync helper in common/utils.cuh.
  - tests/test_attention.py: extend enumerate_paged_mqa_logits to
    cover block_kv=32 and next_n=4 on SM90; call
    get_paged_mqa_logits_metadata with num_clusters = num_sms /
    num_kv_multicast; initialize block_table with torch.zeros so
    unused slots hold a valid block id.
Barry-Delaney added a commit to Barry-Delaney/DeepGEMM that referenced this pull request Apr 21, 2026
Brings in upstream commits d30fc36 (TMEM alloc/dealloc sync fix, deepseek-ai#292)
and 7f2a703 (public release 26/04: Mega MoE, FP4 indexer, and related
refactors, deepseek-ai#304), plus the two NV-side feature commits:

  - Enable bf16 bias output for FP8 GEMM on SM100
  - SM90 paged MQA logits: support kv_block=32 and next_n=4

The upstream deepseek-ai#304 refactor reshapes the attention/MQA layer; conflicts
with the pre-merge nv_dev code are resolved by taking the post-deepseek-ai#304
version. Also drops the redundant __syncthreads() in
sm100_bmk_bnk_mn.cuh that was added by nv_dev's 8e866e4; deepseek-ai#304 removed
the equivalent sync upstream (replaced with NamedBarrier.sync in
sibling SM100 kernels) so keeping the nv_dev version would diverge
from upstream's final sync design.
Barry-Delaney pushed a commit to Barry-Delaney/DeepGEMM that referenced this pull request Apr 21, 2026
Extend the SM90 FP8 paged MQA logits kernel and host wrapper to cover
block_kv=32 and next_n=4 on post-deepseek-ai#304 code.

Kernel (deep_gemm/include/deep_gemm/impls/sm90_fp8_paged_mqa_logits.cuh):
  - Keep MMA tile at kComputeBlockKV=64 and let block_kv vary (32 or
    64). Each WGMMA iteration covers kNumBlocksPerMMA = 64 / block_kv
    physical blocks with one TMA copy each.
  - Load block_table entries with a loop of scalar __ldg's so
    block_table_stride no longer has to be a multiple of
    kNumBlocksPerMMA.
  - Add kNumKVMulticast template parameter for next_n=4; compute
    kNextNPerCTA = kNextN / kNumKVMulticast and split Q rows across
    two CTAs in a cluster via cta_rank_in_cluster, keeping smem and
    register budgets within SM90 limits (WGMMA N=256 otherwise
    exceeds kNumMathRegisters=104).
  - Use cute::cluster_id_in_grid().x for the scheduler index.

Host wrapper (csrc/jit_kernels/impls/smxx_fp8_fp4_paged_mqa_logits.hpp):
  - Select num_kv_multicast = 2 when SM90+next_n=4 else 1. Build
    Q/weights TMA descriptors with next_n_per_cta * num_heads box,
    size smem from next_n_per_cta, launch with
    cluster_dim=num_kv_multicast.
  - Emit the extra kNumKVMulticast template arg only for SM90; SM100
    instantiation is unchanged.

Other:
  - Add utils::shfl_sync helper in common/utils.cuh.
  - tests/test_attention.py: extend enumerate_paged_mqa_logits to
    cover block_kv=32 and next_n=4 on SM90; call
    get_paged_mqa_logits_metadata with num_clusters = num_sms /
    num_kv_multicast; initialize block_table with torch.zeros so
    unused slots hold a valid block id.
Barry-Delaney added a commit to Barry-Delaney/DeepGEMM that referenced this pull request Apr 21, 2026
…API mismatch

The post-deepseek-ai#304 metadata kernel hard-codes num_next_n_atoms from next_n,
which assumes the SM100 time-atomization scheme. On SM90 next_n=4 the
cluster-multicast design processes one full q per cluster with no
atomization, so the metadata ends up producing 2x the tasks the kernel
expects and stepping past context_lens when initializing the scheduler
(observed as illegal memory access in sm90_fp8_paged_mqa_logits).

Fixes and hardening:
  - Pipe num_next_n_atoms as a runtime parameter into the metadata
    kernel (deep_gemm/include/deep_gemm/scheduler/paged_mqa_logits.cuh)
    and let the host decide. attention.hpp computes it based on arch +
    next_n + num_kv_multicast: SM90 multicast passes 1 (one entry per
    cluster), SM100 keeps 2 for next_n=4.
  - Relax the schedule_meta size assertion in attention.hpp to
    'num_sms / num_kv_multicast + 1' so the check matches the cluster
    grid rather than the SM grid.
  - Extract 'get_paged_mqa_logits_num_kv_multicast' helper in the
    host wrapper so the formula lives in one place.
  - Add 'DG_HOST_ASSERT(num_sms % num_kv_multicast == 0)' guard in the
    cluster launch path for non-H200/H100-like num_sms.
  - Document the block_table unused-slot contract on the public API:
    callers MUST initialize unused block_table slots to a valid block
    id (e.g. 0) because the kernel may read slots up to num_kv *
    kNumBlocksPerMMA when context_len is not a multiple of
    kComputeBlockKV=64.
Barry-Delaney added a commit to Barry-Delaney/DeepGEMM that referenced this pull request Apr 21, 2026
Brings in upstream commits d30fc36 (TMEM alloc/dealloc sync fix, deepseek-ai#292)
and 7f2a703 (public release 26/04: Mega MoE, FP4 indexer, and related
refactors, deepseek-ai#304), plus three NV-side commits:

  - Enable bf16 bias output for FP8 GEMM on SM100 (xueweil)
  - SM90 paged MQA logits: support kv_block=32 and next_n=4 (Ray)
  - Fix SM90 paged MQA logits for post-deepseek-ai#304 metadata/schedule API mismatch

The upstream deepseek-ai#304 refactor reshapes the attention/MQA layer; conflicts
with the pre-merge nv_dev code are resolved by taking the post-deepseek-ai#304
version. Also drops the redundant __syncthreads() in
sm100_bmk_bnk_mn.cuh that was added by nv_dev's 8e866e4; deepseek-ai#304 removed
the equivalent sync upstream (replaced with NamedBarrier.sync in
sibling SM100 kernels) so keeping the nv_dev version would diverge
from upstream's final sync design.
Barry-Delaney pushed a commit to Barry-Delaney/DeepGEMM that referenced this pull request Apr 21, 2026
Extend the SM90 FP8 paged MQA logits kernel and host wrapper to cover
block_kv=32 and next_n=4 on post-deepseek-ai#304 code.

Kernel (deep_gemm/include/deep_gemm/impls/sm90_fp8_paged_mqa_logits.cuh):
  - Keep MMA tile at kComputeBlockKV=64 and let block_kv vary (32 or
    64). Each WGMMA iteration covers kNumBlocksPerMMA = 64 / block_kv
    physical blocks with one TMA copy each.
  - Load block_table entries with a loop of scalar __ldg's so
    block_table_stride no longer has to be a multiple of
    kNumBlocksPerMMA.
  - Add kNumKVMulticast template parameter for next_n=4; compute
    kNextNPerCTA = kNextN / kNumKVMulticast and split Q rows across
    two CTAs in a cluster via cta_rank_in_cluster, keeping smem and
    register budgets within SM90 limits (WGMMA N=256 otherwise
    exceeds kNumMathRegisters=104).
  - Use cute::cluster_id_in_grid().x for the scheduler index.

Host wrapper (csrc/jit_kernels/impls/smxx_fp8_fp4_paged_mqa_logits.hpp):
  - Select num_kv_multicast = 2 when SM90+next_n=4 else 1. Build
    Q/weights TMA descriptors with next_n_per_cta * num_heads box,
    size smem from next_n_per_cta, launch with
    cluster_dim=num_kv_multicast.
  - Emit the extra kNumKVMulticast template arg only for SM90; SM100
    instantiation is unchanged.

Other:
  - Add utils::shfl_sync helper in common/utils.cuh.
  - tests/test_attention.py: extend enumerate_paged_mqa_logits to
    cover block_kv=32 and next_n=4 on SM90; call
    get_paged_mqa_logits_metadata with num_clusters = num_sms /
    num_kv_multicast; initialize block_table with torch.zeros so
    unused slots hold a valid block id.
Barry-Delaney added a commit to Barry-Delaney/DeepGEMM that referenced this pull request Apr 21, 2026
…API mismatch

The post-deepseek-ai#304 metadata kernel hard-codes num_next_n_atoms from next_n,
which assumes the SM100 time-atomization scheme. On SM90 next_n=4 the
cluster-multicast design processes one full q per cluster with no
atomization, so the metadata ends up producing 2x the tasks the kernel
expects and stepping past context_lens when initializing the scheduler
(observed as illegal memory access in sm90_fp8_paged_mqa_logits).

Fixes and hardening:
  - Pipe num_next_n_atoms as a runtime parameter into the metadata
    kernel (deep_gemm/include/deep_gemm/scheduler/paged_mqa_logits.cuh)
    and let the host decide. attention.hpp computes it based on arch +
    next_n + num_kv_multicast: SM90 multicast passes 1 (one entry per
    cluster), SM100 keeps 2 for next_n=4.
  - Relax the schedule_meta size assertion in attention.hpp to
    'num_sms / num_kv_multicast + 1' so the check matches the cluster
    grid rather than the SM grid.
  - Extract 'get_paged_mqa_logits_num_kv_multicast' helper in the
    host wrapper so the formula lives in one place.
  - Add 'DG_HOST_ASSERT(num_sms % num_kv_multicast == 0)' guard in the
    cluster launch path for non-H200/H100-like num_sms.
  - Document the block_table unused-slot contract on the public API:
    callers MUST initialize unused block_table slots to a valid block
    id (e.g. 0) because the kernel may read slots up to num_kv *
    kNumBlocksPerMMA when context_len is not a multiple of
    kComputeBlockKV=64.
Barry-Delaney added a commit to Barry-Delaney/DeepGEMM that referenced this pull request Apr 21, 2026
Brings in upstream commits d30fc36 (TMEM alloc/dealloc sync fix, deepseek-ai#292)
and 7f2a703 (public release 26/04: Mega MoE, FP4 indexer, and related
refactors, deepseek-ai#304), plus three NV-side commits:

  - Enable bf16 bias output for FP8 GEMM on SM100 (xueweil)
  - SM90 paged MQA logits: support kv_block=32 and next_n=4 (Ray)
  - Fix SM90 paged MQA logits for post-deepseek-ai#304 metadata/schedule API mismatch

The upstream deepseek-ai#304 refactor reshapes the attention/MQA layer; conflicts
with the pre-merge nv_dev code are resolved by taking the post-deepseek-ai#304
version. Also drops the redundant __syncthreads() in
sm100_bmk_bnk_mn.cuh that was added by nv_dev's 8e866e4; deepseek-ai#304 removed
the equivalent sync upstream (replaced with NamedBarrier.sync in
sibling SM100 kernels) so keeping the nv_dev version would diverge
from upstream's final sync design.
@huangzhilin-hzl

Copy link
Copy Markdown

Hi @LyricZhao tests/test_mega_moe.py currently depends on deep_ep.ElasticBuffer, but I could not find ElasticBuffer in the public deepseek-ai/DeepEP branches; was this test written against a private/internal DeepEP version?

@LyricZhao

Copy link
Copy Markdown
Collaborator Author

@huangzhilin-hzl deep_ep.ElasticBuffer (DeepEP V2) will be released in this week. And, you can still skip the legacy path test and just look at the perf numbers.

@CyberQin

Copy link
Copy Markdown

deep_ep基础设施还没更新完啊,那model release不是更佳悬了,yifan的爆料要失真了

Barry-Delaney added a commit to Barry-Delaney/DeepGEMM that referenced this pull request Apr 22, 2026
…API mismatch

The post-deepseek-ai#304 metadata kernel hard-codes num_next_n_atoms from next_n,
which assumes the SM100 time-atomization scheme. On SM90 next_n=4 the
cluster-multicast design processes one full q per cluster with no
atomization, so the metadata ends up producing 2x the tasks the kernel
expects and stepping past context_lens when initializing the scheduler
(observed as illegal memory access in sm90_fp8_paged_mqa_logits).

Fixes and hardening:
  - Pipe num_next_n_atoms as a runtime parameter into the metadata
    kernel (deep_gemm/include/deep_gemm/scheduler/paged_mqa_logits.cuh)
    and let the host decide. attention.hpp computes it based on arch +
    next_n + num_kv_multicast: SM90 multicast passes 1 (one entry per
    cluster), SM100 keeps 2 for next_n=4.
  - Relax the schedule_meta size assertion in attention.hpp to
    'num_sms / num_kv_multicast + 1' so the check matches the cluster
    grid rather than the SM grid.
  - Extract 'get_paged_mqa_logits_num_kv_multicast' helper in the
    host wrapper so the formula lives in one place.
  - Add 'DG_HOST_ASSERT(num_sms % num_kv_multicast == 0)' guard.
  - Normalize the public 'get_paged_mqa_logits_metadata' helper so
    callers always pass the natural 'get_num_sms()' and physical
    'block_kv': the helper now computes num_clusters internally, sizes
    the schedule as '{num_clusters + 1, 2}', and uses the compute-block
    size on SM90 regardless of the physical block_kv the caller supplies.
  - Tighten the SM90 block_table load guard to a per-physical-block
    'phys_idx < ceil_div(context_len, BLOCK_KV)' check so BLOCK_KV=32
    with odd physical block counts cannot read past the natural row
    end.
Barry-Delaney added a commit to Barry-Delaney/DeepGEMM that referenced this pull request Apr 22, 2026
Brings in upstream commits d30fc36 (TMEM alloc/dealloc sync fix, deepseek-ai#292)
and 7f2a703 (public release 26/04: Mega MoE, FP4 indexer, and related
refactors, deepseek-ai#304), plus three NV-side commits:

  - Enable bf16 bias output for FP8 GEMM on SM100 (xueweil)
  - SM90 paged MQA logits: support kv_block=32 and next_n=4 (Ray)
  - Fix SM90 paged MQA logits for post-deepseek-ai#304 metadata/schedule API mismatch

The upstream deepseek-ai#304 refactor reshapes the attention/MQA layer; conflicts
with the pre-merge nv_dev code are resolved by taking the post-deepseek-ai#304
version. Also drops the redundant __syncthreads() in
sm100_bmk_bnk_mn.cuh that was added by nv_dev's 8e866e4; deepseek-ai#304 removed
the equivalent sync upstream (replaced with NamedBarrier.sync in
sibling SM100 kernels) so keeping the nv_dev version would diverge
from upstream's final sync design.
Barry-Delaney pushed a commit to Barry-Delaney/DeepGEMM that referenced this pull request Apr 22, 2026
Extend the SM90 FP8 paged MQA logits kernel and host wrapper to cover
block_kv=32 and next_n=4 on post-deepseek-ai#304 code.

Kernel (deep_gemm/include/deep_gemm/impls/sm90_fp8_paged_mqa_logits.cuh):
  - Keep MMA tile at kComputeBlockKV=64 and let block_kv vary (32 or
    64). Each WGMMA iteration covers kNumBlocksPerMMA = 64 / block_kv
    physical blocks with one TMA copy each.
  - Add kNumKVMulticast template parameter for next_n=4; compute
    kNextNPerCTA = kNextN / kNumKVMulticast and split Q rows across
    two CTAs in a cluster via cta_rank_in_cluster, keeping smem and
    register budgets within SM90 limits (WGMMA N=256 otherwise
    exceeds kNumMathRegisters=104).
  - Guard the block_table load per physical block so BLOCK_KV=32 with
    an odd physical block count cannot read past the natural row end.
  - Use cute::cluster_id_in_grid().x for the scheduler index.

Post-deepseek-ai#304 metadata/schedule integration
(deep_gemm/include/deep_gemm/scheduler/paged_mqa_logits.cuh,
 csrc/apis/attention.hpp):
  - Pipe num_next_n_atoms as a runtime parameter so SM90 cluster
    multicast (no atomization) passes 1 while SM100 time-atomization
    keeps 2 for next_n=4, matching the kernel's scheduling.
  - Relax the schedule_meta size assertion to
    num_sms / num_kv_multicast + 1 to match the cluster grid.
  - Normalize the public get_paged_mqa_logits_metadata helper so
    callers pass the natural get_num_sms() and physical block_kv; the
    helper now computes num_clusters internally, sizes the schedule
    as {num_clusters + 1, 2}, and uses the compute-block size on SM90
    regardless of the physical block_kv the caller supplies.
  - Extract get_paged_mqa_logits_num_kv_multicast helper and add a
    num_sms % num_kv_multicast == 0 host-side guard.

Host wrapper (csrc/jit_kernels/impls/smxx_fp8_fp4_paged_mqa_logits.hpp):
  - Select num_kv_multicast = 2 when SM90+next_n=4 else 1. Build
    Q/weights TMA descriptors with next_n_per_cta * num_heads box,
    size smem from next_n_per_cta, launch with
    cluster_dim=num_kv_multicast.
  - Emit the extra kNumKVMulticast template arg only for SM90; SM100
    instantiation is unchanged.

Other:
  - Add utils::shfl_sync helper in common/utils.cuh.
  - tests/test_attention.py: extend enumerate_paged_mqa_logits to
    cover block_kv=32 and next_n=4 on SM90.
Barry-Delaney added a commit to Barry-Delaney/DeepGEMM that referenced this pull request Apr 22, 2026
The post-deepseek-ai#304 metadata kernel hard-codes num_next_n_atoms from next_n,
which assumes the SM100 time-atomization scheme. On SM90 next_n=4 the
cluster-multicast design processes one full q per cluster with no
atomization, so the hard-coded value over-schedules by 2x; when the
kernel scheduler (kNumNextNAtoms=1 in the SM90 template) decodes those
entries it steps past context_lens and triggers an illegal memory access
in sm90_fp8_paged_mqa_logits.

Fixes:
  - Pipe num_next_n_atoms as a runtime parameter into the metadata
    kernel (deep_gemm/include/deep_gemm/scheduler/paged_mqa_logits.cuh)
    and let the host decide. attention.hpp computes it based on arch +
    next_n + num_kv_multicast so SM90 multicast passes 1 while SM100
    keeps 2 for next_n=4, matching the respective kernel template.
  - Relax the schedule_meta size assertion in attention.hpp to
    'num_sms / num_kv_multicast + 1' so the check matches the cluster
    grid rather than the SM grid for SM90 next_n=4.
  - Relax the block_kv assertion in get_paged_mqa_logits_metadata to
    accept block_kv in {32, 64} and internally use the compute block
    size (64) on SM90, so callers can pass the physical block_kv the
    matching fp8_fp4_paged_mqa_logits call uses.
Barry-Delaney added a commit to Barry-Delaney/DeepGEMM that referenced this pull request Apr 22, 2026
Brings in upstream commits d30fc36 (TMEM alloc/dealloc sync fix, deepseek-ai#292)
and 7f2a703 (public release 26/04: Mega MoE, FP4 indexer, and related
refactors, deepseek-ai#304), plus three NV-side commits:

  - Enable bf16 bias output for FP8 GEMM on SM100 (xueweil)
  - SM90 paged MQA logits: support kv_block=32 and next_n=4 (Ray)
  - Fix SM90 paged MQA logits for post-deepseek-ai#304 metadata/schedule API mismatch

The upstream deepseek-ai#304 refactor reshapes the attention/MQA layer; conflicts
with the pre-merge nv_dev code are resolved by taking the post-deepseek-ai#304
version. Also drops the redundant __syncthreads() in
sm100_bmk_bnk_mn.cuh that was added by nv_dev's 8e866e4; deepseek-ai#304 removed
the equivalent sync upstream (replaced with NamedBarrier.sync in
sibling SM100 kernels) so keeping the nv_dev version would diverge
from upstream's final sync design.
@qinqinwo

Copy link
Copy Markdown

请问有支持其他卡型的计划吗?比如H100。

@yy16432

yy16432 commented May 1, 2026

Copy link
Copy Markdown

Any plan to add support to RTX PRO 6000(SM 120)? Thanks!

zyongye added a commit to zyongye/DeepGEMM that referenced this pull request Sep 14, 2026
Extend the SM90 FP8 paged MQA logits kernel and host wrapper to cover
block_kv=32 and next_n=4 on post-deepseek-ai#304 code.

Kernel (deep_gemm/include/deep_gemm/impls/sm90_fp8_paged_mqa_logits.cuh):
  - Keep MMA tile at kComputeBlockKV=64 and let block_kv vary (32 or
    64). Each WGMMA iteration covers kNumBlocksPerMMA = 64 / block_kv
    physical blocks with one TMA copy each.
  - Load block_table entries with a loop of scalar __ldg's so
    block_table_stride no longer has to be a multiple of
    kNumBlocksPerMMA.
  - Add kNumKVMulticast template parameter for next_n=4; compute
    kNextNPerCTA = kNextN / kNumKVMulticast and split Q rows across
    two CTAs in a cluster via cta_rank_in_cluster, keeping smem and
    register budgets within SM90 limits (WGMMA N=256 otherwise
    exceeds kNumMathRegisters=104).
  - Use cute::cluster_id_in_grid().x for the scheduler index.

Host wrapper (csrc/jit_kernels/impls/smxx_fp8_fp4_paged_mqa_logits.hpp):
  - Select num_kv_multicast = 2 when SM90+next_n=4 else 1. Build
    Q/weights TMA descriptors with next_n_per_cta * num_heads box,
    size smem from next_n_per_cta, launch with
    cluster_dim=num_kv_multicast.
  - Emit the extra kNumKVMulticast template arg only for SM90; SM100
    instantiation is unchanged.

Other:
  - Add utils::shfl_sync helper in common/utils.cuh.
  - tests/test_attention.py: extend enumerate_paged_mqa_logits to
    cover block_kv=32 and next_n=4 on SM90; call
    get_paged_mqa_logits_metadata with num_clusters = num_sms /
    num_kv_multicast; initialize block_table with torch.zeros so
    unused slots hold a valid block id.
(cherry picked from commit 12cb7c0)
Signed-off-by: Yongye Zhu <zyy1102000@gmail.com>
zyongye added a commit to vllm-project/DeepGEMM that referenced this pull request Sep 14, 2026
Extend the SM90 FP8 paged MQA logits kernel and host wrapper to cover
block_kv=32 and next_n=4 on post-deepseek-ai#304 code.

Kernel (deep_gemm/include/deep_gemm/impls/sm90_fp8_paged_mqa_logits.cuh):
  - Keep MMA tile at kComputeBlockKV=64 and let block_kv vary (32 or
    64). Each WGMMA iteration covers kNumBlocksPerMMA = 64 / block_kv
    physical blocks with one TMA copy each.
  - Load block_table entries with a loop of scalar __ldg's so
    block_table_stride no longer has to be a multiple of
    kNumBlocksPerMMA.
  - Add kNumKVMulticast template parameter for next_n=4; compute
    kNextNPerCTA = kNextN / kNumKVMulticast and split Q rows across
    two CTAs in a cluster via cta_rank_in_cluster, keeping smem and
    register budgets within SM90 limits (WGMMA N=256 otherwise
    exceeds kNumMathRegisters=104).
  - Use cute::cluster_id_in_grid().x for the scheduler index.

Host wrapper (csrc/jit_kernels/impls/smxx_fp8_fp4_paged_mqa_logits.hpp):
  - Select num_kv_multicast = 2 when SM90+next_n=4 else 1. Build
    Q/weights TMA descriptors with next_n_per_cta * num_heads box,
    size smem from next_n_per_cta, launch with
    cluster_dim=num_kv_multicast.
  - Emit the extra kNumKVMulticast template arg only for SM90; SM100
    instantiation is unchanged.

Other:
  - Add utils::shfl_sync helper in common/utils.cuh.
  - tests/test_attention.py: extend enumerate_paged_mqa_logits to
    cover block_kv=32 and next_n=4 on SM90; call
    get_paged_mqa_logits_metadata with num_clusters = num_sms /
    num_kv_multicast; initialize block_table with torch.zeros so
    unused slots hold a valid block id.
(cherry picked from commit 12cb7c0)

Signed-off-by: Yongye Zhu <zyy1102000@gmail.com>
cleonard530 pushed a commit to cleonard530/DeepGEMM that referenced this pull request Sep 24, 2026
Extend the SM90 FP8 paged MQA logits kernel and host wrapper to cover
block_kv=32 and next_n=4 on post-deepseek-ai#304 code.

Kernel (deep_gemm/include/deep_gemm/impls/sm90_fp8_paged_mqa_logits.cuh):
  - Keep MMA tile at kComputeBlockKV=64 and let block_kv vary (32 or
    64). Each WGMMA iteration covers kNumBlocksPerMMA = 64 / block_kv
    physical blocks with one TMA copy each.
  - Load block_table entries with a loop of scalar __ldg's so
    block_table_stride no longer has to be a multiple of
    kNumBlocksPerMMA.
  - Add kNumKVMulticast template parameter for next_n=4; compute
    kNextNPerCTA = kNextN / kNumKVMulticast and split Q rows across
    two CTAs in a cluster via cta_rank_in_cluster, keeping smem and
    register budgets within SM90 limits (WGMMA N=256 otherwise
    exceeds kNumMathRegisters=104).
  - Use cute::cluster_id_in_grid().x for the scheduler index.

Host wrapper (csrc/jit_kernels/impls/smxx_fp8_fp4_paged_mqa_logits.hpp):
  - Select num_kv_multicast = 2 when SM90+next_n=4 else 1. Build
    Q/weights TMA descriptors with next_n_per_cta * num_heads box,
    size smem from next_n_per_cta, launch with
    cluster_dim=num_kv_multicast.
  - Emit the extra kNumKVMulticast template arg only for SM90; SM100
    instantiation is unchanged.

Other:
  - Add utils::shfl_sync helper in common/utils.cuh.
  - tests/test_attention.py: extend enumerate_paged_mqa_logits to
    cover block_kv=32 and next_n=4 on SM90; call
    get_paged_mqa_logits_metadata with num_clusters = num_sms /
    num_kv_multicast; initialize block_table with torch.zeros so
    unused slots hold a valid block id.
(cherry picked from commit 12cb7c0)

Signed-off-by: Yongye Zhu <zyy1102000@gmail.com>
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.

7 participants