From 860255d24e3711115f4f44c6dae70fc9faa7f95d Mon Sep 17 00:00:00 2001 From: niehen6174 Date: Thu, 20 Aug 2026 07:18:08 +0000 Subject: [PATCH] Round fused SwiGLU to the storage dtype before ConvRot. 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. --- comfy_kitchen/backends/cuda/ops/int8_linear.cu | 9 +++++++-- comfy_kitchen/backends/hip/hadamard.h | 16 ++++++++++++++-- tests/test_int8_input_act.py | 15 ++++++++++++--- 3 files changed, 33 insertions(+), 7 deletions(-) diff --git a/comfy_kitchen/backends/cuda/ops/int8_linear.cu b/comfy_kitchen/backends/cuda/ops/int8_linear.cu index 4c7d49e7..05fd9730 100644 --- a/comfy_kitchen/backends/cuda/ops/int8_linear.cu +++ b/comfy_kitchen/backends/cuda/ops/int8_linear.cu @@ -930,10 +930,15 @@ __device__ __forceinline__ float load_input_act( const InputType* __restrict__ x, int64_t in_row, int col, int K) { if constexpr (ACT == kActSwiGLU) { - // Matches torch silu(gate) * up. + // Match F.silu(gate) * up in the storage dtype. Leaving the product in + // fp32 changes the ConvRot row absmax versus the eager chain, which + // materializes silu and the multiply before quantization. One LSB of + // scale moves the whole INT8 row. const float gate = to_float(x[in_row + col]); const float up = to_float(x[in_row + K + col]); - return (gate / (1.0f + expf(-gate))) * up; + const float silu = gate / (1.0f + expf(-gate)); + const float silu_r = to_float(from_float(silu)); + return to_float(from_float(silu_r * up)); } else { return apply_input_act(to_float(x[in_row + col])); } diff --git a/comfy_kitchen/backends/hip/hadamard.h b/comfy_kitchen/backends/hip/hadamard.h index 1cde6c82..d4b707b7 100644 --- a/comfy_kitchen/backends/hip/hadamard.h +++ b/comfy_kitchen/backends/hip/hadamard.h @@ -18,6 +18,8 @@ #include #include +#include "rope_math.h" + namespace comfy::hip_backend { // convrot_quant_kernel handles 256/G groups per pass and rotates in log4(G) @@ -112,10 +114,20 @@ template __forceinline__ __device__ float load_input_act( const void* x, int64_t in_row, int col, int K, int code) { if constexpr (ACT == kActSwiGLU) { - // Matches torch silu(gate) * up. + // Match F.silu(gate) * up in the storage dtype. See the CUDA + // load_input_act comment: an fp32 product changes the row absmax. const float gate = load_in(x, in_row + col, code); const float up = load_in(x, in_row + K + col, code); - return (gate / (1.0f + expf(-gate))) * up; + const float silu = gate / (1.0f + expf(-gate)); + if (code == 0) { + return silu * up; + } + if (code == 1) { + const float silu_r = round_fp16(silu); + return round_fp16(silu_r * up); + } + const float silu_r = round_bf16(silu); + return round_bf16(silu_r * up); } else { return apply_input_act(load_in(x, in_row + col, code)); } diff --git a/tests/test_int8_input_act.py b/tests/test_int8_input_act.py index 9105cf6b..45e06460 100644 --- a/tests/test_int8_input_act.py +++ b/tests/test_int8_input_act.py @@ -214,6 +214,12 @@ def test_at_least_as_accurate_as_eager_chain(self, dtype, shape, seed, cuda_avai assert err_fused <= max(err_chain * 1.05, 1e-4), ( f"fused ({err_fused:.6f}) less accurate than chain ({err_chain:.6f})" ) + # Storage-dtype rounding must match gelu/silu-then-quantize, otherwise + # the row scale moves and every INT8 code on the row can change. + if dtype in (torch.float16, torch.bfloat16): + assert torch.equal(fused_q, chain_q), ( + f"{dtype} fused INT8 codes differ from the eager chain" + ) assert fused_q.shape == (m, k) assert fused_q.dtype == torch.int8 assert fused_s.shape == (m, 1) @@ -241,9 +247,12 @@ def test_matches_eager_activation(self, backend, seed, cuda_available): ) assert got.shape == (m, n) - denom = ref.float().abs().max() - rel = ((got.float() - ref.float()).abs().max() / denom).item() - assert rel < 0.05, f"{backend}: rel={rel:.3e}" + if backend == "cuda" and device == "cuda": + assert torch.equal(got, ref), f"{backend}: fused output is not bit-exact" + else: + denom = ref.float().abs().max() + rel = ((got.float() - ref.float()).abs().max() / denom).item() + assert rel < 0.05, f"{backend}: rel={rel:.3e}" @pytest.mark.parametrize( "tag,shape,kwargs",