deepseek-ai / deepseek-ai/DeepGEMM

Feature Request: Support sm_120 ( 5090 and blackwell 6000 pro )

Open
#236 28 comments 9 reactions 0 assignees View on GitHub
Dominant language
Cuda
Stars
7.8k
Forks
1.3k
Avg merge
3d 7h
Merged PRs (30d)
3

Description

I can run the following successfully with my blackwell 6000 pro:

```bash
./develop.sh
./install.sh
```

This works too:

```bash
python ./tests/test_layout.py
```

However, when I try to run a custom kernel benchmarking script:

**bench_fp8_paged_mqa_logits.py**

```py
#!/usr/bin/env python3
import os
import time

import torch
import deep_gemm

# ---------------------------------------------
# JIT logging / debug knobs
# ---------------------------------------------
# Print the full nvcc/nvrtc compile command once when the kernel is JIT-compiled.
os.environ.setdefault("DG_JIT_PRINT_COMPILER_COMMAND", "1")
# Optionally get more debug info:
# os.environ.setdefault("DG_JIT_DEBUG", "1")
# os.environ.setdefault("DG_JIT_PTXAS_VERBOSE", "1")
# os.environ.setdefault("DG_JIT_PTXAS_CHECK", "1")

device = "cuda"

def ceil_div(a: int, b: int) -> int:
return (a + b - 1) // b

@torch.no_grad()
def bench_fp8_paged_mqa_logits(
batch_size: int = 64,
next_n: int = 2,
heads: int = 64,
index_dim: int = 128,
avg_kv: int = 8192,
blocksize: int = 64,
max_model_len: int = 111 * 1000,
is_context_lens_2d: bool = False,
iters: int = 50,
):
torch.manual_seed(0)

# Shapes follow tests/test_attention.py
q = torch.randn(
(batch_size, next_n, heads, index_dim),
device=device,
dtype=torch.bfloat16,
)
kv_cache = torch.randn(
(max_model_len * 3, blocksize, 1, index_dim),
device=device,
dtype=torch.bfloat16,
)
weights = torch.randn(
(batch_size * next_n, heads),
device=device,
dtype=torch.float32,
)

# FP8 casts (same helpers as tests, but simple version here).
q_fp8 = q.to(torch.float8_e4m3fn)
# Quantize kv_cache per-token (very simple scale=1 for benchmarking only)
kv_cache_fp8 = kv_cache.to(torch.float8_e4m3fn)

# Random context lengths around avg_kv
context_lens = torch.randint(
int(0.7 * avg_kv),
int(1.3 * avg_kv),
(batch_size,),
device=device,
dtype=torch.int32,
)
context_lens_list = context_lens.tolist()
max_block_len = ceil_div(max(context_lens_list), blocksize) * blocksize

# Block tables
num_blocks_total = kv_cache.shape[0]
block_tables = torch.zeros(
(batch_size, max_block_len),
device=device,
dtype=torch.int32,
)
counter = 0
block_idx_pool = torch.randperm(
num_blocks_total, device=device, dtype=torch.int32
)
for i in range(batch_size):
nb = ceil_div(context_lens_list[i], blocksize)
block_tables[i, :nb] = block_idx_pool[counter : counter + nb]
counter += nb

if is_context_lens_2d:
# Build per-(batch,next_n) context lengths as in tests
context_lens_2d = (
(context_lens.unsqueeze(1) + 1)
* torch.rand(batch_size, next_n, device=device)
).to(torch.int32)
context_lens_2d[:, next_n - 1] = context_lens
ctx_for_meta = context_lens_2d
else:
ctx_for_meta = context_lens

# JIT metadata (this triggers paged_mqa_logits metadata specialization)
schedule_metadata = deep_gemm.get_paged_mqa_logits_metadata(
ctx_for_meta, blocksize, deep_gemm.get_num_sms()
)

# One warmup to trigger JIT & capture compile cmd
if is_context_lens_2d:
logits = deep_gemm.fp8_paged_mqa_logits(
q_fp8,
kv_cache_fp8,
weights,
ctx_for_meta,
block_tables,
schedule_metadata,
max_model_len,
clean_logits=False,
)
else:
logits = deep_gemm.fp8_paged_mqa_logits(
q_fp8,
kv_cache_fp8,
weights,
ctx_for_meta,
block_tables,
schedule_metadata,
max_model_len,
clean_logits=True,
)

torch.cuda.synchronize()

# Benchmark loop
times = []
for _ in range(iters):
t0 = time.perf_counter()
if is_context_lens_2d:
logits = deep_gemm.fp8_paged_mqa_logits(
q_fp8,
kv_cache_fp8,
weights,
ctx_for_meta,
block_tables,
schedule_metadata,
max_model_len,
clean_logits=False,
)
else:
logits = deep_gemm.fp8_paged_mqa_logits(
q_fp8,
kv_cache_fp8,
weights,
ctx_for_meta,
block_tables,
schedule_metadata,
max_model_len,
clean_logits=True,
)
torch.cuda.synchronize()
t1 = time.perf_counter()
times.append(t1 - t0)

avg_us = 1e6 * sum(times) / len(times)

# Rough TFLOPS estimate, same idea as test_paged_mqa_logits
sum_lens = sum(context_lens.to(torch.int64))
tflops = 2.0 * sum_lens * next_n * heads * index_dim / 1e12 / (avg_us * 1e-6)

print(
f"fp8_paged_mqa_logits benchmark:\n"
f" BSZ={batch_size}, NextN={next_n}, H={heads}, D={index_dim}, avg_kv={avg_kv}\n"
f" iters={iters}, avg time = {avg_us:.1f} us, approx {tflops:.1f} TFLOPS"
)

if __name__ == "__main__":
# You can toggle these to exercise both metadata modes.
bench_fp8_paged_mqa_logits(is_context_lens_2d=False)
bench_fp8_paged_mqa_logits(is_context_lens_2d=True)
```

