deepseek-ai / deepseek-ai/DeepGEMM
Potential bug in get_num_blocks calculation for specific grouped GEMM types
- Dominant language
- Cuda
- Stars
- 7.8k
- Forks
- 1.3k
- Avg merge
- 3d 7h
- Merged PRs (30d)
- 3
Description
Bug Description
在 csrc/jit_kernels/heuristics/common.hpp 的 get_best_config 函数中,计算总线程块数量的逻辑 get_num_blocks 似乎没有正确处理所有grouped GEMM 类型。
get_num_blocks 的实现为:
ceil_div(m, block_m) * ceil_div(n, block_n) * num_groups;
Problem
1. 对于 m_grouped_fp8_gemm_nt_contiguous 类型的 GEMM (调用入口在 csrc/python_api.cpp),传递给 get_best_config 的参数 m 是被所有 group 共享的总 M 维度。在这种情况下,get_num_blocks 的计算结果似乎被错误地放大了 num_groups 倍。
2. 而对于 fp8_m_grouped_gemm_nt_masked 等其他 grouped GEMM 类型,传递的 m 是单个 group 的 M 维度,此时 get_num_blocks 的计算是正确的。
这种不一致的处理方式可能导致在 MGroupedContiguous 场景下启动过多的线程块,造成资源浪费。
不知道我的理解是否正确,期待您的解答!
Contributor guide
No contributing guide indexed for this repository
Assessment
This issue has not been assessed yet.