deepseek-ai / deepseek-ai/TileKernels

[Correctness][Multi-GPU] MoE APIs allocate and launch on the current CUDA device instead of the input device

Open
#28 1 comment 0 reactions 0 assignees View on GitHub
Dominant language
Python
Stars
1.8k
Forks
162
PR merge metrics
No merged PRs in 30d

Description

## Summary

Several public MoE APIs allocate outputs using an unqualified `device="cuda"` instead of the device of their input tensors.

When an input tensor is on a non-current CUDA device, for example:

- input tensor: `cuda:1`
- current CUDA device: `cuda:0`

the outputs are allocated on `cuda:0`. Kernel compilation and launch also occur without an explicit device guard, which can result in program failure.

Affected APIs include at least:

- [`get_fused_mapping`](https://github.com/deepseek-ai/TileKernels/blob/main/tile_kernels/moe/get_fused_mapping_kernel.py)
- [`expand_to_fused`](https://github.com/deepseek-ai/TileKernels/blob/main/tile_kernels/moe/expand_to_fused_kernel.py)
- [`reduce_fused`](https://github.com/deepseek-ai/TileKernels/blob/main/tile_kernels/moe/reduce_fused_kernel.py)
- [`group_count`](https://github.com/deepseek-ai/TileKernels/blob/main/tile_kernels/moe/group_count_kernel.py)
- [`aux_fi`](https://github.com/deepseek-ai/TileKernels/blob/main/tile_kernels/moe/aux_fi_kernel.py)

For example, `get_fused_mapping()` currently allocates all output tensors with `device="cuda"`:

```python
pos_to_expert = torch.empty(..., device="cuda")
pos_to_token = torch.empty(..., device="cuda")
pos_to_token_topk = torch.empty(..., device="cuda")
token_topk_to_pos = torch.empty(..., device="cuda")
expert_start = torch.empty(..., device="cuda")
expert_end = torch.empty(..., device="cuda")
num_tokens_per_expert = torch.empty(..., device="cuda")
num_experts_per_sm = torch.empty(..., device="cuda")
```

Relevant code: [get_fused_mapping output allocation](https://github.com/deepseek-ai/TileKernels/blob/main/tile_kernels/moe/get_fused_mapping_kernel.py#L204-L229)

## Root cause

### 1. Unqualified CUDA allocations

In PyTorch, `device="cuda"` resolves to the current CUDA device rather than the device of an existing input tensor.

Therefore, the following is not device-safe: `out = torch.empty(shape, device="cuda")` when the input may be on a non-current device. The output should instead be allocated using the relevant input device: `out = torch.empty(shape, device=input.device)`

### 2. Missing device guards around compilation and launch

Even after allocating outputs on the correct device, JIT compilation and kernel launch must run with the input device active as current CUDA device. By code, this is as follow:

```python
device = input.device

with torch.cuda.device(device):
kernel = get_kernel(...)
kernel(input, output, ...)
```

Without this guard, device-specific compilation state and launch state can still be derived from the previously current device.

### 3. Device-property caches do not include the device index

The following functions are zero-argument `lru_cache` functions:

```python
@functools.lru_cache(maxsize=None)
def get_device_num_sms() -> int:
prop = torch.cuda.get_device_properties(torch.cuda.current_device())
return prop.multi_processor_count

@functools.lru_cache(maxsize=None)
def get_max_smem_per_sm() -> int:
prop = torch.cuda.get_device_properties(torch.cuda.current_device())
return prop.shared_memory_per_multiprocessor
```

Relevant code: [`config.py`](https://github.com/deepseek-ai/TileKernels/blob/main/tile_kernels/config.py#L7-L29)

The cache key is always the empty argument tuple. After the first call, switching to another CUDA device can still return the first device's SM count or shared-memory limit.

This is observable when the process uses GPUs with different properties.

## Minimal reproduction

### 1. Multi-GPU bug

The reproduction of bug requires at least two GPU devices. The following uses `group_count()` since it has a small input and a single output:

```python
import torch
from tile_kernels.moe import group_count

if torch.cuda.device_count() < 2:
raise RuntimeError("This reproduction requires two visible CUDA devices")

# Keep cuda:0 as the current device.
torch.cuda.set_device(0)

# Allocate the input on cuda:1 without leaving cuda:1 current.
with torch.cuda.device(1):
group_idx = torch.tensor(
[[0, 1], [1, 0]],
dtype=torch.int64,
device="cuda:1",
).contiguous()

assert torch.cuda.current_device() == 0
assert group_idx.device == torch.device("cuda:1")

out = group_count(group_idx, num_groups=2)

# Surface asynchronous launch errors.
for device_index in (0, 1):
torch.cuda.synchronize(device_index)

print("current device:", torch.cuda.current_device())
print("input device:", group_idx.device)
print("output device:", out.device)

assert out.device == group_idx.device
```

Run with synchronous CUDA error reporting:

```bash
CUDA_VISIBLE_DEVICES=0,1 CUDA_LAUNCH_BLOCKING=1 python repro.py
```

On the current implementation, the call will fail.

### 2. Device-property cache reproduction

This reproduction requires two visible devices with different SM counts:

```python
import torch
import tile_kernels.config as config

if torch.cuda.device_count() < 2:
raise RuntimeError("This reproduction requires two visible CUDA devices")

config.get_device_num_sms.cache_clear()

torch.cuda.set_device(0)
sm0 = config.get_device_num_sms()
actual0 = torch.cuda.get_device_properties(0).multi_processor_count

torch.cuda.set_device(1)
sm1 = config.get_device_num_sms()
actual1 = torch.cuda.get_device_properties(1).multi_processor_count

print("cuda:0:", sm0, actual0)
print("cuda:1:", sm1, actual1)

assert sm0 == actual0
assert sm1 == actual1
```

If the two devices have different SM counts, `sm1` incorrectly retains the value returned for device 0. On identical GPUs, the cache is still device-insensitive, but the problem is numerically hidden because both devices happen to have the same property value.

## Why existing tests did not catch this

Existing tests generally create inputs using an unqualified `device="cuda"` and run them on the current CUDA device. As a result, the input device, output device, compilation device, and launch device are normally all the same.

## Proposed fix

### 1. Output allocation and device guard

Derive the device from the primary input:

```python
device = input.device

if device.type != "cuda":
raise ValueError(f"input must be a CUDA tensor, got {device}")

with torch.cuda.device(device):
out = torch.empty(..., device=device)
kernel = get_kernel(...)
kernel(input, out, ...)
```

For APIs with multiple tensor inputs, validate their device relationship before entering the guard:

```python
if mapping.device != input.device:
raise ValueError(
f"input and mapping must be on the same device, "
f"got {input.device} and {mapping.device}"
)
```

Optional preallocated outputs should be validated rather than moved implicitly.

### 2. Device-aware property caches

Make the device index part of the cache key(Callers can resolve the relevant device index from their primary input):

```python
@functools.lru_cache(maxsize=None)
def get_device_num_sms(device_index: int) -> int:
return torch.cuda.get_device_properties(
device_index
).multi_processor_count

@functools.lru_cache(maxsize=None)
def get_max_smem_per_sm(device_index: int) -> int:
return torch.cuda.get_device_properties(
device_index
).shared_memory_per_multiprocessor
```

If `set_num_sms()` is intended to override device properties, its currently process-global `_num_sms` state may also need to be made device-specific or explicitly documented as a global override.

## Suggested test coverage

Most of the fix can be tested without specialized multi-GPU hardware:

1. Mock output allocation and verify that every automatically created tensor receives `input.device`.
2. Mock CUDA device contexts and verify that kernel construction and invocation occur under the input device.
3. Mock two devices with different SM/shared-memory properties and verify independent cached results.
4. Run all existing single-GPU correctness tests.

An optional two-GPU integration test should additionally verify:

```python
torch.cuda.set_device(0)
input = make_input(device="cuda:1")
output = api(input)
assert output.device == input.device
```

A heterogeneous pair with different SM is only required to reproduce the device-property cache issue on real hardware; deterministic mocked tests can cover this behavior in regular CI.

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.