Skip to content

fix(kunlunxin): repair cholesky solve launch and coverage - #6864

Open
llaboon wants to merge 80 commits into
flagos-ai:masterfrom
llaboon:klx/sdnn-verified-cholesky
Open

llaboon wants to merge 80 commits into
flagos-ai:masterfrom
llaboon:klx/sdnn-verified-cholesky

Conversation

@llaboon

@llaboon llaboon commented Oct 5, 2026

Copy link
Copy Markdown
Contributor

Summary

  • Fix the Kunlunxin cholesky_solve kernel launch argument collision.
  • Pass N positionally so Triton does not receive it twice through the dynamic launcher.
  • Keep Kunlunxin real-dtype coverage aligned with the backend capability.

Verification

Verified with the official FlagGems scheduler:

python tools/run_tests.py --ops cholesky_solve --gpus 0 --dump-output
  • Accuracy: passed
  • Benchmark: passed
  • Commit: 5d7bfde5dc8f657a32d781ae8079c172eef83f45

🤖 Generated with Claude Code

llaboon added 30 commits August 17, 2026 11:42
gcd crashed on XPU with the generic kernel (data-dependent while loop ->
LLVM ERROR getWarpsPerCTA; INT_MIN abs overflow). Vendor override:
- signed-modulo Euclid with fixed-iteration bound (NITER) and active gating
- BLOCK capped at 128 (XPU compiler fails at 256: uni_sram PassManager
  failure; at 512: arith.cmpi shape mismatch)
- materialize broadcast/non-contiguous inputs via a custom triton
  gather kernel: torch-level copy_ is broken on this XPU for strided
  sources (CUDA error: invalid device function), and the gem
  torch.broadcast_tensors override mishandles zero-size dims
The generic kernel declares total/inner_size/dim_stride as tl.constexpr;
on the XPU compiler the masked store with the constexpr-folded computed
offset is miscompiled (grad values land at wrong output positions, e.g.
140/141/... instead of the selected slice). Keeping them as runtime args
fixes it; verified across shapes (2,3,4,5)/(3,7,11)/(2,1,4) x dim {0,1}
x all index cases and dtypes.
aminmax_kernel_1 loaded the same pointer range twice with different
other= (+inf for min, -inf for max). The XPU compiler fuses the two
identical masked loads and reuses the first load other value for both,
so the max pass got +inf in masked lanes -> tl.max returned inf.

Load once with other=0.0 and apply tl.where(mask, val, sentinel) with
per-pass acc_type sentinels, mirroring the dim-path aminmax_kernel.
Multi-dim amin/amax went through dim_compress -> permute().contiguous()
which dispatches to the broken XPU copy_ (invalid device function).
Redirect len(dim)>1 to the native CUDA kernel via torch.library.get_kernel
+ call_boxed(CUDA keyset), which the XPU CUDA-ABI runtime supports.
The dimless (global) two-stage path was inaccurate for large f32 inputs
(test_norm dtype1-shape2 x4: rel 1.9e-4..3.7e-4 vs 1.3e-6 allowed):
l0 overcounted by exactly +9216 elements (9x1024), and sum-based ords
showed the same ~9216 extra data samples worth of magnitude. The error
is deterministic and heap-content independent -- the tail block's masked
loads mis-apply the mask at 1024-lane granularity and duplicate real
elements into the reduction (same masked-load hazard family as the
aminmax no-dim double-load bug).

Fix: rewrite the ten dimless stage kernels to avoid masked loads
entirely (clamp indices in-bounds with an explicit int64 clamp, select
the contribution with tl.where), and cap BLOCK_SIZE at 8192: 32768-lane
tiles are additionally corrupted by silent tl.sum lane dropping (~75% of
lanes lost in a minimal unchunked kernel), while 1024/8192-lane tiles
are exact (matrix-probed at m=1e3..2.5e7).

Validation: tests/test_norm.py 234 (incl. the 4 dimless failures),
tests/test_vector_norm.py 360, aminmax/amin/amax 57 -- 654/654 passed.
The generic kernel fails on XPU for four independent reasons: uint16\narity promotes to int32 (uint16 bitcast error), libdevice nextafter\nfails elfconv, ORing int-derived and float-compare booleans makes the\ntriton compiler exit(1) silently, and any kernel-side bf16 bitcast\ntrips an LLVM Invalid-cast assertion.\n\nThe override computes everything in integer domain (sign-magnitude key\ntransform for ordering, int arithmetic for direction and zero-cross,\nnested tl.where for selection) and routes bf16 through a same-width\nfp16 view with bf16 NaN masks from the wrapper.\n\nVerification: 173/188 (master 23/188). Remaining 15 are external:\n11 fp64 cases (backend silently downcasts float64 to float32, same\nplatform gap as special_log1p) and 4 scalar_x cases (test setup hits\nthe kunlunxin copy_ strided-source bug, copy-family fix branch).
…ve sums

Root causes fixed in the dim-reduction family:
- Non-inner-dim reductions hit dim_compress -> permute().contiguous() ->
  strided copy_, which fails with "CUDA error: invalid device function"
  on XPU. Gate those cases (utils/reduce_native.dim_compress_materializes)
  and redispatch to the accurate native ATen kernels instead.
- The vendor aten::var.correction / var_mean.correction kernels are
  inaccurate for large multi-dim reductions (~1e-2 absolute at
  (200,40999,3) x dim=[0,1] while sum is ~1e-7 relative), so var, std and
  var_mean are COMPOSED from two exact sums: var = (ss - s^2/n)/(n-corr).
- std.dim redispatch re-enters the vendor op (wrong overload, 644-frame
  recursion); the composition removes that handle entirely.
- The generic welford global-var kernel fails to lower on XPU
  (arith.maxnumf MLIR error); global var now goes through the composition.
- mean_dim bmm fast path returns wrong results for tiny (M1=3) and huge
  (M1=122997, bf16) free dims; skip it there and fall through to the
  accurate native path.
- New ops/var.py, ops/norm.py and shared utils/reduce_native.py; kernel
  handles are captured at import time (before use_gems registration) to
  avoid dispatching through the flag_gems wrappers.

Leftovers (other branches): test_norm*.py routes through
ops/vector_norm.py (linalg_vector_norm); partial-dim is fixed on
klx/aminmax-fix, the dimless f32 accuracy gap needs a follow-up there.
…, combine/stats OOB leak, eager fp32 backward, impl_index arg order, _batch_norm_no_update
Verified conv2d_padding and cudnn_convolution with official tools/run_tests.py; all reported dtype means exceed the 0.8 speedup gate for the tested cases.

Co-Authored-By: Claude <noreply@anthropic.com>

This branch has not been deployed

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

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant