Skip to content

Commit 101e6af

Browse files
committed
issue/1565 fix(nvidia): complete runtime support for Qwen MTP
Map E4M3 and BOOL through the existing ATen adaptor and preserve the caller's CUDA device across NCCL communicator destruction. Reuse the existing paged Prefill warp kernel for NVIDIA head size 256, without changing other vendors' default dispatch. Extend existing multi-page/long-context coverage and add finite FP8/mask cast checks. Validation: fresh SM86 build, 88 paged Prefill cases, 2 cast tests, and TP2 communicator teardown from both caller devices. Closes #1565
1 parent 1ab85ef commit 101e6af

6 files changed

Lines changed: 76 additions & 3 deletions

File tree

‎include/infinicore/adaptor/aten_adaptor.hpp‎

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -45,6 +45,8 @@ inline at::ScalarType to_at_dtype(DataType dtype) {
4545
return at::kHalf;
4646
case DataType::BF16:
4747
return at::kBFloat16;
48+
case DataType::F8:
49+
return at::kFloat8_e4m3fn;
4850
case DataType::I8:
4951
return at::kChar;
5052
case DataType::U8:
@@ -53,6 +55,8 @@ inline at::ScalarType to_at_dtype(DataType dtype) {
5355
return at::kInt;
5456
case DataType::I64:
5557
return at::kLong;
58+
case DataType::BOOL:
59+
return at::kBool;
5660
default:
5761
throw std::runtime_error("Unsupported dtype for ATen");
5862
}

‎src/infiniccl/cuda/infiniccl_cuda.cu‎

Lines changed: 7 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -119,7 +119,13 @@ infiniStatus_t commInitRank(
119119
}
120120

121121
infiniStatus_t commDestroy(infinicclComm_t comm) {
122-
CHECK_NCCL(ncclCommDestroy(getNcclComm(comm)));
122+
// NCCL teardown may activate the communicator's device. Preserve the
123+
// caller's device so its existing stream and runtime context stay valid.
124+
int previous_device;
125+
CHECK_INTERNAL(cudaGetDevice(&previous_device), cudaSuccess);
126+
const auto status = ncclCommDestroy(getNcclComm(comm));
127+
CHECK_INTERNAL(cudaSetDevice(previous_device), cudaSuccess);
128+
CHECK_NCCL(status);
123129
delete comm;
124130
return INFINI_STATUS_SUCCESS;
125131
}

‎src/infiniop/ops/paged_attention_prefill/cuda/kernel_v2.cuh‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -280,7 +280,7 @@ __global__ void PagedAttentionPrefillWarpGlobalKernel(
280280
ptrdiff_t o_head_stride) {
281281

282282
constexpr int kWarpSize = 32;
283-
static_assert(HEAD_SIZE == 64 || HEAD_SIZE == 128 || HEAD_SIZE == 192, "Only head_size 64/128/192 supported in v0.4.");
283+
static_assert(HEAD_SIZE == 64 || HEAD_SIZE == 128 || HEAD_SIZE == 192 || HEAD_SIZE == 256, "Unsupported head_size.");
284284
static_assert(HEAD_SIZE % kWarpSize == 0, "HEAD_SIZE must be divisible by 32.");
285285
constexpr int DIMS_PER_THREAD = HEAD_SIZE / kWarpSize;
286286

‎src/infiniop/ops/paged_attention_prefill/nvidia/paged_attention_prefill_nvidia.cu‎

Lines changed: 15 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -27,7 +27,11 @@ inline const char *default_prefill_kernel(const PagedAttentionPrefillInfo &info)
2727
return "ref";
2828
}
2929
if (info.head_size == 256) {
30+
#if defined(ENABLE_NVIDIA_API)
31+
return "warp";
32+
#else
3033
return "ref";
34+
#endif
3135
}
3236
// Iluvatar/Hygon: use warp for the non-MLA shapes where it is the stable path.
3337
#if defined(ENABLE_ILUVATAR_API) || defined(ENABLE_HYGON_API)
@@ -1020,6 +1024,17 @@ infiniStatus_t launch_prefill_warp(
10201024
v_batch_stride, v_row_stride, v_head_stride,
10211025
o_stride, o_head_stride);
10221026
return INFINI_STATUS_SUCCESS;
1027+
case 256:
1028+
op::paged_attention_prefill::cuda::PagedAttentionPrefillWarpGlobalKernel<Tindex, Tdata, 256>
1029+
<<<grid, block, 0, stream>>>(
1030+
out, q, k_cache, v_cache, block_tables, total_kv_lens, cu_seqlens_q, alibi_slopes,
1031+
num_heads, num_seqs, num_kv_heads, total_q_tokens, scale, max_num_blocks_per_seq,
1032+
page_block_size, block_table_batch_stride,
1033+
q_stride, q_head_stride,
1034+
k_batch_stride, k_row_stride, k_head_stride,
1035+
v_batch_stride, v_row_stride, v_head_stride,
1036+
o_stride, o_head_stride);
1037+
return INFINI_STATUS_SUCCESS;
10231038
default:
10241039
return INFINI_STATUS_BAD_TENSOR_SHAPE;
10251040
}

