deepseek-ai / deepseek-ai/FlashMLA

[Bug/Correctness] Hardcoded device_id=0 + missing CUDAGuard can break multi-GPU correctness (wrong hw_info / stream mismatch)

Open
#158 2 comments 0 reactions 0 assignees View on GitHub
Dominant language
C++
Stars
12.9k
Forks
1.2k
Avg merge
4h 20m
Merged PRs (30d)
2

Description

Hi maintainers,

While reviewing the FMHA forward runner integration, I noticed two correctness issues that can break execution on multi-GPU setups (and can also create subtle stream/device mismatches):

1) Hardcoded GPU selection (device_id = 0)

In run_fmha_fwd, the hardware info is pinned to GPU0:

cutlass::KernelHardwareInfo hw_info;
hw_info.device_id = 0;
hw_info.sm_count =
cutlass::KernelHardwareInfo::query_device_multiprocessor_count(hw_info.device_id);

If tensors (q/k/v/o/lse) live on a non-zero device, this will query the wrong SM count and may also lead to launching with incorrect hardware assumptions.

2) Missing CUDA device guard (at::cuda::CUDAGuard) and stream/device alignment

The code uses at::cuda::getCurrentCUDAStream() at the end:

CUTLASS_CHECK(op.run(at::cuda::getCurrentCUDAStream()));

but does not guard/set the current device to match q.device() (or any input tensor). In multi-GPU scenarios, the “current device” may differ from the tensor’s device, leading to:

wrong stream/device being used

incorrect hw_info.device_id / sm_count

potential launch failures or silent misbehavior

Suggested fix

Use a device guard based on an input tensor (e.g., q) and set hw_info.device_id accordingly:

#include

at::cuda::CUDAGuard device_guard(q.device());
const int dev = at::cuda::current_device();

cutlass::KernelHardwareInfo hw_info;
hw_info.device_id = dev;
hw_info.sm_count =
cutlass::KernelHardwareInfo::query_device_multiprocessor_count(dev);

CUTLASS_CHECK(op.run(at::cuda::getCurrentCUDAStream()));

This ensures:

correct device is active

stream matches the tensor device context

hardware info queries the right GPU

Why this matters

Even if most users run single-GPU, multi-GPU is common in training/inference servers. Hardcoding GPU0 + missing guards can produce correctness issues that are hard to diagnose (especially when the failure is not immediate).

If you'd like, I can provide a small repro snippet that places q/k/v/o on cuda:1 and shows the mismatch.

Thanks!

Contributor guide

No contributing guide indexed for this repository

Research direction

Start at the FMHA forward runner's run_fmha_fwd entry point and inspect how q.device(), hardware information, and the current CUDA stream are selected. Done means non-zero-device tensors use matching device hardware information and stream context without the hardcoded GPU0 assumption; the payload mentions no file or test to run.

Written by the indexing model from the issue text.

Assessment

Tech stack
cpp
Domain
backend, performance
Issue type
Bug
Difficulty
4/5
Estimated time
3-5 days
Activity status
Stale
Clarity
Mostly clear
Newbie friendliness
35/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.