```bash
DG_JIT_PRINT_COMPILER_COMMAND=1 DG_JIT_PTXAS_VERBOSE=1 python bench_fp8_paged_mqa_logits.py
```

... first it fails with `Unsupported architecture`, then if I add `arch_major == 12`
it fails with this:

```
$ DG_JIT_PRINT_COMPILER_COMMAND=1 DG_JIT_PTXAS_VERBOSE=1 python bench_fp8_paged_mqa_logits.py
Warning: please use at least NVCC 12.9 for the best DeepGEMM performance
Running NVCC command: /usr/local/cuda/bin/nvcc /home/jesse/.deep_gemm/cache/kernel.smxx_paged_mqa_logits_metadata.a8d04824495f4b393f3e37abf95794af/kernel.cu -o /home/jesse/.deep_gemm/tmp/775171-79e89b64-d5c1d7aa-f29af1ce -std=c++20 --diag-suppress=39,161,174,177,186,940 --ptxas-options=--register-usage-level=10 --ptxas-options=--verbose,--warn-on-local-memory-usage -I/home/jesse/sandbox/DeepGEMM/deep_gemm/include --gpu-architecture=sm_120a --compiler-options=-fPIC,-O3,-fconcepts,-Wno-deprecated-declarations,-Wno-abi -cubin -O3 --expt-relaxed-constexpr --expt-extended-lambda
NVCC compilation failed: /home/jesse/.deep_gemm/cache/kernel.smxx_paged_mqa_logits_metadata.a8d04824495f4b393f3e37abf95794af/kernel.cu:2:10: fatal error: deep_gemm/impls/sm120_fp8_paged_mqa_logits.cuh: No such file or directory
2 | #include
| ^~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
compilation terminated.

Traceback (most recent call last):
File "/home/jesse/sandbox/DeepGEMM/bench_fp8_paged_mqa_logits.py", line 175, in
bench_fp8_paged_mqa_logits(is_context_lens_2d=False)
File "/data/conda-envs/vllm_cu128/lib/python3.11/site-packages/torch/utils/_contextlib.py", line 120, in decorate_context
return func(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^
File "/home/jesse/sandbox/DeepGEMM/bench_fp8_paged_mqa_logits.py", line 100, in bench_fp8_paged_mqa_logits
schedule_metadata = deep_gemm.get_paged_mqa_logits_metadata(
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
RuntimeError: Assertion error (csrc/apis/../jit_kernels/impls/../../jit/compiler.hpp:183): false and "NVCC compilation failed"
```

Please support `sm120`. I'd love to add support myself, but I'm no kernel wizard.

Contributor guide

No contributing guide indexed for this repository

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.