‎test/infinicore/ops/fp8_cast.py‎

Lines changed: 44 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,44 @@
1+
"""Check FP8 weights and boolean acceptance masks across the ATen cast bridge."""
2+
3+
import torch
4+
from infinicore.lib import _infinicore
5+
6+
import infinicore
7+
8+
9+
def test_fp8_cast():
10+
bits = torch.arange(256, dtype=torch.int16)
11+
bits = bits[(bits != 127) & (bits != 255)].to(torch.uint8)
12+
source = bits.view(torch.float8_e4m3fn).cuda()
13+
for dtype in (torch.float32, torch.bfloat16):
14+
output = torch.empty(source.shape, dtype=dtype, device="cuda")
15+
_infinicore.cast_(
16+
infinicore.from_torch(output)._underlying,
17+
infinicore.from_torch(source)._underlying,
18+
)
19+
infinicore.sync_device()
20+
torch.testing.assert_close(output, source.to(dtype), rtol=0, atol=0)
21+
22+
23+
def test_bool_cast():
24+
candidates = torch.tensor([3, 5, 7, 9], device="cuda")
25+
expected = torch.tensor([4, 5, 7, 8], device="cuda")
26+
source = infinicore.equal(
27+
infinicore.from_torch(candidates), infinicore.from_torch(expected)
28+
)
29+
for dtype in (torch.float32, torch.int64):
30+
output = torch.empty(source.shape, dtype=dtype, device="cuda")
31+
_infinicore.cast_(
32+
infinicore.from_torch(output)._underlying,
33+
source._underlying,
34+
)
35+
infinicore.sync_device()
36+
torch.testing.assert_close(
37+
output, (candidates == expected).to(dtype), rtol=0, atol=0
38+
)
39+
40+
41+
if __name__ == "__main__":
42+
test_fp8_cast()
43+
test_bool_cast()
44+
print("Finite E4M3 and boolean acceptance mask casts passed")

‎test/infinicore/ops/paged_attention_prefill.py‎

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -26,6 +26,8 @@
2626
(1, 24, 4, 256, 8, 8, 1),
2727
(1, 12, 2, 256, 8, 8, 1),
2828
(1, 6, 1, 256, 8, 8, 1),
29+
(2, 12, 2, 256, 64, 128, 2),
30+
(1, 24, 4, 256, 64, 1024, 1),
2931
# New DeepSeek MLA wrapper case: verifies prefill supports q/k head
3032
# size 576 with value head size 512.
3133
(1, 16, 1, 576, 8, 8, 1, 512),
@@ -94,7 +96,9 @@ def parse_test_cases():
9496
value_size,
9597
) = case
9698
scale = head_size**-0.5
97-
num_blocks = 8192
99+
num_blocks = num_seqs * (
100+
(max_step_len * num_rounds + block_size - 1) // block_size
101+
)
98102
manager = SimpleCacheManager(num_blocks, block_size)
99103
kv_lens = torch.zeros(num_seqs, dtype=torch.int32)
100104

0 commit comments

Comments
 (0)