[Fix] Correct dense FP8 Marlin bias ordering (#35020)

This commit is contained in:
Enrique Shockwave
2026-08-17 03:43:36 +00:00
committed by GitHub
parent b6d7602914
commit 92b1d382c7
2 changed files with 58 additions and 2 deletions
@@ -186,10 +186,12 @@ def prepare_fp8_layer_for_marlin(
marlin_scales = fp8_fused_exponent_bias_into_scales(marlin_scales)
layer.weight_scale = torch.nn.Parameter(marlin_scales, requires_grad=False)
# The dense FP8 Marlin wrapper adds bias after the kernel returns, so the
# bias must remain in logical output-channel order. Only scales need the
# Marlin tile permutation.
if hasattr(layer, "bias") and layer.bias is not None:
assert layer.bias.shape == (part_size_n,)
bias = marlin_permute_bias(layer.bias)
layer.bias = torch.nn.Parameter(bias, requires_grad=False)
layer.bias = torch.nn.Parameter(layer.bias.detach(), requires_grad=False)
def prepare_moe_fp8_layer_for_marlin(
@@ -0,0 +1,54 @@
"""Tests for FP8 Marlin utilities."""
import unittest
from unittest.mock import patch
import torch
from sglang.srt.layers.quantization import marlin_utils_fp8
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
class TestFp8MarlinBias(CustomTestCase):
def test_dense_bias_remains_in_logical_output_order(self):
size_k = 32
size_n = 32
layer = torch.nn.Module()
layer.input_size_per_partition = size_k
layer.output_size_per_partition = size_n
layer.orig_dtype = torch.float16
layer.weight_block_size = None
layer.weight = torch.nn.Parameter(
torch.zeros((size_k, size_n), dtype=torch.float8_e4m3fn),
requires_grad=False,
)
layer.weight_scale = torch.nn.Parameter(
torch.ones((size_n,), dtype=torch.float32), requires_grad=False
)
original_bias = torch.arange(size_n, dtype=torch.float16)
layer.bias = torch.nn.Parameter(original_bias.clone(), requires_grad=False)
with (
patch.object(
marlin_utils_fp8,
"marlin_make_workspace",
return_value=torch.empty(0, dtype=torch.int32),
),
patch.object(
marlin_utils_fp8,
"gptq_marlin_repack",
return_value=torch.empty(0, dtype=torch.int32),
create=True,
),
):
marlin_utils_fp8.prepare_fp8_layer_for_marlin(layer)
torch.testing.assert_close(layer.bias, original_bias)
if __name__ == "__main__":
unittest.main(verbosity=3)