Repository navigation
Round fused SwiGLU to the storage dtype before ConvRot. - #124
niehen6174 wants to merge 1 commit into
Conversation
The int8_linear(input_act=swiglu) path left silu*up in fp32, so the per-row absmax diverged from F.silu(gate)*up and rewrote the INT8 row.
|
✅ All contributors have signed the CLA. Thank you! This PR is ready to be merged. |
|
No actionable comments were generated in the recent review. 🎉 ℹ️ Recent review info⚙️ Run configurationConfiguration used: Organization UI Review profile: ASSERTIVE Plan: Pro Plus Run ID: 📒 Files selected for processing (3)
Included review availability: Your plan provides up to 2 included reviews per hour; 1 remains after this review. 📝 WalkthroughWalkthroughChangesThe fused SwiGLU paths now round SiLU and product values according to FP16 or BF16 storage precision. FP32 behavior remains unchanged. Tests now require exact CUDA parity with eager activation results. Suggested reviewers: Merge Risk: ⚪ Minimal · up to The PR is merge-ready after normal checks and review; no actionable merge-blocking risk remains. 🚥 Pre-merge checks | ✅ 2✅ Passed checks (2 passed)
✨ Finishing Touches 💡 1🛠️ Fix failing CI checks 💡
🧪 Generate unit tests (beta)
✨ Simplify code
Thanks for using CodeRabbit! It's free for OSS, and your support helps us grow. If you like it, consider giving us a shout-out. Comment |
|
I have read and agree to the Contributor License Agreement |
|
@christian-byrne Could you take a look at this PR when you have a chance? Thanks! |
Summary
int8_linear(..., input_act="swiglu")is supposed to equalint8_linear(silu(gate)*up, ...).The fused CUDA/HIP quantizer left
silu*upin fp32, so ConvRot's per-row absmaxdiverged from the eager chain (which materializes BF16/FP16). One LSB of scale
rewrites the whole INT8 row.
This PR rounds silu and the multiply to the storage dtype before the rotation.
Accuracy
Kitchen unit shapes + MiniMax-H3 fc2 (
N=5376,K=14336), isolated build of this patch,F.silu(gate)*upthenint8_linearvs fused:Same H3 t2va, 50 NFE, seed=42, INT8+FA, old fused kernel (fp32 product) vs
unfused INT8 (eager SwiGLU then quantize):
That is a different diffusion sample, not a slightly noisier tensor.
After this PR, fc2 is bit-exact, so that pair should be identical (PSNR 99).
Speed (why keep the fuse)
H3 fc2, M=32700, RTX 4090 D, CUDA event, 50 DiT blocks/step:
silu*mul+int8_linearinput_act=swiglu(this PR)20 NFE exclusive 4090 D (SGLang, earlier run with the old fuse; wall-clock
should match because the extra rounding is register-only):
Test plan
CUDA+bf16
int8_linear(x, input_act="swiglu")musttorch.equalthe eager chain.FP16/BF16 fused INT8 codes must equal
quantize(silu(gate)*up).Files
comfy_kitchen/backends/cuda/ops/int8_linear.cu— CUDA quantizercomfy_kitchen/backends/hip/hadamard.h— HIP mirrortests/test_int8_input_act.py— require bit-exact on the fused CUDA path