[1/N] Quantization Refactor: remove dead code and dedup the FP4 marlin helpers (#37552)
This commit is contained in:
@@ -62,7 +62,6 @@ WEIGHT_LOADER_V2_SUPPORTED = [
|
|||||||
"GPTQMarlinLinearMethod",
|
"GPTQMarlinLinearMethod",
|
||||||
"Fp8LinearMethod",
|
"Fp8LinearMethod",
|
||||||
"BlockInt8LinearMethod",
|
"BlockInt8LinearMethod",
|
||||||
"MarlinLinearMethod",
|
|
||||||
"QQQLinearMethod",
|
"QQQLinearMethod",
|
||||||
"GPTQMarlin24LinearMethod",
|
"GPTQMarlin24LinearMethod",
|
||||||
"TPUInt8LinearMethod",
|
"TPUInt8LinearMethod",
|
||||||
|
|||||||
@@ -3,20 +3,7 @@
|
|||||||
# Adapted from https://raw.githubusercontent.com/vllm-project/vllm/v0.5.5/vllm/model_executor/layers/quantization/__init__.py
|
# Adapted from https://raw.githubusercontent.com/vllm-project/vllm/v0.5.5/vllm/model_executor/layers/quantization/__init__.py
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import builtins
|
from typing import Dict, Type
|
||||||
import inspect
|
|
||||||
from typing import TYPE_CHECKING, Dict, Optional, Type
|
|
||||||
|
|
||||||
import torch
|
|
||||||
|
|
||||||
|
|
||||||
# Define empty classes as placeholders when vllm is not available
|
|
||||||
class DummyConfig:
|
|
||||||
def override_quantization_method(self, *args, **kwargs):
|
|
||||||
return None
|
|
||||||
|
|
||||||
|
|
||||||
CompressedTensorsConfig = DummyConfig
|
|
||||||
|
|
||||||
from sglang.srt.layers.quantization.auto_round import AutoRoundConfig
|
from sglang.srt.layers.quantization.auto_round import AutoRoundConfig
|
||||||
from sglang.srt.layers.quantization.awq import (
|
from sglang.srt.layers.quantization.awq import (
|
||||||
@@ -72,9 +59,6 @@ from sglang.srt.utils import (
|
|||||||
|
|
||||||
_is_gfx95_supported = is_gfx95_supported()
|
_is_gfx95_supported = is_gfx95_supported()
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
|
||||||
from sglang.srt.layers.moe.topk import TopKOutput
|
|
||||||
|
|
||||||
# Base quantization methods
|
# Base quantization methods
|
||||||
BASE_QUANTIZATION_METHODS: Dict[str, Type[QuantizationConfig]] = {
|
BASE_QUANTIZATION_METHODS: Dict[str, Type[QuantizationConfig]] = {
|
||||||
"fp8": Fp8Config,
|
"fp8": Fp8Config,
|
||||||
@@ -184,6 +168,3 @@ def get_quantization_config(quantization: str) -> Type[QuantizationConfig]:
|
|||||||
return config
|
return config
|
||||||
|
|
||||||
return QUANTIZATION_METHODS[quantization]
|
return QUANTIZATION_METHODS[quantization]
|
||||||
|
|
||||||
|
|
||||||
original_isinstance = builtins.isinstance
|
|
||||||
|
|||||||
@@ -1597,10 +1597,6 @@ def quant_weight_ue8m0(
|
|||||||
return out_w, out_s
|
return out_w, out_s
|
||||||
|
|
||||||
|
|
||||||
def transform_scale_ue8m0_inplace(param, mn):
|
|
||||||
param.data = transform_scale_ue8m0(param.data, mn=mn)
|
|
||||||
|
|
||||||
|
|
||||||
# NOTE copy and modified from DeepGEMM
|
# NOTE copy and modified from DeepGEMM
|
||||||
def transform_scale_ue8m0(sf, mn, use_torch_impl: bool = False):
|
def transform_scale_ue8m0(sf, mn, use_torch_impl: bool = False):
|
||||||
import deep_gemm.utils.layout
|
import deep_gemm.utils.layout
|
||||||
|
|||||||
@@ -6,7 +6,6 @@ from humming.layer import HummingInputSchema, HummingMethod
|
|||||||
from humming.schema import BaseWeightSchema
|
from humming.schema import BaseWeightSchema
|
||||||
|
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.layers.linear import LinearBase
|
|
||||||
from sglang.srt.layers.moe import get_moe_a2a_backend
|
from sglang.srt.layers.moe import get_moe_a2a_backend
|
||||||
from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
|
from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
|
||||||
from sglang.srt.runtime_context import get_exec
|
from sglang.srt.runtime_context import get_exec
|
||||||
@@ -81,46 +80,6 @@ def make_humming_deepep_input_schema(
|
|||||||
return HummingInputSchema(a_dtype="float8e4m3", input_scale_group_size=128)
|
return HummingInputSchema(a_dtype="float8e4m3", input_scale_group_size=128)
|
||||||
|
|
||||||
|
|
||||||
def prepare_humming_layer(layer: LinearBase, quant_config: dict):
|
|
||||||
weight_schema = BaseWeightSchema.from_config(quant_config)
|
|
||||||
input_schema = HummingInputSchema()
|
|
||||||
|
|
||||||
shape_k_stacks = [layer.input_size_per_partition]
|
|
||||||
shape_n_stacks = layer.output_partition_sizes
|
|
||||||
|
|
||||||
# Step 1: convert weight to humming standard format
|
|
||||||
weight_schema, tensors = weight_schema.convert_humming(
|
|
||||||
tensors=layer.named_parameters(),
|
|
||||||
shape_n_stacks=shape_n_stacks,
|
|
||||||
shape_k_stacks=shape_k_stacks,
|
|
||||||
param_dtype=layer.params_dtype,
|
|
||||||
)
|
|
||||||
|
|
||||||
layer.weight_schema = weight_schema
|
|
||||||
|
|
||||||
for name, _ in list(layer.named_parameters()):
|
|
||||||
delattr(layer, name)
|
|
||||||
|
|
||||||
for name, tensor in tensors.items():
|
|
||||||
param = torch.nn.Parameter(tensor, requires_grad=False)
|
|
||||||
setattr(layer, name, param)
|
|
||||||
|
|
||||||
# Step 2: transform weight (humming standard format) for forwarding
|
|
||||||
HummingMethod.prepare_layer_meta(
|
|
||||||
layer=layer,
|
|
||||||
shape_n=layer.output_partition_sizes_sum,
|
|
||||||
shape_k=layer.input_size_per_partition,
|
|
||||||
weight_schema=weight_schema,
|
|
||||||
input_schema=input_schema,
|
|
||||||
pad_n_to_multiple=256,
|
|
||||||
pad_k_to_multiple=128,
|
|
||||||
has_bias=layer.has_bias,
|
|
||||||
torch_dtype=layer.param_dtype,
|
|
||||||
)
|
|
||||||
|
|
||||||
HummingMethod.transform_humming_layer(layer)
|
|
||||||
|
|
||||||
|
|
||||||
def prepare_humming_moe_layer(layer: FusedMoE, quant_config: dict):
|
def prepare_humming_moe_layer(layer: FusedMoE, quant_config: dict):
|
||||||
weight_schema = BaseWeightSchema.from_config(quant_config)
|
weight_schema = BaseWeightSchema.from_config(quant_config)
|
||||||
input_quant_config = envs.SGLANG_HUMMING_INPUT_QUANT_CONFIG.get() or {}
|
input_quant_config = envs.SGLANG_HUMMING_INPUT_QUANT_CONFIG.get() or {}
|
||||||
|
|||||||
@@ -1,4 +1,4 @@
|
|||||||
from typing import List, Optional, Tuple
|
from typing import List, Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
@@ -31,19 +31,6 @@ def apply_w8a8_block_int8_linear(
|
|||||||
return output.to(dtype=input.dtype).view(*output_shape)
|
return output.to(dtype=input.dtype).view(*output_shape)
|
||||||
|
|
||||||
|
|
||||||
def input_to_int8(
|
|
||||||
x: torch.Tensor, dtype: torch.dtype = torch.int8
|
|
||||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
|
||||||
"""This function quantizes input values to int8 values with tensor-wise quantization."""
|
|
||||||
iinfo = torch.iinfo(dtype)
|
|
||||||
min_val, max_val = x.aminmax()
|
|
||||||
amax = torch.maximum(min_val.abs(), max_val.abs()).clamp(min=1e-12)
|
|
||||||
int8_min, int8_max = iinfo.min, iinfo.max
|
|
||||||
scale = int8_max / amax
|
|
||||||
x_scl_sat = (x * scale).clamp(min=int8_min, max=int8_max)
|
|
||||||
return x_scl_sat.to(dtype).contiguous(), scale.float().reciprocal()
|
|
||||||
|
|
||||||
|
|
||||||
def block_dequant(
|
def block_dequant(
|
||||||
x_q_block: torch.Tensor,
|
x_q_block: torch.Tensor,
|
||||||
x_s: torch.Tensor,
|
x_s: torch.Tensor,
|
||||||
|
|||||||
@@ -12,19 +12,10 @@
|
|||||||
# limitations under the License.
|
# limitations under the License.
|
||||||
# ==============================================================================
|
# ==============================================================================
|
||||||
|
|
||||||
# Define a enum class for FP4 formats, including MXFP4, NVFP4 and future formats
|
|
||||||
from enum import Enum
|
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.runtime_context import get_platform
|
from sglang.srt.runtime_context import get_platform
|
||||||
|
|
||||||
|
|
||||||
class FP4KVCacheRecipe(Enum):
|
|
||||||
MXFP4 = 1 # KVFP4: block-wise scaling
|
|
||||||
NVFP4 = 2 # two-level scaling: global FP32 + block FP8 E4M3
|
|
||||||
|
|
||||||
|
|
||||||
E2M1_MAX = 6.0
|
E2M1_MAX = 6.0
|
||||||
MAX_BLOCK_SCALE_FP8 = 448.0 # Maximum FP8 E4M3 value
|
MAX_BLOCK_SCALE_FP8 = 448.0 # Maximum FP8 E4M3 value
|
||||||
# E2M1 format: 1 sign bit + 2 exponent bits + 1 mantissa bit = 4 bits
|
# E2M1 format: 1 sign bit + 2 exponent bits + 1 mantissa bit = 4 bits
|
||||||
|
|||||||
@@ -6,21 +6,11 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import logging
|
import logging
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import TYPE_CHECKING, Any, Optional
|
from typing import TYPE_CHECKING, Optional
|
||||||
|
|
||||||
import numpy
|
import numpy
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.layers.parameter import (
|
|
||||||
BasevLLMParameter,
|
|
||||||
ChannelQuantScaleParameter,
|
|
||||||
GroupQuantScaleParameter,
|
|
||||||
PackedvLLMParameter,
|
|
||||||
)
|
|
||||||
from sglang.srt.layers.quantization.base_config import (
|
|
||||||
LinearMethodBase,
|
|
||||||
QuantizationConfig,
|
|
||||||
)
|
|
||||||
from sglang.srt.layers.quantization.utils import (
|
from sglang.srt.layers.quantization.utils import (
|
||||||
get_scalar_types,
|
get_scalar_types,
|
||||||
pack_cols,
|
pack_cols,
|
||||||
@@ -37,12 +27,6 @@ from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph impo
|
|||||||
get_tc_piecewise_forward_context,
|
get_tc_piecewise_forward_context,
|
||||||
)
|
)
|
||||||
|
|
||||||
try:
|
|
||||||
from vllm import _custom_ops as ops
|
|
||||||
except ImportError:
|
|
||||||
ops = None
|
|
||||||
|
|
||||||
|
|
||||||
_is_cuda = is_cuda()
|
_is_cuda = is_cuda()
|
||||||
|
|
||||||
if _is_cuda:
|
if _is_cuda:
|
||||||
@@ -311,12 +295,6 @@ def marlin_make_empty_g_idx(device: torch.device) -> torch.Tensor:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def marlin_make_empty_zp(device: torch.device) -> torch.Tensor:
|
|
||||||
return torch.nn.Parameter(
|
|
||||||
torch.empty(0, dtype=torch.int, device=device), requires_grad=False
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def marlin_sort_g_idx(g_idx: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
|
def marlin_sort_g_idx(g_idx: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
|
||||||
g_idx_sort_indices = torch.argsort(g_idx).to(torch.int)
|
g_idx_sort_indices = torch.argsort(g_idx).to(torch.int)
|
||||||
return g_idx[g_idx_sort_indices], g_idx_sort_indices
|
return g_idx[g_idx_sort_indices], g_idx_sort_indices
|
||||||
@@ -623,267 +601,6 @@ def apply_awq_marlin_linear(
|
|||||||
return output.reshape(out_shape)
|
return output.reshape(out_shape)
|
||||||
|
|
||||||
|
|
||||||
class MarlinConfig(QuantizationConfig):
|
|
||||||
"""Config class for Marlin.
|
|
||||||
|
|
||||||
Reference: https://github.com/IST-DASLab/marlin/tree/master
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
group_size: int,
|
|
||||||
lm_head_quantized: bool,
|
|
||||||
) -> None:
|
|
||||||
super().__init__()
|
|
||||||
|
|
||||||
# Group size for the quantization.
|
|
||||||
self.group_size = group_size
|
|
||||||
self.lm_head_quantized = lm_head_quantized
|
|
||||||
if self.group_size != 128 and self.group_size != -1:
|
|
||||||
raise ValueError(
|
|
||||||
"Currently, only group size 128 and -1 (channelwise) "
|
|
||||||
"is supported for Marlin, but got group_size of "
|
|
||||||
f"{self.group_size}"
|
|
||||||
)
|
|
||||||
|
|
||||||
# 4 Bits packed into 32 bit datatype.
|
|
||||||
self.pack_factor = 32 // 4
|
|
||||||
|
|
||||||
# Tile size used by marlin kernels.
|
|
||||||
self.tile_size = 16
|
|
||||||
|
|
||||||
# Min out_features dim
|
|
||||||
self.min_n_threads = 64
|
|
||||||
|
|
||||||
# Min in_features dim
|
|
||||||
self.min_k_threads = 128
|
|
||||||
|
|
||||||
# Max parallel problems to solve at once (improves large
|
|
||||||
# batch performance)
|
|
||||||
self.max_parallel = 16
|
|
||||||
|
|
||||||
# Permutation length used by the marlin kernels.
|
|
||||||
self.perm_len = 1024
|
|
||||||
|
|
||||||
def __repr__(self) -> str:
|
|
||||||
return (
|
|
||||||
f"MarlinConfig(group_size={self.group_size}, "
|
|
||||||
f"lm_head_quantized={self.lm_head_quantized})"
|
|
||||||
)
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def get_name(cls) -> str:
|
|
||||||
return "marlin"
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def get_supported_act_dtypes(cls) -> list[torch.dtype]:
|
|
||||||
return [torch.half]
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
# Need to figure it out
|
|
||||||
def get_min_capability(cls) -> int:
|
|
||||||
return 80
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def get_config_filenames(cls) -> list[str]:
|
|
||||||
return ["quantize_config.json"]
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def from_config(cls, config: dict[str, Any]) -> MarlinConfig:
|
|
||||||
group_size = cls.get_from_keys(config, ["group_size"])
|
|
||||||
lm_head_quantized = cls.get_from_keys_or(config, ["lm_head"], default=False)
|
|
||||||
return cls(group_size, lm_head_quantized)
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def override_quantization_method(cls, hf_quant_cfg, user_quant) -> Optional[str]:
|
|
||||||
# compat: autogptq >=0.8.0 use checkpoint_format: str
|
|
||||||
# compat: autogptq <=0.7.1 is_marlin_format: bool
|
|
||||||
is_marlin_format = hf_quant_cfg.get(
|
|
||||||
"checkpoint_format"
|
|
||||||
) == "marlin" or hf_quant_cfg.get("is_marlin_format", False)
|
|
||||||
|
|
||||||
is_valid_user_quant = (
|
|
||||||
user_quant is None or user_quant == "gptq" or user_quant == "marlin"
|
|
||||||
)
|
|
||||||
|
|
||||||
if is_marlin_format and is_valid_user_quant:
|
|
||||||
msg = "The model is serialized in {} format. Using {} kernel.".format(
|
|
||||||
cls.get_name(), cls.get_name()
|
|
||||||
)
|
|
||||||
logger.info(msg)
|
|
||||||
return cls.get_name()
|
|
||||||
|
|
||||||
return None
|
|
||||||
|
|
||||||
def get_quant_method(
|
|
||||||
self, layer: torch.nn.Module, prefix: str
|
|
||||||
) -> Optional[MarlinLinearMethod]:
|
|
||||||
from sglang.srt.layers.linear import LinearBase
|
|
||||||
from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead
|
|
||||||
|
|
||||||
if isinstance(layer, LinearBase) or (
|
|
||||||
isinstance(layer, ParallelLMHead) and self.lm_head_quantized
|
|
||||||
):
|
|
||||||
return MarlinLinearMethod(self)
|
|
||||||
return None
|
|
||||||
|
|
||||||
|
|
||||||
class MarlinLinearMethod(LinearMethodBase):
|
|
||||||
"""Linear method for Marlin.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
quant_config: The Marlin quantization config.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(self, quant_config: MarlinConfig):
|
|
||||||
self.quant_config = quant_config
|
|
||||||
|
|
||||||
def create_weights(
|
|
||||||
self,
|
|
||||||
layer: torch.nn.Module,
|
|
||||||
input_size_per_partition: int,
|
|
||||||
output_partition_sizes: list[int],
|
|
||||||
input_size: int,
|
|
||||||
output_size: int,
|
|
||||||
params_dtype: torch.dtype,
|
|
||||||
**extra_weight_attrs,
|
|
||||||
):
|
|
||||||
del output_size # Unused.
|
|
||||||
weight_loader = extra_weight_attrs["weight_loader"]
|
|
||||||
|
|
||||||
if params_dtype != torch.float16:
|
|
||||||
raise ValueError(
|
|
||||||
f"The params dtype must be float16, but got {params_dtype}"
|
|
||||||
)
|
|
||||||
|
|
||||||
# Validate output_size_per_partition
|
|
||||||
output_size_per_partition = sum(output_partition_sizes)
|
|
||||||
if output_size_per_partition % self.quant_config.min_n_threads != 0:
|
|
||||||
raise ValueError(
|
|
||||||
f"Weight output_size_per_partition = "
|
|
||||||
f"{output_size_per_partition} is not divisible by "
|
|
||||||
f"min_n_threads = {self.quant_config.min_n_threads}."
|
|
||||||
)
|
|
||||||
if output_size_per_partition % self.quant_config.pack_factor != 0:
|
|
||||||
raise ValueError(
|
|
||||||
f"Weight output_size_per_partition = "
|
|
||||||
f"{output_size_per_partition} is not divisible by "
|
|
||||||
f"pack_factor = {self.quant_config.pack_factor}."
|
|
||||||
)
|
|
||||||
|
|
||||||
# Validate input_size_per_partition
|
|
||||||
if input_size_per_partition % self.quant_config.min_k_threads != 0:
|
|
||||||
raise ValueError(
|
|
||||||
f"Weight input_size_per_partition = "
|
|
||||||
f"{input_size_per_partition} is not divisible by "
|
|
||||||
f"min_k_threads = {self.quant_config.min_k_threads}."
|
|
||||||
)
|
|
||||||
if (
|
|
||||||
self.quant_config.group_size != -1
|
|
||||||
and input_size_per_partition % self.quant_config.group_size != 0
|
|
||||||
):
|
|
||||||
raise ValueError(
|
|
||||||
f"Weight input_size_per_partition = "
|
|
||||||
f"{input_size_per_partition} is not divisible by "
|
|
||||||
f"group_size = {self.quant_config.group_size}."
|
|
||||||
)
|
|
||||||
|
|
||||||
# Check that we have at least 4 tiles horizontally in the shard
|
|
||||||
num_tiles_per_perm = self.quant_config.perm_len // (
|
|
||||||
self.quant_config.tile_size**2
|
|
||||||
)
|
|
||||||
if output_size_per_partition % num_tiles_per_perm != 0:
|
|
||||||
raise ValueError("Each permutation group must reside on the same gpu")
|
|
||||||
|
|
||||||
# Quantized 4Bit weights packed into Int32.
|
|
||||||
qweight = PackedvLLMParameter(
|
|
||||||
data=torch.empty(
|
|
||||||
input_size_per_partition // self.quant_config.tile_size,
|
|
||||||
output_size_per_partition
|
|
||||||
* self.quant_config.tile_size
|
|
||||||
// self.quant_config.pack_factor,
|
|
||||||
device="cuda",
|
|
||||||
dtype=torch.int32,
|
|
||||||
),
|
|
||||||
input_dim=0,
|
|
||||||
output_dim=1,
|
|
||||||
packed_dim=1,
|
|
||||||
packed_factor=self.quant_config.pack_factor,
|
|
||||||
marlin_tile_size=self.quant_config.tile_size,
|
|
||||||
weight_loader=weight_loader,
|
|
||||||
)
|
|
||||||
|
|
||||||
# Determine if channelwise or not
|
|
||||||
input_groups = (
|
|
||||||
1
|
|
||||||
if self.quant_config.group_size == -1
|
|
||||||
else input_size_per_partition // self.quant_config.group_size
|
|
||||||
)
|
|
||||||
|
|
||||||
weight_scale_args = {
|
|
||||||
"data": torch.empty(
|
|
||||||
input_groups,
|
|
||||||
output_size_per_partition,
|
|
||||||
device="cuda",
|
|
||||||
dtype=params_dtype,
|
|
||||||
),
|
|
||||||
"weight_loader": weight_loader,
|
|
||||||
}
|
|
||||||
if input_groups == 1:
|
|
||||||
scales = ChannelQuantScaleParameter(output_dim=1, **weight_scale_args)
|
|
||||||
else:
|
|
||||||
scales = GroupQuantScaleParameter(
|
|
||||||
output_dim=1, input_dim=0, **weight_scale_args
|
|
||||||
)
|
|
||||||
|
|
||||||
# Allocate workspace (Used for internal locking mechanism)
|
|
||||||
max_workspace_size = (
|
|
||||||
output_size_per_partition // self.quant_config.min_n_threads
|
|
||||||
) * self.quant_config.max_parallel
|
|
||||||
|
|
||||||
workspace = BasevLLMParameter(
|
|
||||||
data=torch.zeros(max_workspace_size, device="cuda", dtype=torch.int),
|
|
||||||
weight_loader=weight_loader,
|
|
||||||
)
|
|
||||||
|
|
||||||
layer.register_parameter("B", qweight)
|
|
||||||
layer.register_parameter("s", scales)
|
|
||||||
layer.register_parameter("workspace", workspace)
|
|
||||||
|
|
||||||
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
|
|
||||||
# required by torch.compile
|
|
||||||
layer.B = torch.nn.Parameter(layer.B.data, requires_grad=False)
|
|
||||||
layer.s = torch.nn.Parameter(layer.s.data, requires_grad=False)
|
|
||||||
layer.workspace = torch.nn.Parameter(layer.workspace.data, requires_grad=False)
|
|
||||||
|
|
||||||
def apply(
|
|
||||||
self,
|
|
||||||
layer: torch.nn.Module,
|
|
||||||
x: torch.Tensor,
|
|
||||||
bias: Optional[torch.Tensor] = None,
|
|
||||||
) -> torch.Tensor:
|
|
||||||
qweight = layer.B
|
|
||||||
scales = layer.s
|
|
||||||
workspace = layer.workspace
|
|
||||||
|
|
||||||
x_2d = x.view(-1, x.shape[-1])
|
|
||||||
|
|
||||||
size_m = x_2d.shape[0]
|
|
||||||
size_k = x_2d.shape[1]
|
|
||||||
size_n = scales.shape[1]
|
|
||||||
|
|
||||||
output_2d = ops.marlin_gemm(
|
|
||||||
x_2d, qweight, scales, workspace, size_m, size_n, size_k
|
|
||||||
)
|
|
||||||
|
|
||||||
output = output_2d.view(x.shape[:-1] + (output_2d.shape[1],))
|
|
||||||
|
|
||||||
if bias is not None:
|
|
||||||
output.add_(bias) # In-place add
|
|
||||||
|
|
||||||
return output
|
|
||||||
|
|
||||||
|
|
||||||
def fake_unified_apply_gptq_marlin_gemm(
|
def fake_unified_apply_gptq_marlin_gemm(
|
||||||
input: torch.Tensor,
|
input: torch.Tensor,
|
||||||
weight: torch.Tensor,
|
weight: torch.Tensor,
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
|
from typing import Callable
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
@@ -296,6 +297,52 @@ def deinterleave_moe_mxfp4_w13_for_marlin(layer: torch.nn.Module) -> None:
|
|||||||
w13_bias.data = bias.view(e, n // 2, 2).permute(0, 2, 1).contiguous().view(e, n)
|
w13_bias.data = bias.view(e, n // 2, 2).permute(0, 2, 1).contiguous().view(e, n)
|
||||||
|
|
||||||
|
|
||||||
|
def _repack_moe_fp4_weight_for_marlin(
|
||||||
|
weight: torch.Tensor,
|
||||||
|
*,
|
||||||
|
num_experts: int,
|
||||||
|
size_n: int,
|
||||||
|
size_k: int,
|
||||||
|
perm: torch.Tensor,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
assert weight.shape == (num_experts, size_n, size_k // 2)
|
||||||
|
|
||||||
|
tensor_list = []
|
||||||
|
for i in range(num_experts):
|
||||||
|
qweight = weight[i].view(torch.int32).T.contiguous()
|
||||||
|
marlin_qweight = gptq_marlin_repack(
|
||||||
|
b_q_weight=qweight,
|
||||||
|
perm=perm,
|
||||||
|
size_k=size_k,
|
||||||
|
size_n=size_n,
|
||||||
|
num_bits=4,
|
||||||
|
)
|
||||||
|
tensor_list.append(marlin_qweight)
|
||||||
|
return torch.stack(tensor_list)
|
||||||
|
|
||||||
|
|
||||||
|
def _permute_moe_fp4_scales_for_marlin(
|
||||||
|
scales: torch.Tensor,
|
||||||
|
*,
|
||||||
|
num_experts: int,
|
||||||
|
size_n: int,
|
||||||
|
size_k: int,
|
||||||
|
group_size: int,
|
||||||
|
process_scales: Callable[[torch.Tensor], torch.Tensor],
|
||||||
|
) -> torch.Tensor:
|
||||||
|
tensor_list = []
|
||||||
|
for i in range(num_experts):
|
||||||
|
scale = scales[i].T.contiguous()
|
||||||
|
marlin_scales = marlin_permute_scales(
|
||||||
|
s=scale,
|
||||||
|
size_k=size_k,
|
||||||
|
size_n=size_n,
|
||||||
|
group_size=group_size,
|
||||||
|
)
|
||||||
|
tensor_list.append(process_scales(marlin_scales))
|
||||||
|
return torch.stack(tensor_list)
|
||||||
|
|
||||||
|
|
||||||
def prepare_moe_mxfp4_layer_for_marlin(layer: torch.nn.Module) -> None:
|
def prepare_moe_mxfp4_layer_for_marlin(layer: torch.nn.Module) -> None:
|
||||||
group_size = 32
|
group_size = 32
|
||||||
w13 = layer.w13_weight.data
|
w13 = layer.w13_weight.data
|
||||||
@@ -361,48 +408,11 @@ def prepare_moe_mxfp4_layer_for_marlin(layer: torch.nn.Module) -> None:
|
|||||||
if w13_bias_data is not None:
|
if w13_bias_data is not None:
|
||||||
w13_bias_data = _pad_w13(w13_bias_data.unsqueeze(-1)).squeeze(-1)
|
w13_bias_data = _pad_w13(w13_bias_data.unsqueeze(-1)).squeeze(-1)
|
||||||
|
|
||||||
def _repack_weight(weight: torch.Tensor, is_w13: bool) -> torch.Tensor:
|
w13_size_n, w13_size_k = padded_intermediate_size * 2, hidden_size
|
||||||
if is_w13:
|
w2_size_n, w2_size_k = hidden_size, padded_intermediate_size
|
||||||
size_n, size_k = padded_intermediate_size * 2, hidden_size
|
|
||||||
else:
|
|
||||||
size_n, size_k = hidden_size, padded_intermediate_size
|
|
||||||
assert weight.shape == (num_experts, size_n, size_k // 2)
|
|
||||||
|
|
||||||
tensor_list = []
|
def _process_scales(marlin_scales: torch.Tensor) -> torch.Tensor:
|
||||||
for i in range(num_experts):
|
return mxfp4_marlin_process_scales(marlin_scales, input_dtype=param_dtype)
|
||||||
qweight = weight[i].view(torch.int32).T.contiguous()
|
|
||||||
marlin_qweight = gptq_marlin_repack(
|
|
||||||
b_q_weight=qweight,
|
|
||||||
perm=perm,
|
|
||||||
size_k=size_k,
|
|
||||||
size_n=size_n,
|
|
||||||
num_bits=4,
|
|
||||||
)
|
|
||||||
tensor_list.append(marlin_qweight)
|
|
||||||
return torch.stack(tensor_list)
|
|
||||||
|
|
||||||
def _permute_scales(scales: torch.Tensor, is_w13: bool) -> torch.Tensor:
|
|
||||||
if is_w13:
|
|
||||||
size_n, size_k = padded_intermediate_size * 2, hidden_size
|
|
||||||
else:
|
|
||||||
size_n, size_k = hidden_size, padded_intermediate_size
|
|
||||||
|
|
||||||
tensor_list = []
|
|
||||||
for i in range(num_experts):
|
|
||||||
scale = scales[i].T.contiguous()
|
|
||||||
marlin_scales = marlin_permute_scales(
|
|
||||||
s=scale,
|
|
||||||
size_k=size_k,
|
|
||||||
size_n=size_n,
|
|
||||||
group_size=group_size,
|
|
||||||
)
|
|
||||||
tensor_list.append(
|
|
||||||
mxfp4_marlin_process_scales(
|
|
||||||
marlin_scales,
|
|
||||||
input_dtype=param_dtype,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
return torch.stack(tensor_list)
|
|
||||||
|
|
||||||
def _permute_bias(bias: torch.Tensor | None) -> torch.Tensor | None:
|
def _permute_bias(bias: torch.Tensor | None) -> torch.Tensor | None:
|
||||||
if bias is None:
|
if bias is None:
|
||||||
@@ -412,10 +422,28 @@ def prepare_moe_mxfp4_layer_for_marlin(layer: torch.nn.Module) -> None:
|
|||||||
tensor_list.append(marlin_permute_bias(bias[i].to(param_dtype)))
|
tensor_list.append(marlin_permute_bias(bias[i].to(param_dtype)))
|
||||||
return torch.stack(tensor_list)
|
return torch.stack(tensor_list)
|
||||||
|
|
||||||
w13_marlin = _repack_weight(w13, True)
|
w13_marlin = _repack_moe_fp4_weight_for_marlin(
|
||||||
w2_marlin = _repack_weight(w2, False)
|
w13, num_experts=num_experts, size_n=w13_size_n, size_k=w13_size_k, perm=perm
|
||||||
w13_scale_marlin = _permute_scales(w13_scale_data, True)
|
)
|
||||||
w2_scale_marlin = _permute_scales(w2_scale_data, False)
|
w2_marlin = _repack_moe_fp4_weight_for_marlin(
|
||||||
|
w2, num_experts=num_experts, size_n=w2_size_n, size_k=w2_size_k, perm=perm
|
||||||
|
)
|
||||||
|
w13_scale_marlin = _permute_moe_fp4_scales_for_marlin(
|
||||||
|
w13_scale_data,
|
||||||
|
num_experts=num_experts,
|
||||||
|
size_n=w13_size_n,
|
||||||
|
size_k=w13_size_k,
|
||||||
|
group_size=group_size,
|
||||||
|
process_scales=_process_scales,
|
||||||
|
)
|
||||||
|
w2_scale_marlin = _permute_moe_fp4_scales_for_marlin(
|
||||||
|
w2_scale_data,
|
||||||
|
num_experts=num_experts,
|
||||||
|
size_n=w2_size_n,
|
||||||
|
size_k=w2_size_k,
|
||||||
|
group_size=group_size,
|
||||||
|
process_scales=_process_scales,
|
||||||
|
)
|
||||||
|
|
||||||
layer.w13_weight = torch.nn.Parameter(w13_marlin, requires_grad=False)
|
layer.w13_weight = torch.nn.Parameter(w13_marlin, requires_grad=False)
|
||||||
layer.w2_weight = torch.nn.Parameter(w2_marlin, requires_grad=False)
|
layer.w2_weight = torch.nn.Parameter(w2_marlin, requires_grad=False)
|
||||||
@@ -480,44 +508,8 @@ def prepare_moe_nvfp4_layer_for_marlin(layer: torch.nn.Module) -> None:
|
|||||||
w13_bias = torch.nn.functional.pad(w13_bias, (0, intermediate_size_pad))
|
w13_bias = torch.nn.functional.pad(w13_bias, (0, intermediate_size_pad))
|
||||||
intermediate_size = padded_intermediate_size
|
intermediate_size = padded_intermediate_size
|
||||||
|
|
||||||
def _repack_weight(weight: torch.Tensor, is_w13: bool) -> torch.Tensor:
|
w13_size_n, w13_size_k = intermediate_size * num_shards, hidden_size
|
||||||
if is_w13:
|
w2_size_n, w2_size_k = hidden_size, intermediate_size
|
||||||
size_n, size_k = intermediate_size * num_shards, hidden_size
|
|
||||||
else:
|
|
||||||
size_n, size_k = hidden_size, intermediate_size
|
|
||||||
assert weight.shape == (num_experts, size_n, size_k // 2)
|
|
||||||
|
|
||||||
tensor_list = []
|
|
||||||
for i in range(num_experts):
|
|
||||||
qweight = weight[i].view(torch.int32).T.contiguous()
|
|
||||||
marlin_qweight = gptq_marlin_repack(
|
|
||||||
b_q_weight=qweight,
|
|
||||||
perm=perm,
|
|
||||||
size_k=size_k,
|
|
||||||
size_n=size_n,
|
|
||||||
num_bits=4,
|
|
||||||
)
|
|
||||||
tensor_list.append(marlin_qweight)
|
|
||||||
return torch.stack(tensor_list)
|
|
||||||
|
|
||||||
def _permute_scales(scales: torch.Tensor, is_w13: bool) -> torch.Tensor:
|
|
||||||
scales = scales.to(param_dtype)
|
|
||||||
if is_w13:
|
|
||||||
size_n, size_k = intermediate_size * num_shards, hidden_size
|
|
||||||
else:
|
|
||||||
size_n, size_k = hidden_size, intermediate_size
|
|
||||||
|
|
||||||
tensor_list = []
|
|
||||||
for i in range(num_experts):
|
|
||||||
scale = scales[i].T.contiguous()
|
|
||||||
marlin_scales = marlin_permute_scales(
|
|
||||||
s=scale,
|
|
||||||
size_k=size_k,
|
|
||||||
size_n=size_n,
|
|
||||||
group_size=16,
|
|
||||||
)
|
|
||||||
tensor_list.append(nvfp4_marlin_process_scales(marlin_scales))
|
|
||||||
return torch.stack(tensor_list)
|
|
||||||
|
|
||||||
def _process_global_scale(global_scale: torch.Tensor) -> torch.Tensor:
|
def _process_global_scale(global_scale: torch.Tensor) -> torch.Tensor:
|
||||||
return nvfp4_marlin_process_global_scale(global_scale.to(param_dtype))
|
return nvfp4_marlin_process_global_scale(global_scale.to(param_dtype))
|
||||||
@@ -531,14 +523,42 @@ def prepare_moe_nvfp4_layer_for_marlin(layer: torch.nn.Module) -> None:
|
|||||||
return torch.stack(tensor_list)
|
return torch.stack(tensor_list)
|
||||||
|
|
||||||
layer.w13_weight = torch.nn.Parameter(
|
layer.w13_weight = torch.nn.Parameter(
|
||||||
_repack_weight(w13, True), requires_grad=False
|
_repack_moe_fp4_weight_for_marlin(
|
||||||
|
w13,
|
||||||
|
num_experts=num_experts,
|
||||||
|
size_n=w13_size_n,
|
||||||
|
size_k=w13_size_k,
|
||||||
|
perm=perm,
|
||||||
|
),
|
||||||
|
requires_grad=False,
|
||||||
|
)
|
||||||
|
layer.w2_weight = torch.nn.Parameter(
|
||||||
|
_repack_moe_fp4_weight_for_marlin(
|
||||||
|
w2, num_experts=num_experts, size_n=w2_size_n, size_k=w2_size_k, perm=perm
|
||||||
|
),
|
||||||
|
requires_grad=False,
|
||||||
)
|
)
|
||||||
layer.w2_weight = torch.nn.Parameter(_repack_weight(w2, False), requires_grad=False)
|
|
||||||
layer.w13_weight_scale = torch.nn.Parameter(
|
layer.w13_weight_scale = torch.nn.Parameter(
|
||||||
_permute_scales(w13_scale, True), requires_grad=False
|
_permute_moe_fp4_scales_for_marlin(
|
||||||
|
w13_scale.to(param_dtype),
|
||||||
|
num_experts=num_experts,
|
||||||
|
size_n=w13_size_n,
|
||||||
|
size_k=w13_size_k,
|
||||||
|
group_size=16,
|
||||||
|
process_scales=nvfp4_marlin_process_scales,
|
||||||
|
),
|
||||||
|
requires_grad=False,
|
||||||
)
|
)
|
||||||
layer.w2_weight_scale = torch.nn.Parameter(
|
layer.w2_weight_scale = torch.nn.Parameter(
|
||||||
_permute_scales(w2_scale, False), requires_grad=False
|
_permute_moe_fp4_scales_for_marlin(
|
||||||
|
w2_scale.to(param_dtype),
|
||||||
|
num_experts=num_experts,
|
||||||
|
size_n=w2_size_n,
|
||||||
|
size_k=w2_size_k,
|
||||||
|
group_size=16,
|
||||||
|
process_scales=nvfp4_marlin_process_scales,
|
||||||
|
),
|
||||||
|
requires_grad=False,
|
||||||
)
|
)
|
||||||
layer.w13_weight_scale_2 = torch.nn.Parameter(
|
layer.w13_weight_scale_2 = torch.nn.Parameter(
|
||||||
_process_global_scale(w13_global_scale), requires_grad=False
|
_process_global_scale(w13_global_scale), requires_grad=False
|
||||||
|
|||||||
@@ -8,7 +8,6 @@ import torch
|
|||||||
from sglang.srt.layers.quantization.marlin_utils import (
|
from sglang.srt.layers.quantization.marlin_utils import (
|
||||||
USE_FP32_REDUCE_DEFAULT,
|
USE_FP32_REDUCE_DEFAULT,
|
||||||
marlin_make_workspace,
|
marlin_make_workspace,
|
||||||
marlin_permute_bias,
|
|
||||||
marlin_permute_scales,
|
marlin_permute_scales,
|
||||||
should_use_atomic_add_reduce,
|
should_use_atomic_add_reduce,
|
||||||
)
|
)
|
||||||
@@ -194,136 +193,6 @@ def prepare_fp8_layer_for_marlin(
|
|||||||
layer.bias = torch.nn.Parameter(layer.bias.detach(), requires_grad=False)
|
layer.bias = torch.nn.Parameter(layer.bias.detach(), requires_grad=False)
|
||||||
|
|
||||||
|
|
||||||
def prepare_moe_fp8_layer_for_marlin(
|
|
||||||
layer: torch.nn.Module, size_k_first: bool = True
|
|
||||||
) -> None:
|
|
||||||
logger.warning_once(
|
|
||||||
"Your GPU does not have native support for FP8 computation but "
|
|
||||||
"FP8 quantization is being used. Weight-only FP8 compression will "
|
|
||||||
"be used leveraging the Marlin kernel. This may degrade "
|
|
||||||
"performance for compute-heavy workloads."
|
|
||||||
)
|
|
||||||
|
|
||||||
e = layer.num_experts
|
|
||||||
k = layer.hidden_size
|
|
||||||
n = layer.intermediate_size_per_partition
|
|
||||||
weight_block_size = getattr(layer, "weight_block_size", None)
|
|
||||||
|
|
||||||
# WORKSPACE
|
|
||||||
device = layer.w13_weight.device
|
|
||||||
layer.workspace = marlin_make_workspace(device, 4)
|
|
||||||
perm = torch.empty(0, dtype=torch.int, device=device)
|
|
||||||
|
|
||||||
# WEIGHT
|
|
||||||
# Repack weights to marlin format
|
|
||||||
for name in ["w13_weight", "w2_weight"]:
|
|
||||||
weight = getattr(layer, name)
|
|
||||||
tensor_list = []
|
|
||||||
if "w13" in name:
|
|
||||||
size_n, size_k = n * 2, k
|
|
||||||
else:
|
|
||||||
size_n, size_k = k, n
|
|
||||||
|
|
||||||
if size_k_first:
|
|
||||||
assert weight.shape == (e, size_k, size_n)
|
|
||||||
else:
|
|
||||||
assert weight.shape == (e, size_n, size_k)
|
|
||||||
|
|
||||||
for i in range(e):
|
|
||||||
qweight = pack_fp8_to_int32(weight[i], size_k_first)
|
|
||||||
if not size_k_first:
|
|
||||||
qweight = qweight.T.contiguous()
|
|
||||||
|
|
||||||
marlin_qweight = gptq_marlin_repack(
|
|
||||||
b_q_weight=qweight, perm=perm, size_k=size_k, size_n=size_n, num_bits=8
|
|
||||||
)
|
|
||||||
tensor_list.append(marlin_qweight)
|
|
||||||
|
|
||||||
weight = torch.cat([x.unsqueeze(0) for x in tensor_list], 0)
|
|
||||||
weight = torch.nn.Parameter(weight, requires_grad=False)
|
|
||||||
|
|
||||||
setattr(layer, name, weight)
|
|
||||||
|
|
||||||
# WEIGHT SCALES
|
|
||||||
# Permute scales
|
|
||||||
group_size = -1 if weight_block_size is None else weight_block_size[1]
|
|
||||||
|
|
||||||
for name in ["w13", "w2"]:
|
|
||||||
if name + "_weight_scale" in dir(layer):
|
|
||||||
new_name = name + "_weight_scale"
|
|
||||||
scales = getattr(layer, new_name).to(layer.orig_dtype)
|
|
||||||
delattr(layer, new_name)
|
|
||||||
elif name + "_weight_scale_inv" in dir(layer):
|
|
||||||
new_name = name + "_weight_scale_inv"
|
|
||||||
scales = getattr(layer, new_name).to(layer.orig_dtype)
|
|
||||||
delattr(layer, new_name)
|
|
||||||
|
|
||||||
tensor_list = []
|
|
||||||
if "w13" in name:
|
|
||||||
size_n, size_k = n * 2, k
|
|
||||||
else:
|
|
||||||
size_n, size_k = k, n
|
|
||||||
|
|
||||||
# marlin kernel only support channel-wise and group-wise quantization
|
|
||||||
# we need to convert the scales
|
|
||||||
if weight_block_size is None:
|
|
||||||
if scales.nelement() == e:
|
|
||||||
# tensor-wise quantization -> channel-wise quantization
|
|
||||||
# (e, 1, 1) =>(repeat)=> (e, 1, size_n)
|
|
||||||
scales = scales.view(e, 1, 1).repeat_interleave(size_n, 2)
|
|
||||||
elif scales.nelement() > e and scales.nelement() != e * size_n:
|
|
||||||
assert (e * size_n) % scales.nelement() == 0
|
|
||||||
s_size = scales.nelement() // e
|
|
||||||
# tensor-wise quantization (for gate-up proj)
|
|
||||||
# -> channel-wise quantization
|
|
||||||
# (e, 1, s_size) =>(repeat)=> (e, 1, size_n)
|
|
||||||
scales = scales.view(e, 1, s_size)
|
|
||||||
scales = scales.repeat_interleave(size_n // s_size, 2)
|
|
||||||
else:
|
|
||||||
# channel-wise quantization
|
|
||||||
# (e, 1, size_n)
|
|
||||||
scales = scales.view(e, 1, size_n)
|
|
||||||
else:
|
|
||||||
# block-wise quantization -> group-wise quantization
|
|
||||||
# (e, size_k // block_size[1], ceil(size_n / block_size[0]))
|
|
||||||
# =>(repeat)=> (e, size_k // block_size[1], size_n)
|
|
||||||
if not size_k_first:
|
|
||||||
scales = scales.permute(0, 2, 1)
|
|
||||||
block_n = weight_block_size[0]
|
|
||||||
scales = scales.repeat_interleave(block_n, 2)
|
|
||||||
# size_n may not divisible by block_size[0]
|
|
||||||
scales = scales[..., :size_n].contiguous()
|
|
||||||
|
|
||||||
for i in range(e):
|
|
||||||
marlin_scales = marlin_permute_scales(
|
|
||||||
s=scales[i], size_k=size_k, size_n=size_n, group_size=group_size
|
|
||||||
)
|
|
||||||
tensor_list.append(marlin_scales)
|
|
||||||
|
|
||||||
scales = torch.cat([x.unsqueeze(0) for x in tensor_list], 0)
|
|
||||||
scales = fp8_fused_exponent_bias_into_scales(scales)
|
|
||||||
scales = torch.nn.Parameter(scales, requires_grad=False)
|
|
||||||
|
|
||||||
setattr(layer, name + "_weight_scale", scales)
|
|
||||||
|
|
||||||
# BIAS
|
|
||||||
# Permute bias
|
|
||||||
for name in ["w13_bias", "w2_bias"]:
|
|
||||||
if not hasattr(layer, name):
|
|
||||||
continue
|
|
||||||
bias = getattr(layer, name).to(layer.orig_dtype)
|
|
||||||
|
|
||||||
tensor_list = []
|
|
||||||
for i in range(e):
|
|
||||||
expert_bias = bias[i]
|
|
||||||
|
|
||||||
tensor_list.append(marlin_permute_bias(expert_bias))
|
|
||||||
|
|
||||||
bias = torch.cat([x.unsqueeze(0) for x in tensor_list], 0)
|
|
||||||
bias = torch.nn.Parameter(bias, requires_grad=False)
|
|
||||||
setattr(layer, name, bias)
|
|
||||||
|
|
||||||
|
|
||||||
def pack_fp8_to_int32(
|
def pack_fp8_to_int32(
|
||||||
fp8_tensor: torch.Tensor, size_k_first: bool = True
|
fp8_tensor: torch.Tensor, size_k_first: bool = True
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
@@ -339,36 +208,3 @@ def pack_fp8_to_int32(
|
|||||||
# with `.view(torch.int32)`, it become (N, K // 4)
|
# with `.view(torch.int32)`, it become (N, K // 4)
|
||||||
int32_tensor = fp8_tensor.view(torch.int32)
|
int32_tensor = fp8_tensor.view(torch.int32)
|
||||||
return int32_tensor.T.contiguous() if size_k_first else int32_tensor
|
return int32_tensor.T.contiguous() if size_k_first else int32_tensor
|
||||||
|
|
||||||
|
|
||||||
def marlin_quant_fp8_torch(weight, group_size):
|
|
||||||
size_n, size_k = weight.shape
|
|
||||||
device = weight.device
|
|
||||||
|
|
||||||
if group_size != -1:
|
|
||||||
scales = weight.view(size_n, -1, group_size).abs().max(-1)[0] / 448
|
|
||||||
repeated_scales = scales.repeat_interleave(group_size, 1)
|
|
||||||
fp8_weight = (weight / repeated_scales).to(torch.float8_e4m3fn)
|
|
||||||
weight_ref = fp8_weight.to(weight.dtype) * repeated_scales
|
|
||||||
else:
|
|
||||||
scales = weight.view(size_n, 1, group_size).abs().max(-1)[0] / 448
|
|
||||||
repeated_scales = scales.repeat_interleave(size_k, 1)
|
|
||||||
fp8_weight = (weight / repeated_scales).to(torch.float8_e4m3fn)
|
|
||||||
weight_ref = fp8_weight.to(weight.dtype) * repeated_scales
|
|
||||||
|
|
||||||
packed_weight = pack_fp8_to_int32(fp8_weight, False).T.contiguous()
|
|
||||||
marlin_qweight = gptq_marlin_repack(
|
|
||||||
b_q_weight=packed_weight,
|
|
||||||
perm=torch.empty(0, dtype=torch.int, device=device),
|
|
||||||
size_k=size_k,
|
|
||||||
size_n=size_n,
|
|
||||||
num_bits=8,
|
|
||||||
)
|
|
||||||
|
|
||||||
marlin_scales = marlin_permute_scales(
|
|
||||||
s=scales.T, size_k=size_k, size_n=size_n, group_size=group_size
|
|
||||||
)
|
|
||||||
|
|
||||||
marlin_scales = fp8_fused_exponent_bias_into_scales(marlin_scales)
|
|
||||||
|
|
||||||
return weight_ref.T, marlin_qweight, marlin_scales
|
|
||||||
|
|||||||
@@ -295,22 +295,6 @@ def dequant_mxfp4(
|
|||||||
return mx.dq_mxfp4(x, scale, float_dtype)
|
return mx.dq_mxfp4(x, scale, float_dtype)
|
||||||
|
|
||||||
|
|
||||||
@register_custom_op(out_shape="x")
|
|
||||||
def quant_dequant_mxfp4(
|
|
||||||
x: torch.Tensor, scale_calculation_mode: str = "even"
|
|
||||||
) -> torch.Tensor:
|
|
||||||
try:
|
|
||||||
from quark.torch.kernel import mx
|
|
||||||
except ImportError as err:
|
|
||||||
raise ImportError(
|
|
||||||
"The package `amd-quark` is required to use "
|
|
||||||
"MX-FP4 models. Please install it with `pip install "
|
|
||||||
"amd-quark`."
|
|
||||||
) from err
|
|
||||||
|
|
||||||
return mx.qdq_mxfp4(x, scale_calculation_mode)
|
|
||||||
|
|
||||||
|
|
||||||
class Mxfp4Config(QuantizationConfig):
|
class Mxfp4Config(QuantizationConfig):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
|
|||||||
@@ -219,29 +219,6 @@ def replace_parameter(
|
|||||||
mod.register_parameter(name, torch.nn.Parameter(new, requires_grad=False))
|
mod.register_parameter(name, torch.nn.Parameter(new, requires_grad=False))
|
||||||
|
|
||||||
|
|
||||||
def assert_fp8_all_close(a: torch.Tensor, b: torch.Tensor):
|
|
||||||
assert a.shape == b.shape
|
|
||||||
assert a.dtype == b.dtype == torch.float8_e4m3fn
|
|
||||||
|
|
||||||
a_u8 = a.view(torch.uint8)
|
|
||||||
b_u8 = b.view(torch.uint8)
|
|
||||||
diff_u8 = (a_u8.to(torch.int16) - b_u8.to(torch.int16)).abs()
|
|
||||||
|
|
||||||
numel = a.numel()
|
|
||||||
|
|
||||||
count_diff_sign = ((a_u8 >= 0) & (b_u8 < 0)).sum().item()
|
|
||||||
count_tiny_diff = (diff_u8 >= 1).sum().item()
|
|
||||||
count_large_diff = (diff_u8 >= 2).sum().item()
|
|
||||||
|
|
||||||
assert (
|
|
||||||
(count_diff_sign == 0)
|
|
||||||
and (count_tiny_diff / numel < 0.005)
|
|
||||||
and (count_large_diff == 0)
|
|
||||||
), f"{count_diff_sign=} {count_tiny_diff=} {count_large_diff=} {numel=}"
|
|
||||||
|
|
||||||
|
|
||||||
# Match dynamic rules with module name (prefix) and override quantize
|
|
||||||
# config if module (prefix) matches a rule
|
|
||||||
def override_config(config: QuantizationConfig, prefix: str):
|
def override_config(config: QuantizationConfig, prefix: str):
|
||||||
weight_bits = get_dynamic_override(config, prefix, "bits", config.weight_bits)
|
weight_bits = get_dynamic_override(config, prefix, "bits", config.weight_bits)
|
||||||
if isinstance(weight_bits, int):
|
if isinstance(weight_bits, int):
|
||||||
|
|||||||
@@ -132,7 +132,6 @@ QUANTIZATION_CHOICES = [
|
|||||||
"fp8", # MOE + linear online quantization.
|
"fp8", # MOE + linear online quantization.
|
||||||
"mxfp8", # MOE + linear online quantization.
|
"mxfp8", # MOE + linear online quantization.
|
||||||
"gptq",
|
"gptq",
|
||||||
"marlin",
|
|
||||||
"gptq_marlin",
|
"gptq_marlin",
|
||||||
"awq_marlin",
|
"awq_marlin",
|
||||||
"bitsandbytes",
|
"bitsandbytes",
|
||||||
|
|||||||
Reference in New Issue
Block a user