[BUG] Grouped MXFP8 quantization is not concurrency safe with multiple streams
维护者通常 2 天内回复
评估
- 难度
- 4/5
- 预计耗时
- 3-5 天
- 新手友好度
- 25/100
- Issue 类型
- 缺陷
- 描述清晰度
- 描述清楚
- 活跃度
- 停滞
调研方向
Start by reading transformer_engine/common/cast/core/grouped_tma.cuh, especially the global g_tensor_maps storage described in the issue, and review linked pull request #3632 to see what work is already underway. Run the supplied two-stream reproduction on a B200 and compare its outputs with the serial references; done means overlapping quantizations produce the same outputs as serial execution.
由索引模型根据 Issue 内容生成。
描述
Describe the bug
In transformer_engine/common/cast/core/grouped_tma.cuh, TensorMapStorage is a global static variable
struct alignas(128) TensorMapStorage {
alignas(128) CUtensorMap input[MAX_SUPPORTED_TENSOR_DESCRIPTORS];
alignas(128) CUtensorMap act_input[MAX_SUPPORTED_TENSOR_DESCRIPTORS];
alignas(128) CUtensorMap output_rowwise[MAX_SUPPORTED_TENSOR_DESCRIPTORS];
alignas(128) CUtensorMap output_colwise[MAX_SUPPORTED_TENSOR_DESCRIPTORS];
size_t rows[MAX_SUPPORTED_TENSOR_DESCRIPTORS];
size_t cols[MAX_SUPPORTED_TENSOR_DESCRIPTORS];
size_t offsets[MAX_SUPPORTED_TENSOR_DESCRIPTORS];
};
// Internal linkage avoids device-link ODR issues when this header is included by multiple .cu TUs.
static __device__ TensorMapStorage g_tensor_maps;
which would be not concurrency safe if multiple kernels are launched and their execution overlaps
Steps/Code to reproduce bug
import torch
import transformer_engine.pytorch as te
import transformer_engine_torch as tex
NUM_GROUPS, ROWS, COLS_PER_GROUP = 8, 8192, 2048
def group_quantize(x):
quantizer = te.MXFP8Quantizer(te.DType.kFloat8E4M3, rowwise=True, columnwise=False)
last_dims = torch.full((NUM_GROUPS,), COLS_PER_GROUP, dtype=torch.int64, device="cuda")
out = tex.group_quantize(x, quantizer, NUM_GROUPS, None, last_dims)
return out.rowwise_data, out.scale_inv
x_a = torch.randn(ROWS, NUM_GROUPS * COLS_PER_GROUP, dtype=torch.bfloat16, device="cuda")
x_b = torch.randn_like(x_a) * 1000
ref_a = [t.clone() for t in group_quantize(x_a)]
ref_b = [t.clone() for t in group_quantize(x_b)]
s1, s2 = torch.cuda.Stream(), torch.cuda.Stream()
failures, iters = 0, 100
for _ in range(iters):
torch.cuda.synchronize()
with torch.cuda.stream(s1):
out_a = group_quantize(x_a)
with torch.cuda.stream(s2):
out_b = group_quantize(x_b)
torch.cuda.synchronize()
if any(not torch.equal(o, r) for o, r in zip(out_a + out_b, ref_a + ref_b)):
failures += 1
# Poison freed outputs: the allocator reuses these blocks next iteration, and stale correct
# bytes would hide tiles that were never written.
for t in out_a + out_b:
t.view(torch.uint8).fill_(0xA5)
print(f"{failures}/{iters} iterations produced wrong output")
Expected behavior
This should output the same result as when they are quantized in serial, but now it's 90/100 iterations produced wrong output
Environment overview (please complete the following information)
Not related
Environment details
Not related
Device details
- B200
Additional context
N/A
- 主要语言
- Python
- 星标
- 3.6k
- 派生
- 851
- 平均合并
- 4 天 15 小时
- 30 天内合并 PR
- 51
环境准备
- 没有 Dockerfile 或 Docker Compose 文件
- 有 Pull Request 模板
- 阅读贡献指南
从这里开始
- 先读完整个 Issue,再读项目的贡献指南。
- 在 Issue 下留言说明你要接手 —— 这能避免两个人做同样的事。
- Fork 仓库,在一个分支上完成修改。
- 提交 Pull Request,并在描述里引用这个 Issue 编号。
NVIDIA/TransformerEngine 的其他 Issue
-
[PyTorch] fp8_cs_quantize fake implementation returns a vector inverse scale instead of a scalar未关闭
难度 2/5 1-3 小时 新手友好度 82/100
NVIDIA/TransformerEngine#3636 ·
维护者通常 2 天内回复
-
[Bug] Backend selection picks FA3 for training with head_dim_qk=192 / v_head_dim=128, but FA3 backward cannot run it可能已有人在做 @yuweih205 于 31 天前认领。 未关闭attention
难度 2/5 1-3 小时 新手友好度 85/100
NVIDIA/TransformerEngine#3481 · 4 条评论 ·
维护者通常 2 天内回复
-
bug
难度 2/5 1-3 小时 新手友好度 68/100
NVIDIA/TransformerEngine#2189 · 7 条评论 · 5 个 reaction ·
维护者通常 2 天内回复
-
难度 4/5 3-5 天 新手友好度 50/100
NVIDIA/TransformerEngine#3645 ·
维护者通常 2 天内回复
-
enhancement
难度 5/5 一周以上 新手友好度 35/100
NVIDIA/TransformerEngine#3644 ·
维护者通常 2 天内回复
查看 NVIDIA/TransformerEngine 的全部 Issue
相似的 Issue
-
难度 1/5 1-3 小时 新手友好度 85/100
pytest-dev/pluggy#757 ·
维护者通常 1 天内回复
-
难度 1/5 1-3 小时 新手友好度 85/100
NousResearch/hermes-agent#134960 ·
维护者通常 1 天内回复
-
HTML backend: `<br>` leaks the internal sentinel U+E000 into list items, headings and captions可能已有人在做 @morten-lagabote 今天认领。 未关闭
难度 2/5 1-3 小时 新手友好度 67/100
docling-project/docling#4671 ·
维护者通常 1 天内回复
-
难度 2/5 1-3 小时 新手友好度 70/100
维护者通常 1 天内回复
-
good first issue hacktoberfest infra
难度 2/5 1-3 小时 新手友好度 78/100
维护者通常 1 天内回复