Skip to content

[Fix] Honor deterministic algorithms in SM90 bmk,bnk->mn einsum - #457

Open
LiRunGuo wants to merge 2 commits into
deepseek-ai:mainfrom
LiRunGuo:deterministic-bmk-bnk-mn
Open

LiRunGuo wants to merge 2 commits into
deepseek-ai:mainfrom
LiRunGuo:deterministic-bmk-bnk-mn

Conversation

@LiRunGuo

Copy link
Copy Markdown

Problem

deep_gemm.einsum('bmk,bnk->mn', ...) returns a different result on every call for the same inputs, even with deep_gemm.use_deterministic_algorithms(True).

On SM90, sm90_bmn_bnk_mn_gemm splits the s * k reduction across CTAs (split_factor) and adds each CTA's partial sum into D with float2 atomicAdd. The order of those additions changes between runs. With DeepGEMM main (78b6900) on an H200, 50 identical calls give:

Shape (s, m, n, k) Default use_deterministic_algorithms(True)
121, 128, 1024, 64 50 distinct outputs 50 distinct outputs
4096, 128, 384, 128 50 distinct outputs 49 distinct outputs

The flag is already honored by the other einsum paths, the BF16 GEMM, mega_mhc and mega_gate, but not by this one.

Fix (SM90)

In deterministic mode, and only when there is more than one split:

  • The launcher allocates a zeroed [num_splits, m, n] FP32 workspace. Its size is at most about num_sms * 128 * 128 floats, since num_splits * num_mn_blocks is bounded by the SM count.
  • A new kUseSplitWorkspace template 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.
  • The slices are summed into D in a fixed order with 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)

  • The new test_bmk_bnk_mn_deterministic (SM90 only) fails on main and passes with this change. It checks 10 identical calls for bitwise equality and compares them with the reference.
  • Distinct outputs over 50 identical calls, with this change:
Shape (s, m, n, k) Default use_deterministic_algorithms(True)
121, 128, 1024, 64 50 1
4096, 128, 384, 128 50 1
  • python tests/test_einsum.py passes.
  • End-to-end time per call, in µs, including the workspace zeroing and the final reduction. The default mode is unchanged:
s, m, n, k (FP32 D) main this PR, default this PR, deterministic
129, 128, 384, 128 9.1 9.1 23.9
129, 384, 128, 384 17.8 17.5 35.2
4096, 128, 384, 128 134.8 133.7 149.2
4096, 256, 256, 256 254.1 254.2 281.2
8192, 384, 128, 384 777.4 774.0 806.9

The 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.

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)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🔵 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;

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🔵 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

Copy link
Copy Markdown
Collaborator

🤖 ds-review-bot Code Review

v6

⚠️ 未完成评审(budget_exceeded:模型额度已用尽)

v5

⚠️ 未完成评审(upstream_error:模型上游服务不可用)

v4

This MR fixes non-deterministic results in the SM90 bmk,bnk-&gt;mn einsum path. The original kernel splits the s*k reduction across CTAs and reduces the partial sums into D with float2 atomicAdd, whose completion order varies between runs. The fix adds a kUseSplitWorkspace template flag: when deterministic algorithms are enabled and there is more than one split, the launcher allocates a zeroed [num_splits, m, n] FP32 workspace, each split accumulates into its own [m, n] slice (single writer per element), and the slices are reduced into D in a fixed order via split_workspace.sum(0). When the flag is off or num_splits == 1, the generated kernel and launch are equivalent to before, so default performance is unchanged. The implementation is small, self-contained, and the three-file change is coherent. The dtype/pointer handling is sound: sm90_bmn_bnk_mn_gemm always receives an FP32 destination (the BF16 API path first routes through an FP32 workspace), so d.options() yields the required FP32 workspace, and d.add_(...) preserves any initial accumulator carried in c. The uint64_t cast correctly avoids a 32-bit overflow in the split offset, and within a single split each output element is written exactly once (distinct mn-blocks cover disjoint output tiles), so the remaining atomicAdd into the zeroed slice is exact and order-independent. The new SM90-only regression test is well targeted: it forces num_splits &gt; 1, runs 10 identical calls with the flag on, asserts torch.equal bitwise equality, validates against a bmm reference for FP32 and BF16, and restores the deterministic flag in a finally block. The main residual gap, explicitly acknowledged by the author, is that the SM100 kernel still uses TMA reduce-add and remains non-deterministic.

Files reviewed: 3
Issues found: 🔵 3 suggestion
Inline comments posted: 2
General comments (无法定位到 diff): 1

⚠️ Parse warning: [v6] budget_exceeded:模型额度已用尽; [v5] upstream_error:模型上游服务不可用


📍 未定位到 diff 的评论

🔵 suggestion csrc/jit_kernels/impls/sm100_bmk_bnk_mn.hpp:L20: The SM100 counterpart still reduces splits with TMA reduce-add, so use_deterministic_algorithms(True) remains violated on SM100 for this expression. The description acknowledges this and states the same per-split workspace approach should apply. Since the deterministic contract is user-visible and architecture-independent, please file/track a follow-up so SM100 does not silently regress this guarantee, even though the code change itself is reasonably out of scope for this SM90 fix. 🤖 v4

Signed-off-by: RunguoLi <li19107254665@gmail.com>
@LiRunGuo

Copy link
Copy Markdown
Author

@RayWang96 Could you review this when you have a chance? It makes the SM90 bmk,bnk->mn einsum honor use_deterministic_algorithms(True) with a per-split workspace, and leaves the default path unchanged. I also added the FP32 assertion the review bot suggested (ad03662). SM100 still uses TMA reduce-add because I have no SM100 machine to test on, and I'm happy to follow up there if the approach looks right to you. #456 adds a test at the same spot in tests/test_einsum.py, so I'll rebase whichever one lands second. Thanks!

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants