Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 7 additions & 2 deletions comfy_kitchen/backends/cuda/ops/int8_linear.cu
Original file line number Diff line number Diff line change
Expand Up @@ -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<InputType>(silu));
return to_float(from_float<InputType>(silu_r * up));
} else {
return apply_input_act<ACT>(to_float(x[in_row + col]));
}
Expand Down
16 changes: 14 additions & 2 deletions comfy_kitchen/backends/hip/hadamard.h
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,8 @@
#include <hip/hip_fp16.h>
#include <hip/hip_runtime.h>

#include "rope_math.h"

namespace comfy::hip_backend {

// convrot_quant_kernel handles 256/G groups per pass and rotates in log4(G)
Expand Down Expand Up @@ -112,10 +114,20 @@ template <int ACT>
__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<ACT>(load_in(x, in_row + col, code));
}
Expand Down
15 changes: 12 additions & 3 deletions tests/test_int8_input_act.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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",
Expand Down
Loading