[3/N] Quantization Refactor: ModelSlim MoE schemes (#17993)

Co-authored-by: ronnie_zheng <zl19940307@163.com>
This commit is contained in:
Tamir Baydasov
2026-02-17 21:38:27 +03:00
committed by GitHub
co-authored by ronnie_zheng
parent 10569d04bb
commit aeca7d348c
7 changed files with 338 additions and 183 deletions
@@ -2,7 +2,7 @@ from __future__ import annotations
import logging import logging
from types import MappingProxyType from types import MappingProxyType
from typing import Any, Dict, List, Mapping, Optional, Tuple, Union, cast from typing import TYPE_CHECKING, Any, Dict, List, Mapping, Optional, Tuple, Union, cast
import torch import torch
@@ -10,19 +10,31 @@ from sglang.srt.hardware_backend.npu.quantization.linear_method_npu import (
_NPULinearMethodBase, _NPULinearMethodBase,
) )
from sglang.srt.layers.quantization.base_config import ( from sglang.srt.layers.quantization.base_config import (
FusedMoEMethodBase,
QuantizationConfig, QuantizationConfig,
QuantizeMethodBase,
) )
from sglang.srt.layers.quantization.compressed_tensors.utils import should_ignore_layer from sglang.srt.layers.quantization.compressed_tensors.utils import should_ignore_layer
from sglang.srt.layers.quantization.modelslim.modelslim_moe import ModelSlimMoEMethod
from sglang.srt.layers.quantization.modelslim.schemes import ( from sglang.srt.layers.quantization.modelslim.schemes import (
ModelSlimScheme,
ModelSlimW4A4Int4, ModelSlimW4A4Int4,
ModelSlimW4A8Int8MoE,
ModelSlimW8A8Int8, ModelSlimW8A8Int8,
ModelSlimW8A8Int8MoE,
) )
from sglang.srt.layers.quantization.unquant import UnquantizedLinearMethod from sglang.srt.layers.quantization.unquant import UnquantizedLinearMethod
from sglang.srt.utils import apply_module_patch from sglang.srt.utils import apply_module_patch
if TYPE_CHECKING:
from sglang.srt.layers.moe import MoeRunnerConfig
from sglang.srt.layers.moe.token_dispatcher import (
CombineInput,
StandardDispatchOutput,
)
from sglang.srt.layers.quantization.base_config import QuantizeMethodBase
from sglang.srt.layers.quantization.modelslim.schemes import (
ModelSlimLinearScheme,
ModelSlimMoEScheme,
)
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -150,17 +162,20 @@ class ModelSlimConfig(QuantizationConfig):
if self.is_layer_skipped(prefix, packed_modules_mapping_subset): if self.is_layer_skipped(prefix, packed_modules_mapping_subset):
return UnquantizedLinearMethod() return UnquantizedLinearMethod()
scheme = self.get_scheme(layer=layer, layer_name=prefix_in_quant_config) scheme = self.get_linear_scheme(
layer=layer, layer_name=prefix_in_quant_config
)
layer.scheme = scheme layer.scheme = scheme
return ModelSlimLinearMethod(self) return ModelSlimLinearMethod(self)
elif isinstance(layer, FusedMoE): elif isinstance(layer, FusedMoE):
return ModelSlimMoEMethod.get_moe_method(self, layer, prefix) layer.scheme = self.get_moe_scheme(layer, prefix)
return ModelSlimFusedMoEMethod(self)
return None return None
def _get_scheme_from_parts( def _get_scheme_from_parts(
self, self,
layer_name: str, layer_name: str,
) -> ModelSlimScheme: ) -> ModelSlimLinearScheme:
quant_type = self.quant_description.get(layer_name + ".weight", "") quant_type = self.quant_description.get(layer_name + ".weight", "")
if quant_type == "W8A8_DYNAMIC" or quant_type == "W8A8": if quant_type == "W8A8_DYNAMIC" or quant_type == "W8A8":
@@ -173,9 +188,9 @@ class ModelSlimConfig(QuantizationConfig):
) )
raise NotImplementedError("No modelslim compatible scheme was found.") raise NotImplementedError("No modelslim compatible scheme was found.")
def get_scheme( def get_linear_scheme(
self, layer: torch.nn.Module, layer_name: Optional[str] = None self, layer: torch.nn.Module, layer_name: Optional[str] = None
) -> Optional[ModelSlimScheme]: ) -> Optional[ModelSlimLinearScheme]:
""" """
get_scheme method adjusted for modelslim, taken from get_scheme method adjusted for modelslim, taken from
python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors.py python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors.py
@@ -188,6 +203,37 @@ class ModelSlimConfig(QuantizationConfig):
logger.debug("Using scheme: %s for %s", scheme.__class__.__name__, layer_name) logger.debug("Using scheme: %s for %s", scheme.__class__.__name__, layer_name)
return scheme return scheme
def get_moe_scheme(
self,
layer: torch.nn.Module,
prefix: str,
) -> Optional[ModelSlimMoEScheme]:
# TODO: @dsikka: refactor this to use schemes as other kernels
# are supported + check if the layer is being ignored.
prefix_in_quant_config = prefix + ".0.down_proj.weight"
is_moe_w4a8_dynamic = (
self.quant_description.get(prefix_in_quant_config, "STATIC")
== "W4A8_DYNAMIC"
)
is_moe_w8a8_dynamic = (
self.quant_description.get(prefix_in_quant_config, "STATIC")
== "W8A8_DYNAMIC"
)
if is_moe_w4a8_dynamic:
logger.info_once("Using ModelSlimW4A8Int8MoE")
return ModelSlimW4A8Int8MoE(self)
elif is_moe_w8a8_dynamic:
logger.info_once("Using ModelSlimW8A8Int8MoE")
return ModelSlimW8A8Int8MoE(self)
else:
logger.warning(
f"Unsupported FusedMoe modelslim scheme: "
f"{self.quant_description.get(prefix_in_quant_config.strip())} "
f"in layer: {prefix}"
)
return None
def is_layer_skipped( def is_layer_skipped(
self, prefix: str, fused_mapping: Mapping[str, List[str]] = MappingProxyType({}) self, prefix: str, fused_mapping: Mapping[str, List[str]] = MappingProxyType({})
): ):
@@ -251,7 +297,7 @@ class ModelSlimLinearMethod(_NPULinearMethodBase):
**extra_weight_attrs, **extra_weight_attrs,
): ):
""" """
Use the ModelSlimScheme associated with each layer to create Use the ModelSlimLinearScheme associated with the layer to create
the necessary parameters for the layer. See LinearMethodBase for param the necessary parameters for the layer. See LinearMethodBase for param
details details
""" """
@@ -273,7 +319,7 @@ class ModelSlimLinearMethod(_NPULinearMethodBase):
bias: Optional[torch.Tensor] = None, bias: Optional[torch.Tensor] = None,
): ):
""" """
Use the output of create_weights and the CompressedTensorsScheme Use the output of create_weights and the ModelSlimLinearScheme
associated with the layer to apply the forward pass with the associated with the layer to apply the forward pass with the
layer input. See LinearMethodBase for param details layer input. See LinearMethodBase for param details
@@ -283,3 +329,74 @@ class ModelSlimLinearMethod(_NPULinearMethodBase):
if scheme is None: if scheme is None:
raise ValueError("A scheme must be defined for each layer") raise ValueError("A scheme must be defined for each layer")
return scheme.apply_weights(layer, x, bias=bias) return scheme.apply_weights(layer, x, bias=bias)
class ModelSlimFusedMoEMethod(FusedMoEMethodBase):
def __init__(self, quantization_config: ModelSlimConfig):
self.quantization_config = quantization_config
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
layer.scheme.process_weights_after_loading(layer)
def create_weights(
self,
layer: torch.nn.Module,
num_experts: int,
hidden_size: int,
intermediate_size_per_partition: int,
params_dtype: torch.dtype,
**extra_weight_attrs,
):
"""
Use the ModelSlimMoEScheme associated with the layer to create
the necessary parameters for the layer. See FusedMoEMethodBase for param
details
"""
layer.scheme.create_weights(
layer=layer,
num_experts=num_experts,
hidden_size=hidden_size,
intermediate_size_per_partition=intermediate_size_per_partition,
params_dtype=params_dtype,
**extra_weight_attrs,
)
def create_moe_runner(
self, layer: torch.nn.Module, moe_runner_config: MoeRunnerConfig
):
return layer.scheme.create_moe_runner(layer, moe_runner_config)
def apply(
self,
layer: torch.nn.Module,
dispatch_output: StandardDispatchOutput,
) -> CombineInput:
"""
Use the output of create_weights and the ModelSlimMoEScheme
associated with the layer to apply the forward pass with the
layer input. See FusedMoEMethodBase for param details
"""
scheme = layer.scheme
if scheme is None:
raise ValueError("A scheme must be defined for each layer")
return scheme.apply_weights(layer, dispatch_output)
def apply_without_routing_weights(
self,
layer,
hidden_states,
hidden_states_scale,
group_list_type,
group_list,
output_dtype,
):
return layer.scheme.apply_without_routing_weights(
layer,
hidden_states,
hidden_states_scale,
group_list_type,
group_list,
output_dtype,
)
@@ -1,11 +1,16 @@
# SPDX-License-Identifier: Apache-2.0 # SPDX-License-Identifier: Apache-2.0
from .modelslim_scheme import ModelSlimScheme from .modelslim_scheme import ModelSlimLinearScheme, ModelSlimMoEScheme
from .modelslim_w4a4_int4 import ModelSlimW4A4Int4 from .modelslim_w4a4_int4 import ModelSlimW4A4Int4
from .modelslim_w4a8_int8_moe import ModelSlimW4A8Int8MoE
from .modelslim_w8a8_int8 import ModelSlimW8A8Int8 from .modelslim_w8a8_int8 import ModelSlimW8A8Int8
from .modelslim_w8a8_int8_moe import ModelSlimW8A8Int8MoE
__all__ = [ __all__ = [
"ModelSlimScheme", "ModelSlimLinearScheme",
"ModelSlimMoEScheme",
"ModelSlimW8A8Int8", "ModelSlimW8A8Int8",
"ModelSlimW4A4Int4", "ModelSlimW4A4Int4",
"ModelSlimW4A8Int8MoE",
"ModelSlimW8A8Int8MoE",
] ]
@@ -1,18 +1,24 @@
# Adapted from https://github.com/vllm-project/vllm/tree/main/vllm/model_executor/layers/quantization/compressed_tensors # Adapted from https://github.com/vllm-project/vllm/tree/main/vllm/model_executor/layers/quantization/compressed_tensors
# SPDX-License-Identifier: Apache-2.0 # SPDX-License-Identifier: Apache-2.0
from abc import ABC, abstractmethod from abc import abstractmethod
from typing import Optional from typing import TYPE_CHECKING, Optional
import torch import torch
__all__ = ["ModelSlimScheme"] from sglang.srt.layers.moe import MoeRunnerConfig
from sglang.srt.layers.quantization.base_scheme import BaseLinearScheme, BaseMoEScheme
if TYPE_CHECKING:
from sglang.srt.layers.moe.token_dispatcher import StandardDispatchOutput
__all__ = ["ModelSlimLinearScheme", "ModelSlimMoEScheme"]
class ModelSlimScheme(ABC): class ModelSlimLinearScheme(BaseLinearScheme):
""" """
Abstract class used to describe the weight creation and forward pass Abstract class used to describe the weight creation and forward pass
of different quantization schemes supported by CompressedTensors. of different quantization schemes supported by ModelSlim.
""" """
@abstractmethod @abstractmethod
@@ -23,6 +29,14 @@ class ModelSlimScheme(ABC):
""" """
raise NotImplementedError raise NotImplementedError
@abstractmethod
def process_weights_after_loading(self, layer: torch.nn.Module):
"""
Called after weight loading is complete for any cleanup that
needs to occur.
"""
raise NotImplementedError
@abstractmethod @abstractmethod
def apply_weights( def apply_weights(
self, layer: torch.nn.Module, x: torch.Tensor, bias: Optional[torch.Tensor] self, layer: torch.nn.Module, x: torch.Tensor, bias: Optional[torch.Tensor]
@@ -39,6 +53,21 @@ class ModelSlimScheme(ABC):
""" """
raise NotImplementedError raise NotImplementedError
class ModelSlimMoEScheme(BaseMoEScheme):
"""
Abstract class used to describe the weight creation and forward pass
of different quantization schemes supported by ModelSlim.
"""
@abstractmethod
def create_weights(self, *args, **kwargs):
"""
Weight creation for the particular scheme. Inputs to this function
"""
raise NotImplementedError
@abstractmethod @abstractmethod
def process_weights_after_loading(self, layer: torch.nn.Module): def process_weights_after_loading(self, layer: torch.nn.Module):
""" """
@@ -46,3 +75,26 @@ class ModelSlimScheme(ABC):
needs to occur. needs to occur.
""" """
raise NotImplementedError raise NotImplementedError
def create_moe_runner(
self, layer: torch.nn.Module, moe_runner_config: "MoeRunnerConfig"
):
raise NotImplementedError
@abstractmethod
def apply_weights(
self,
layer,
dispatch_output: "StandardDispatchOutput",
):
"""
Run the forward pass for the particular scheme. This is where
scheme-specific dequant/quant steps/kernels should be applied.
:param layer: torch.nn.Module with the registered weights and
other parameters relevant to the particular scheme.
:param x: input to the layer
:param bias: bias parameter
"""
raise NotImplementedError
@@ -9,11 +9,11 @@ from sglang.srt.hardware_backend.npu.quantization.linear_method_npu import (
NPU_W4A4DynamicLinearMethod, NPU_W4A4DynamicLinearMethod,
) )
from sglang.srt.layers.parameter import PerTensorScaleParameter from sglang.srt.layers.parameter import PerTensorScaleParameter
from sglang.srt.layers.quantization.modelslim.schemes import ModelSlimScheme from sglang.srt.layers.quantization.modelslim.schemes import ModelSlimLinearScheme
from sglang.srt.utils import set_weight_attrs from sglang.srt.utils import set_weight_attrs
class ModelSlimW4A4Int4(ModelSlimScheme): class ModelSlimW4A4Int4(ModelSlimLinearScheme):
def __init__( def __init__(
self, self,
@@ -1,5 +1,3 @@
# Adapted from https://github.com/vllm-project/vllm/tree/v0.8.2/vllm/model_executor/layers/quantization/compressed_tensors
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations from __future__ import annotations
import logging import logging
@@ -9,9 +7,8 @@ import torch
from sglang.srt.hardware_backend.npu.quantization.fused_moe_method_npu import ( from sglang.srt.hardware_backend.npu.quantization.fused_moe_method_npu import (
NPUW4A8Int8DynamicMoEMethod, NPUW4A8Int8DynamicMoEMethod,
NPUW8A8Int8DynamicMoEMethod,
) )
from sglang.srt.layers.quantization.base_config import FusedMoEMethodBase from sglang.srt.layers.quantization.modelslim.schemes import ModelSlimMoEScheme
from sglang.srt.utils import set_weight_attrs from sglang.srt.utils import set_weight_attrs
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -20,56 +17,15 @@ if TYPE_CHECKING:
CombineInput, CombineInput,
StandardDispatchOutput, StandardDispatchOutput,
) )
from sglang.srt.layers.quantization.modelslim.modelslim import ModelSlimConfig
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
__all__ = [ __all__ = [
"ModelSlimMoEMethod",
"ModelSlimW4A8Int8MoE", "ModelSlimW4A8Int8MoE",
"ModelSlimW8A8Int8MoE",
] ]
class ModelSlimMoEMethod(FusedMoEMethodBase): class ModelSlimW4A8Int8MoE(ModelSlimMoEScheme):
def __new__(cls, *args, **kwargs):
if cls is ModelSlimMoEMethod:
return super().__new__(cls)
return super().__new__(cls)
@staticmethod
def get_moe_method(
quant_config: ModelSlimConfig,
layer: torch.nn.Module,
prefix: str,
) -> "ModelSlimMoEMethod":
# TODO: @dsikka: refactor this to use schemes as other kernels
# are supported + check if the layer is being ignored.
prefix_in_quant_config = prefix + ".0.down_proj.weight"
is_moe_w4a8_dynamic = (
quant_config.quant_description.get(prefix_in_quant_config, "STATIC")
== "W4A8_DYNAMIC"
)
is_moe_w8a8_dynamic = (
quant_config.quant_description.get(prefix_in_quant_config, "STATIC")
== "W8A8_DYNAMIC"
)
if is_moe_w4a8_dynamic:
logger.info_once("Using ModelSlimW4A8Int8MoE")
return ModelSlimW4A8Int8MoE(quant_config)
elif is_moe_w8a8_dynamic:
logger.info_once("Using ModelSlimW8A8Int8MoE")
return ModelSlimW8A8Int8MoE(quant_config)
else:
logger.warning(f"Unsupported FusedMoe modelslim scheme: \
{quant_config.quant_description.get(prefix_in_quant_config.strip())} \
in layer: {prefix}")
return None
class ModelSlimW4A8Int8MoE(ModelSlimMoEMethod):
def __init__( def __init__(
self, self,
@@ -234,7 +190,7 @@ class ModelSlimW4A8Int8MoE(ModelSlimMoEMethod):
): ):
self.moe_runner_config = moe_runner_config self.moe_runner_config = moe_runner_config
def apply( def apply_weights(
self, self,
layer, layer,
dispatch_output: "StandardDispatchOutput", dispatch_output: "StandardDispatchOutput",
@@ -259,117 +215,3 @@ class ModelSlimW4A8Int8MoE(ModelSlimMoEMethod):
group_list, group_list,
output_dtype, output_dtype,
) )
class ModelSlimW8A8Int8MoE(ModelSlimMoEMethod):
def __init__(
self,
quant_config: Dict[str, Any],
prefix: str = None,
):
self.quant_config = quant_config
self.kernel = NPUW8A8Int8DynamicMoEMethod()
def create_weights(
self,
layer: torch.nn.Module,
num_experts: int,
hidden_size: int,
intermediate_size_per_partition: int,
params_dtype: torch.dtype,
**extra_weight_attrs,
) -> None:
from sglang.srt.layers.moe.fused_moe_triton import FusedMoeWeightScaleSupported
self.num_experts = num_experts
extra_weight_attrs.update(
{"quant_method": FusedMoeWeightScaleSupported.CHANNEL.value}
)
# weight
w13_weight = torch.nn.Parameter(
torch.empty(
num_experts,
2 * intermediate_size_per_partition,
hidden_size,
dtype=torch.int8,
),
requires_grad=False,
)
layer.register_parameter("w13_weight", w13_weight)
set_weight_attrs(w13_weight, extra_weight_attrs)
w2_weight = torch.nn.Parameter(
torch.empty(
num_experts,
hidden_size,
intermediate_size_per_partition,
dtype=torch.int8,
),
requires_grad=False,
)
layer.register_parameter("w2_weight", w2_weight)
set_weight_attrs(w2_weight, extra_weight_attrs)
# scale
w13_weight_scale = torch.nn.Parameter(
torch.empty(
num_experts, 2 * intermediate_size_per_partition, 1, dtype=torch.float32
),
requires_grad=False,
)
layer.register_parameter("w13_weight_scale", w13_weight_scale)
set_weight_attrs(w13_weight_scale, extra_weight_attrs)
w2_weight_scale = torch.nn.Parameter(
torch.empty(num_experts, hidden_size, 1, dtype=torch.float32),
requires_grad=False,
)
layer.register_parameter("w2_weight_scale", w2_weight_scale)
set_weight_attrs(w2_weight_scale, extra_weight_attrs)
# offset
w13_weight_offset = torch.nn.Parameter(
torch.empty(
num_experts, 2 * intermediate_size_per_partition, 1, dtype=torch.float32
),
requires_grad=False,
)
layer.register_parameter("w13_weight_offset", w13_weight_offset)
set_weight_attrs(w13_weight_offset, extra_weight_attrs)
w2_weight_offset = torch.nn.Parameter(
torch.empty(num_experts, hidden_size, 1, dtype=torch.float32),
requires_grad=False,
)
layer.register_parameter("w2_weight_offset", w2_weight_offset)
set_weight_attrs(w2_weight_offset, extra_weight_attrs)
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
self.kernel.process_weights_after_loading(layer)
def create_moe_runner(
self, layer: torch.nn.Module, moe_runner_config: "MoeRunnerConfig"
):
self.moe_runner_config = moe_runner_config
def apply(
self,
layer,
dispatch_output: "StandardDispatchOutput",
) -> "CombineInput":
return self.kernel.apply(layer, dispatch_output)
def apply_without_routing_weights(
self,
layer,
hidden_states,
hidden_states_scale,
group_list_type,
group_list,
output_dtype,
):
return self.kernel.apply_without_routing_weights(
layer,
hidden_states,
hidden_states_scale,
group_list_type,
group_list,
output_dtype,
)
@@ -14,10 +14,10 @@ from sglang.srt.layers.parameter import (
ModelWeightParameter, ModelWeightParameter,
PerTensorScaleParameter, PerTensorScaleParameter,
) )
from sglang.srt.layers.quantization.modelslim.schemes import ModelSlimScheme from sglang.srt.layers.quantization.modelslim.schemes import ModelSlimLinearScheme
class ModelSlimW8A8Int8(ModelSlimScheme): class ModelSlimW8A8Int8(ModelSlimLinearScheme):
def __init__( def __init__(
self, self,
@@ -0,0 +1,139 @@
from __future__ import annotations
import logging
from typing import TYPE_CHECKING, Any, Dict
import torch
from sglang.srt.hardware_backend.npu.quantization.fused_moe_method_npu import (
NPUW8A8Int8DynamicMoEMethod,
)
from sglang.srt.layers.quantization.modelslim.schemes import ModelSlimMoEScheme
from sglang.srt.utils import set_weight_attrs
if TYPE_CHECKING:
from sglang.srt.layers.moe import MoeRunnerConfig
from sglang.srt.layers.moe.token_dispatcher import (
CombineInput,
StandardDispatchOutput,
)
logger = logging.getLogger(__name__)
__all__ = [
"ModelSlimW8A8Int8MoE",
]
class ModelSlimW8A8Int8MoE(ModelSlimMoEScheme):
def __init__(
self,
quant_config: Dict[str, Any],
prefix: str = None,
):
self.quant_config = quant_config
self.kernel = NPUW8A8Int8DynamicMoEMethod()
def create_weights(
self,
layer: torch.nn.Module,
num_experts: int,
hidden_size: int,
intermediate_size_per_partition: int,
params_dtype: torch.dtype,
**extra_weight_attrs,
) -> None:
from sglang.srt.layers.moe.fused_moe_triton import FusedMoeWeightScaleSupported
self.num_experts = num_experts
extra_weight_attrs.update(
{"quant_method": FusedMoeWeightScaleSupported.CHANNEL.value}
)
# weight
w13_weight = torch.nn.Parameter(
torch.empty(
num_experts,
2 * intermediate_size_per_partition,
hidden_size,
dtype=torch.int8,
),
requires_grad=False,
)
layer.register_parameter("w13_weight", w13_weight)
set_weight_attrs(w13_weight, extra_weight_attrs)
w2_weight = torch.nn.Parameter(
torch.empty(
num_experts,
hidden_size,
intermediate_size_per_partition,
dtype=torch.int8,
),
requires_grad=False,
)
layer.register_parameter("w2_weight", w2_weight)
set_weight_attrs(w2_weight, extra_weight_attrs)
# scale
w13_weight_scale = torch.nn.Parameter(
torch.empty(
num_experts, 2 * intermediate_size_per_partition, 1, dtype=torch.float32
),
requires_grad=False,
)
layer.register_parameter("w13_weight_scale", w13_weight_scale)
set_weight_attrs(w13_weight_scale, extra_weight_attrs)
w2_weight_scale = torch.nn.Parameter(
torch.empty(num_experts, hidden_size, 1, dtype=torch.float32),
requires_grad=False,
)
layer.register_parameter("w2_weight_scale", w2_weight_scale)
set_weight_attrs(w2_weight_scale, extra_weight_attrs)
# offset
w13_weight_offset = torch.nn.Parameter(
torch.empty(
num_experts, 2 * intermediate_size_per_partition, 1, dtype=torch.float32
),
requires_grad=False,
)
layer.register_parameter("w13_weight_offset", w13_weight_offset)
set_weight_attrs(w13_weight_offset, extra_weight_attrs)
w2_weight_offset = torch.nn.Parameter(
torch.empty(num_experts, hidden_size, 1, dtype=torch.float32),
requires_grad=False,
)
layer.register_parameter("w2_weight_offset", w2_weight_offset)
set_weight_attrs(w2_weight_offset, extra_weight_attrs)
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
self.kernel.process_weights_after_loading(layer)
def create_moe_runner(
self, layer: torch.nn.Module, moe_runner_config: "MoeRunnerConfig"
):
self.moe_runner_config = moe_runner_config
def apply_weights(
self,
layer,
dispatch_output: "StandardDispatchOutput",
) -> "CombineInput":
return self.kernel.apply(layer, dispatch_output)
def apply_without_routing_weights(
self,
layer,
hidden_states,
hidden_states_scale,
group_list_type,
group_list,
output_dtype,
):
return self.kernel.apply_without_routing_weights(
layer,
hidden_states,
hidden_states_scale,
group_list_type,
group_list,
output_dtype,
)