Repository navigation
Conversation
The SM90 `bmk,bnk->mn` kernel splits the `s * k` reduction across CTAs and adds the partial sums into D with float2 `atomicAdd`. The order of the additions changes between runs, so repeated identical calls return different results, even with `deep_gemm.use_deterministic_algorithms(True)`. In deterministic mode (and only when there is more than one split), let each split accumulate into its own slice of a zeroed `[num_splits, m, n]` FP32 workspace, so every element has a single writer, and sum the slices into D in a fixed order afterwards. The default mode is unchanged. The SM100 kernel reduces with TMA reduce-add in the same way and is left as is. Signed-off-by: RunguoLi <li19107254665@gmail.com>
| s, tensor_map_a, tensor_map_b, d.data_ptr<float>() | ||
| s, tensor_map_a, tensor_map_b, (use_split_workspace ? split_workspace : d).data_ptr<float>() | ||
| ); | ||
| if (use_split_workspace) |
There was a problem hiding this comment.
🔵 suggestion: split_workspace.sum(0) allocates a fresh [m, n] FP32 tensor on every deterministic call before the in-place add_. This path already pays for zeroing and an extra reduction, so it is acceptable, but preallocating the reduction buffer (or using torch::sum(..., out=...)) would remove one allocation and reduce allocator churn on the small shapes where the relative overhead is largest.
🤖 v4
|
|
||
| // Splits are reduced into D with atomic additions, whose order changes between runs. For deterministic | ||
| // algorithms, each split writes its own slice of a workspace instead, and the slices are summed in a fixed order. | ||
| const bool use_split_workspace = heuristics_runtime->get_deterministic_algorithms() and num_splits > 1; |
There was a problem hiding this comment.
🔵 suggestion: The FP32 workspace relies on d always being FP32, which holds today because the BF16 entry point allocates an FP32 workspace and recurses, but the coupling is implicit and the kernel immediately calls data_ptr<float>(). Consider documenting it with a brief comment or a DG_HOST_ASSERT(d.scalar_type() == torch::kFloat) so a future caller cannot silently pass a non-FP32 destination and corrupt the launch.
🤖 v4
🤖 ds-review-bot Code Reviewv6v5v4This MR fixes non-deterministic results in the SM90 Files reviewed: 3 📍 未定位到 diff 的评论🔵 suggestion |
Signed-off-by: RunguoLi <li19107254665@gmail.com>
|
@RayWang96 Could you review this when you have a chance? It makes the SM90 |
Problem
deep_gemm.einsum('bmk,bnk->mn', ...)returns a different result on every call for the same inputs, even withdeep_gemm.use_deterministic_algorithms(True).On SM90,
sm90_bmn_bnk_mn_gemmsplits thes * kreduction across CTAs (split_factor) and adds each CTA's partial sum into D withfloat2atomicAdd. The order of those additions changes between runs. With DeepGEMMmain(78b6900) on an H200, 50 identical calls give:use_deterministic_algorithms(True)The flag is already honored by the other einsum paths, the BF16 GEMM,
mega_mhcandmega_gate, but not by this one.Fix (SM90)
In deterministic mode, and only when there is more than one split:
[num_splits, m, n]FP32 workspace. Its size is at most aboutnum_sms * 128 * 128floats, sincenum_splits * num_mn_blocksis bounded by the SM count.kUseSplitWorkspacetemplate flag makes each split accumulate into its own slice. Every element then has a single writer, so the atomic additions no longer depend on ordering.split_workspace.sum(0).With the flag off, or with a single split, the generated kernel and launch are the same as before.
The SM100 kernel reduces the splits with TMA reduce-add in the same way, and I left it unchanged because I have no SM100 machine to test on. The same approach should apply there.
Testing (H200, CUDA 13.0, PyTorch 2.11.0+cu130)
test_bmk_bnk_mn_deterministic(SM90 only) fails onmainand passes with this change. It checks 10 identical calls for bitwise equality and compares them with the reference.use_deterministic_algorithms(True)python tests/test_einsum.pypasses.mainThe deterministic mode costs a roughly fixed 12–15 µs for the extra zeroing, reduction and addition. That is about 2x on the smallest problems and 2–11% on the larger ones.