Skip to content

vllm.model_executor.kernels.linear.scaled_mm.BlockScaledMMLinearKernel

Classes:

FP8BlockParams dataclass

Bases: FP8Params

Attributes:

Source code in vllm/model_executor/kernels/linear/scaled_mm/BlockScaledMMLinearKernel.py
@dataclass
class FP8BlockParams(FP8Params):
    weight_scale_inv: torch.Tensor | None
    weight_scale: torch.Tensor | None

    WEIGHT_SCALE_INV: ClassVar[str] = "weight_scale_inv"

    @classmethod
    def from_layer(cls, layer: torch.nn.Module) -> Self:
        return cls(
            weight=getattr(layer, cls.WEIGHT),
            weight_scale_inv=getattr(layer, cls.WEIGHT_SCALE_INV, None),
            weight_scale=getattr(layer, cls.WEIGHT_SCALE, None),
            input_scale=getattr(layer, cls.INPUT_SCALE, None),
            input_scale_ub=getattr(layer, cls.INPUT_SCALE_UB, None),
        )

    @property
    def block_scale_attr(self) -> str:
        """Fp8LinearMethod registers the block scale as ``weight_scale_inv``,
        compressed-tensors as ``weight_scale``."""
        return (
            self.WEIGHT_SCALE
            if self.weight_scale_inv is None
            else self.WEIGHT_SCALE_INV
        )

    @property
    def block_scale(self) -> torch.Tensor:
        scale = getattr(self, self.block_scale_attr)
        assert scale is not None
        return scale

block_scale_attr property

Fp8LinearMethod registers the block scale as weight_scale_inv, compressed-tensors as weight_scale.

Fp8BlockScaledDynamicMMLinearKernel

Bases: Fp8BlockScaledMMLinearKernel, ABC

Dynamic FP8 block-scaled kernel that dispatches at runtime.

Extends Fp8BlockScaledMMLinearKernel to inherit apply_weights and overrides apply_block_scaled_mm to dispatch between two sub-kernels using torch.cond.

Subclasses must define

base_type: The primary kernel class. fallback_type: The fallback kernel class.

Source code in vllm/model_executor/kernels/linear/scaled_mm/BlockScaledMMLinearKernel.py
class Fp8BlockScaledDynamicMMLinearKernel(Fp8BlockScaledMMLinearKernel, ABC):
    """Dynamic FP8 block-scaled kernel that dispatches at runtime.

    Extends Fp8BlockScaledMMLinearKernel to inherit apply_weights and overrides
    apply_block_scaled_mm to dispatch between two sub-kernels using torch.cond.

    Subclasses must define:
        base_type:     The primary kernel class.
        fallback_type: The fallback kernel class.
    """

    base_type: ClassVar[type[Fp8BlockScaledMMLinearKernel]]
    fallback_type: ClassVar[type[Fp8BlockScaledMMLinearKernel]]

    def __init__(self, config: "FP8ScaledMMLinearLayerConfig") -> None:
        super().__init__(config)
        self.base = self.base_type(config)
        self.fallback = self.fallback_type(config)

    @classmethod
    def is_supported(
        cls, compute_capability: int | None = None
    ) -> tuple[bool, str | None]:
        is_base_supported, reason_1 = cls.base_type.is_supported(compute_capability)
        is_fallback_supported, reason_2 = cls.fallback_type.is_supported(
            compute_capability
        )
        if is_base_supported and is_fallback_supported:
            return True, None
        if not is_base_supported and not is_fallback_supported:
            return (
                False,
                f"base is not supported due to {reason_1}; "
                f"fallback is not supported due to {reason_2}",
            )
        if not is_base_supported:
            return False, f"base is not supported due to {reason_1}"
        return False, f"fallback is not supported due to {reason_2}"

    @classmethod
    def can_implement(
        cls, config: "FP8ScaledMMLinearLayerConfig"
    ) -> tuple[bool, str | None]:
        can_implement_base, reason_1 = cls.base_type.can_implement(config)
        can_implement_fallback, reason_2 = cls.fallback_type.can_implement(config)
        if can_implement_base and can_implement_fallback:
            return True, None
        if not can_implement_base and not can_implement_fallback:
            return (
                False,
                f"base cannot implement due to {reason_1}; "
                f"fallback cannot implement due to {reason_2}",
            )
        if not can_implement_base:
            return False, f"base cannot implement due to {reason_1}"
        return False, f"fallback cannot implement due to {reason_2}"