Multi-Card Communication-Computation Fusion Operator Precision Verification
Precision verification for “communication + computation” fusion operators in multi-card (NPU) environments. Core approach: Decompose fusion operators into communication and computation parts for independent verification, then compare against PyTorch distributed official implementation as the overall baseline, refining “end-to-end success” into “every link is correct”.
1 Background and Problem
Fusion operators compress cross-card data movement (collective) and local computation (gemm / element-wise / reordering) into a single kernel in forward(). This approach performs well but creates difficulties in precision determination:
Difficulty identifying error source: When overall output doesn’t match golden, it’s hard to directly determine whether communication moved wrong data or computation produced incorrect results.
Single-pass cannot reproduce: Multi-card scenarios may exhibit inter-process timing-dependent errors that single-card tests mask communication path issues.
Baseline must be trustworthy: PyTorch distributed implementation is our reference baseline. If the baseline itself is misconfigured (world_size, grouping, backend), all comparisons are invalid.
Therefore, precision verification should not only perform “final result comparison”, but should be layered:
Layer 1: Communication path verification — Only concerned with whether data arrives, assembles, and reduces as expected
Layer 2: Single-card computation semantics — Only concerned with whether NPU kernel numerical results are correct
Layer 3: Fusion overall baseline comparison — On real multi-card, full comparison against PyTorch distributed implementation
Only when all three layers pass can the precision conclusion for the fusion operator be given.
2 Overall Verification Framework
2.1 Three Execution Entries
Layer |
Verification Target |
Execution Method |
Determination Goal |
|---|---|---|---|
Communication verification |
Communication sub-module (isolated) |
Single or multi-core, sending without computing is core |
Move/assemble/reduce results correct |
Computation verification |
Local computation kernel |
Single card, fed fixed input |
Numerical semantics correct (including NaN/boundary) |
Overall comparison |
Fusion operator |
Real multi-card distributed |
Consistent with PyTorch distributed baseline |
2.2 Common Verification Tools
The following tools are recommended to be fixed in the project to reduce manual comparison:
Golden generation: Use PyTorch distributed official APIs as truth source (
dist.all_gather/all_reduce/reduce_scatter/all_to_all_singleetc.).Fixed seed: All random inputs use fixed seed to ensure reproducibility.
Relaxed + strict dual threshold: Numerical comparison uses
allclose(default relative/absolute tolerance), and separately performs bit-wise or extremely strict checks when input happens to be binary-exactly representable.
3 Precision Comparison and Error Analysis
All three layers ultimately fall to the same action: numerically compare “calculation result” with “reference result (Golden)”. This section defines the unified implementation, evaluation criteria, and considerations for this action.
3.1 Obtaining Golden
Golden is the source of truth determination, prioritized from high to low:
Equivalent Torch operator results computed on CPU / GPU at high precision (
fp32/fp64);Same Triton operator computation results on CPU / GPU (for verifying implementation semantic consistency).
Note: Golden should be computed at higher precision when possible to avoid including reference implementation’s own rounding errors in tolerance. If target dtype is
fp16/bf16, it’s recommended to compute golden infp32then convert back to target dtype.
3.2 Comparison Determination Entry
In distributed environments, fusion operators are no longer just “local computation” on a single card, but first move data across cards, then complete reordering/computation locally (communication + computation fusion). The following uses HCCL Reverse All-to-All fusion operator as an example: each rank holds a local shard, and moves all ranks’ shards across cards through symmetric memory, reordering them into complete results. Example from
04-ascend-reverse-all2all.py.
Distributed Reverse All-to-All example (fused communication + computation): https://gitcode.com/Ascend/Triton-distributed-ascend/blob/master/tutorials/ascend/04-ascend-reverse-all2all/04-ascend-reverse-all2all.py
Note (example key points):
Use
dist.init_process_group(backend="hccl")to initialize process group, and complete cross-card data movement through ACL symmetric memory (aclshmem_*);Fusion kernel
kernel_hccl_reverse_a2a_pipelinedcompletes both cross-card movement and local reordering within a single@triton.jitkernel (Producer/Consumer phases);Golden constructed using PyTorch distributed official implementation
torch_reverse_a2a(dist.all_to_all_single), then compared with fusion operator output usingassert_close;Precision comparison collects pass flags from each rank via
dist.all_gather, only passes when all pass.
Corresponding run command (multi-card launch, LOCAL_RANK automatically injected by torchrun):
torchrun --nproc_per_node=2 --master_port=29500 04-ascend-reverse-all2all.py
3.3 Precision Comparison Function (Unified Entry)
def compare_precision(cal, ref, rtol=1e-3, atol=1e-3):
"""
Precision comparison function: select appropriate comparison strategy based on data type.
Args:
cal: Calculation result
ref: Reference result
rtol: Relative error tolerance
atol: Absolute error tolerance
Raises:
AssertionError: Raised when precision is insufficient
"""
assert cal.dtype == ref.dtype, f"dtype mismatch: {cal.dtype} vs {ref.dtype}"
tensor_dtype = cal.dtype
if tensor_dtype == torch.float16:
torch.testing.assert_close(ref, cal, rtol=rtol, atol=atol, equal_nan=True)
elif tensor_dtype == torch.bfloat16:
torch.testing.assert_close(ref, cal, rtol=5e-3, atol=5e-3, equal_nan=True)
elif tensor_dtype == torch.float32:
torch.testing.assert_close(ref, cal, rtol=1e-5, atol=1e-5, equal_nan=True)
elif tensor_dtype in [torch.int64, torch.int32, torch.int16, torch.int8]:
assert torch.equal(cal, ref), f"Integer tensors are not equal for dtype {tensor_dtype}"
elif tensor_dtype == torch.bool:
assert torch.equal(cal, ref), "Boolean tensors are not equal"
else:
raise ValueError(f"Unsupported tensor dtype: {tensor_dtype}")
print(f"dtype: {tensor_dtype} — Precision check passed.")
3.4 Determination Criteria
torch.testing.assert_close/torch.equalnot raising exception → pass, otherwise fail.torch.testing.assert_close: Passes if tensors are approximately equal within specified tolerance, otherwise raisesAssertionError.torch.equal: ReturnsTrueonly if shapes are identical and all elements are absolutely equal at binary level.
torch.testing.assert_closeinternal criterion:|cal - ref| <= atol + rtol * |ref|
That is, absolute error must fall within the dynamic boundary formed by “absolute tolerance + relative tolerance × reference value magnitude”. When
refis close to 0,atoldominates; whenrefis large,rtoldominates.
Recommended tolerances by data type:
Data Type |
rtol |
atol |
Description |
|---|---|---|---|
float32 |
1e-5 |
1e-5 |
Strict |
float16 |
1e-3 |
1e-3 |
Lower precision, appropriately relaxed |
bfloat16 |
5e-3 |
5e-3 |
Lower precision, appropriately relaxed |
int8/16/32/64 |
— |
— |
Must be completely consistent ( |
bool |
— |
— |
Must be completely consistent ( |
Note: The above tolerances are engineering defaults. Fusion operators may slightly relax
fp16/bf16due to stacked communication and computation errors, but must be explicitly documented in test documentation, not silently relaxed.
3.5 Considerations
NaN / Inf handling:
equal_nan=Truetreats NaN as equal. If strict detection of NaN differences is needed, should be set toFalse. Whether NaN propagates is operator semantics, testing should choose this switch according to semantics.Integer types:
int/booldo not allow any error, must be strictly consistent.Device alignment: When comparing across devices, must move
calandrefto same device (e.g., CPU) before comparison to avoid misjudgment from underlying representation differences. Notetriton_cal.cpu()should be placed before comparison.dtype alignment: Assert
cal.dtype == ref.dtypebefore comparison, direct failure if dtype mismatch, error message must be clear.
4 Communication Path Verification
The goal of communication verification is to prove cross-card data movement itself is correct, independent of computation. Therefore, communication behavior should be isolated from the fusion kernel and constructed and compared separately.
4.1 Isolation Method
Construct a “communication only, no computation” test kernel: it only
load/storefrom peer rank’s symmetric memory, assembles shards, or performs pureidentityreduction (no gemm / element-wise transformation).If operator structure makes it difficult to export communication kernel separately, can use “identity computation” substitute: replace computation part with
x = x, retain communication and reordering structure, thereby separately exposing communication behavior.
4.2 Verification Points
For each collective type, lock down its key invariants:
Communication Primitive |
Must-Check Invariants |
|---|---|
|
Concatenated result received by each rank is local shards from each rank concatenated in rank order; length = world_size × local length |
|
Each rank’s result is consistent; value = sum of corresponding elements from each rank (for SUM) |
|
Shard held by rank i = reduced result of corresponding shard from each rank |
|
Block received by rank i from rank j = block rank j wants to send to i; block content, block order, block size all correspond one-to-one |
4.3 Determination
For each rank, independently compare communication output with PyTorch corresponding collective’s golden, only passes when every rank passes.
At communication layer, don’t use relaxed element-wise replacement to mask error paths; prioritize using sufficiently distinctive input (e.g., each rank uses blocks with different values), so that any misalignment is immediately exposed.
4.4 Focus Points
rank order / block granularity: Most error-prone is “who sent to whom, spliced at which position”.
Symmetric memory layout: Once offset / stride / shard shape doesn’t match launcher parameters, cross-rank misalignment occurs without process errors.
Synchronization: Communication correctness also depends on barrier / wait / notify order; communication verification simultaneously observes whether hanging or data not ready occurs.
5 Computation Semantics Verification
The goal of computation verification is to prove local numerical computation is correct, independent of communication. On single card, directly give kernel deterministic input for comparison, can quickly locate numerical errors.
5.1 Isolation Method
In single-card environment, use “identity communication” to replace real communication: data doesn’t cross cards, directly use local memory, making kernel’s computation path testable separately.
Maintain same tile / block / num_stages / layout parameters as fusion kernel to avoid “unit test passes, fusion fails” parameter drift.
5.2 Verification Points
Numerical semantics: Allclose with
torchcorresponding operation (matmul/elementwise) under same input.NaN / Inf propagation: If operator semantics require NaN propagation, use input containing NaN to confirm propagation behavior meets expectations.
Boundary and tail: Focus on non-divisible blocks, whether mask / boundary padding correctly handles the last block.
Precision level:
fp16/bf16need to explain tolerance selection,fp32can be tightened.
5.3 Determination
Only after computation layer single-card all pass, proceed to overall comparison, avoiding bringing computation bugs to multi-card troubleshooting.
6 Fusion Overall Baseline Comparison
Execute fusion operator in real multi-card environment, perform full comparison with PyTorch distributed baseline. This is the sole adjudication entry for final precision conclusion.
6.1 Baseline Configuration Key Points
Baseline implementation: Use PyTorch distributed official API (
dist.init_process_group(backend="hccl")+ corresponding collective) to build reference implementation completely consistent with operator input/output.world_size / rank alignment: Baseline and fusion operator must use same process group size and grouping; any inconsistency will invalidate comparison.
Input completely consistent: For comparison, give baseline and fusion operator same input tensor (or copy to independent memory by rank but content bit-wise identical), eliminate input differences.
Layout and dtype consistent: Input/output memory layout (contiguous/non-contiguous, NCHW/NHWC), dtype must all be consistent.
6.2 Comparison Process
1. Initialize multi-card process group, barrier ensures each rank is ready
2. Each rank generates input with fixed seed
3. Separately run PyTorch baseline and fusion operator
4. Each rank independently compares output allclose
5. Aggregate: collect pass flags from each rank, only passes when all pass
6.3 Determination
Per-rank independent allclose: Each rank’s output is compared with that rank’s golden.
Cross-rank consistency check (optional): For operators requiring consistent output across ranks (like all_reduce), additionally confirm each rank’s result is consistent.
Determination criteria: Only when all ranks all allclose, and no hanging, no crash, does overall precision pass.
7 Common Problem Troubleshooting
Phenomenon |
Possible Cause |
Recommended Action |
|---|---|---|
Overall inconsistent, but communication layer passes |
Computation semantics bug |
Return to Layer 2 to troubleshoot computation |
Overall inconsistent, but computation layer passes |
Communication move/assembly error |
Return to Layer 1 to troubleshoot communication |
Single rank passes, multi rank fails |
rank order/block/symmetric memory layout issue |
Focus on Layer 1 Section 4.4 |
Results consistent but differ at extremes |
Precision/tolerance selection inappropriate |
Explain numerical precision level, adjust tolerance |
Process hanging |
Synchronization (barrier/wait/notify) out of order |
Check synchronization order, add barrier if necessary |
8 Implementation Checklist
[ ] Fixed seed, reproducible input generator
[ ] Communication isolation test (identity substitute, per-collective verification)
[ ] Single-card computation semantics test (including NaN / boundary)
[ ] Multi-card overall vs PyTorch distributed baseline comparison
[ ] Allclose pass flag collection and aggregation for each rank
[ ] Precision level and tolerance explanation written into documentation
9 Summary
The correct approach to precision verification is layered:
Separate communication and computation, each using corresponding minimal unit to falsify/verify correctness;
Finally on real multi-card, use PyTorch distributed implementation as authoritative baseline comparison.
Layering can quickly attribute “overall mismatch” to communication or computation, avoiding needle-in-haystack search in multi-card, hard-to-reproduce environments.