[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
- フォーク
- 844
- 平均マージ
- 4日 55分
- マージ済み PR(30日)
- 51
環境構築
- Dockerfile・Docker Compose ファイルなし
- プルリクエストのテンプレートあり
- コントリビューションガイドを読む
はじめの一歩
- issue を最後まで読み、次にプロジェクトのコントリビューションガイドを読みます。
- 着手することを issue にコメントします — 二人が同じ作業をするのを防げます。
- リポジトリをフォークし、ブランチを切って変更します。
- 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 が 30 日前に担当しました。 オープンattention
難易度 2/5 1〜3時間 初心者へのやさしさ 85/100
NVIDIA/TransformerEngine#3481 · コメント 4 件 ·
メンテナーはふだん 2 日以内に返信
-
bug
難易度 2/5 1〜3時間 初心者へのやさしさ 68/100
NVIDIA/TransformerEngine#2189 · コメント 7 件 · リアクション 5 件 ·
メンテナーはふだん 2 日以内に返信
-
bug
難易度 4/5 3〜5日 初心者へのやさしさ 54/100
NVIDIA/TransformerEngine#3640 · コメント 5 件 ·
メンテナーはふだん 2 日以内に返信
-
[bug] NVFP4 + `torch.compile`: errors with 3D input対応中かも @pggPL が 2 日前に担当しました。 オープンbug
難易度 2/5 1〜3時間 初心者へのやさしさ 25/100
NVIDIA/TransformerEngine#3626 ·
メンテナーはふだん 2 日以内に返信
NVIDIA/TransformerEngine の issue をすべて見る
似ている issue
-
難易度 1/5 1時間未満 初心者へのやさしさ 92/100
-
Harmony OPeNDAP SubSetter (HOSS) Geographic LARC_CLOUD PREFIRE_SAT2_AUX-SAT R01 production
難易度 2/5 1〜3時間 初心者へのやさしさ 68/100
nasa/harmony-autotester#245 ·
-
enhancement
難易度 2/5 1〜3時間 初心者へのやさしさ 68/100
Deltares/imod-python#1928 ·
-
難易度 1/5 1時間未満 初心者へのやさしさ 88/100
メンテナーはふだん 1 日以内に返信
-
feature
難易度 2/5 1〜3時間 初心者へのやさしさ 66/100