deepseek-ai / deepseek-ai/DeepGEMM

[QUESTION] dsv32 mqa_logits kernel not considering causal masking?

Open
#209 1 comment 1 reaction 0 assignees View on GitHub
Dominant language
Cuda
Stars
7.8k
Forks
1.3k
Avg merge
3d 7h
Merged PRs (30d)
3

Description

Thanks for dsv32 great work!

By analysis the `fp8_mqa_logits` and `fp8_paged_mqa_logits` function, looks like after the q@k, we don't consider causal masking before topk?

I know the q/k feed to mqa_logits kernel is different from the q/k feed to MLA attention kernel, but we use the output of the mqa_logits kernel (and topk 2048) as indexer into the real MLA's kvcache, hence during the real attention computation we need consider causal, in prefill or decode(MTP) case.

vLLM [prefill dispatch](https://github.com/vllm-project/vllm/blob/08d26a1b7edc200d8d117491eac3e28c0428e571/vllm/model_executor/models/deepseek_v2.py#L654C26-L654C39) using `torch.ops._C.top_k_per_row` seems not considering causal, [decode dispatch](https://github.com/vllm-project/vllm/blob/08d26a1b7edc200d8d117491eac3e28c0428e571/vllm/model_executor/models/deepseek_v2.py#L708) looks like considered causal in MTP case after the logits kernel, before topk.

not sure if it is suppose to let the framework side to do causal before topk, or actually causal is not important during the indexer kernel?

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.