Skip to content

vllm.model_executor.models.transformers.fusers.base

Base classes for the Transformers backend fusers.

Classes:

  • BaseFuser

    A detected fusion and how to apply it.

  • RewriteFuser

    A fuser that rewrites the module's forward and rebinds it.

  • StackedFuser

    A fuser that merges sibling projections into one stacked linear and

Functions:

  • local_output_sizes

    Source for the per-rank widths of the merged linear self.<merged_name>.

BaseFuser dataclass

Bases: ABC

A detected fusion and how to apply it.

match analyses the module class once (cached, see get_fuser); fuse then applies the fusion to an instance in recursive_replace, returning the module to install in its place.

Methods:

  • fuse

    Apply the fusion to an already-validated module, returning the

  • info

    A human-readable description of the fusion at name, for logging.

  • match

    Match the pattern in graph, returning a fuser if found.

  • orig_to_new_stacked

    WeightsMapper.orig_to_new_stacked entries this fuser contributes

  • validate

    Whether this fuser can be applied to this module instance.

Attributes:

Source code in vllm/model_executor/models/transformers/fusers/base.py
@dataclass
class BaseFuser(ABC):
    """A detected fusion and how to apply it.

    `match` analyses the module *class* once (cached, see `get_fuser`); `fuse`
    then applies the fusion to an instance in `recursive_replace`, returning the
    module to install in its place.
    """

    @abstractmethod
    def info(self, name: str) -> str:
        """A human-readable description of the fusion at `name`, for logging."""

    @classmethod
    @abstractmethod
    def match(cls, graph: fx.Graph, module: nn.Module) -> "BaseFuser | None":
        """Match the pattern in `graph`, returning a fuser if found."""

    @abstractmethod
    def validate(self, module: nn.Module, vllm_config: "VllmConfig") -> bool:
        """Whether this fuser can be applied to this `module` instance."""

    @abstractmethod
    def fuse(
        self, module: nn.Module, prefix: str, vllm_config: "VllmConfig"
    ) -> nn.Module:
        """Apply the fusion to an already-validated `module`, returning the
        module to install in its place (mutated in place, or freshly built)."""

    def orig_to_new_stacked(self, prefix: str) -> dict[str, tuple[str, ShardId]]:
        """`WeightsMapper.orig_to_new_stacked` entries this fuser contributes
        (none unless it stacks weights)."""
        return {}

    @property
    def packed_modules_mapping(self) -> dict[str, list[str]]:
        """`packed_modules_mapping` entries this fuser contributes (none unless
        it stacks weights)."""
        return {}

packed_modules_mapping property

packed_modules_mapping entries this fuser contributes (none unless it stacks weights).

fuse(module, prefix, vllm_config) abstractmethod

Apply the fusion to an already-validated module, returning the module to install in its place (mutated in place, or freshly built).

Source code in vllm/model_executor/models/transformers/fusers/base.py
@abstractmethod
def fuse(
    self, module: nn.Module, prefix: str, vllm_config: "VllmConfig"
) -> nn.Module:
    """Apply the fusion to an already-validated `module`, returning the
    module to install in its place (mutated in place, or freshly built)."""

info(name) abstractmethod

A human-readable description of the fusion at name, for logging.

Source code in vllm/model_executor/models/transformers/fusers/base.py
@abstractmethod
def info(self, name: str) -> str:
    """A human-readable description of the fusion at `name`, for logging."""

match(graph, module) abstractmethod classmethod

Match the pattern in graph, returning a fuser if found.

Source code in vllm/model_executor/models/transformers/fusers/base.py
@classmethod
@abstractmethod
def match(cls, graph: fx.Graph, module: nn.Module) -> "BaseFuser | None":
    """Match the pattern in `graph`, returning a fuser if found."""

orig_to_new_stacked(prefix)

WeightsMapper.orig_to_new_stacked entries this fuser contributes (none unless it stacks weights).

Source code in vllm/model_executor/models/transformers/fusers/base.py
def orig_to_new_stacked(self, prefix: str) -> dict[str, tuple[str, ShardId]]:
    """`WeightsMapper.orig_to_new_stacked` entries this fuser contributes
    (none unless it stacks weights)."""
    return {}

validate(module, vllm_config) abstractmethod

Whether this fuser can be applied to this module instance.

Source code in vllm/model_executor/models/transformers/fusers/base.py
@abstractmethod
def validate(self, module: nn.Module, vllm_config: "VllmConfig") -> bool:
    """Whether this fuser can be applied to this `module` instance."""

RewriteFuser dataclass

Bases: BaseFuser

A fuser that rewrites the module's forward and rebinds it.

match and update_forward analyse the class once; fuse swaps the submodules and binds the compiled forward on an instance in place, so it keeps its class and any attribute the fusion does not consume.

Methods:

  • fuse

    Fuse an already-validated module in place (see Fusers.__getitem__).

  • update_attrs

    Replace module's submodules with their vLLM equivalents.

  • update_forward

    Rewrite and compile type(module)'s forward source.

Attributes:

  • fused_forward (Callable) –

    The compiled rewritten forward, set by update_forward.

  • source_cls (str) –

    Class of the HF module the fused projections belonged to (for logging).

Source code in vllm/model_executor/models/transformers/fusers/base.py
@dataclass
class RewriteFuser(BaseFuser):
    """A fuser that rewrites the module's forward and rebinds it.

    `match` and `update_forward` analyse the class once; `fuse` swaps the
    submodules and binds the compiled forward on an instance in place, so it
    keeps its class and any attribute the fusion does not consume.
    """

    source_cls: str
    """Class of the HF module the fused projections belonged to (for logging)."""

    fused_forward: Callable = field(init=False, repr=False)
    """The compiled rewritten forward, set by `update_forward`."""

    @abstractmethod
    def update_forward(self, module: nn.Module) -> None:
        """Rewrite and compile `type(module)`'s forward source.

        Raises if the source does not admit the rewrite (fusion is then skipped).
        """

    @abstractmethod
    def update_attrs(
        self, module: nn.Module, prefix: str, vllm_config: "VllmConfig"
    ) -> None:
        """Replace `module`'s submodules with their vLLM equivalents."""

    def fuse(
        self, module: nn.Module, prefix: str, vllm_config: "VllmConfig"
    ) -> nn.Module:
        """Fuse an already-validated `module` in place (see `Fusers.__getitem__`).

        Builds the merged submodule and binds the compiled forward."""
        self.update_attrs(module, prefix, vllm_config)
        module.forward = types.MethodType(self.fused_forward, module)
        return module

fused_forward = field(init=False, repr=False) class-attribute instance-attribute

The compiled rewritten forward, set by update_forward.

source_cls instance-attribute

Class of the HF module the fused projections belonged to (for logging).

fuse(module, prefix, vllm_config)

Fuse an already-validated module in place (see Fusers.__getitem__).

Builds the merged submodule and binds the compiled forward.

Source code in vllm/model_executor/models/transformers/fusers/base.py
def fuse(
    self, module: nn.Module, prefix: str, vllm_config: "VllmConfig"
) -> nn.Module:
    """Fuse an already-validated `module` in place (see `Fusers.__getitem__`).

    Builds the merged submodule and binds the compiled forward."""
    self.update_attrs(module, prefix, vllm_config)
    module.forward = types.MethodType(self.fused_forward, module)
    return module

update_attrs(module, prefix, vllm_config) abstractmethod

Replace module's submodules with their vLLM equivalents.

Source code in vllm/model_executor/models/transformers/fusers/base.py
@abstractmethod
def update_attrs(
    self, module: nn.Module, prefix: str, vllm_config: "VllmConfig"
) -> None:
    """Replace `module`'s submodules with their vLLM equivalents."""

update_forward(module) abstractmethod

Rewrite and compile type(module)'s forward source.

Raises if the source does not admit the rewrite (fusion is then skipped).

Source code in vllm/model_executor/models/transformers/fusers/base.py
@abstractmethod
def update_forward(self, module: nn.Module) -> None:
    """Rewrite and compile `type(module)`'s forward source.

    Raises if the source does not admit the rewrite (fusion is then skipped).
    """

StackedFuser dataclass

Bases: RewriteFuser

A fuser that merges sibling projections into one stacked linear and rewrites the forward to call it.

Methods:

Attributes:

Source code in vllm/model_executor/models/transformers/fusers/base.py
@dataclass
class StackedFuser(RewriteFuser):
    """A fuser that merges sibling projections into one stacked linear and
    rewrites the forward to call it."""

    merged_name: ClassVar[str]
    """Attribute name of the merged module created by `update_attrs`."""
    merged_cls: ClassVar[str]
    """Name of the vLLM class the merged projection becomes (for logging)."""

    def info(self, name: str) -> str:
        sources = " + ".join(shard for shard, _ in self.shards)
        return (
            f"Fused: {sources} ({name}: {self.source_cls}) -> "
            f"{self.merged_name} ({self.merged_cls})"
        )

    @property
    @abstractmethod
    def shards(self) -> list[tuple[str, ShardId]]:
        """Each projection's original name and its shard id in the merged module.

        Source for both `orig_to_new_stacked` and `packed_modules_mapping`."""

    def orig_to_new_stacked(self, prefix: str) -> dict[str, tuple[str, ShardId]]:
        """`WeightsMapper.orig_to_new_stacked` entries for one fused instance.

        Maps each checkpoint name to `(merged_name, shard_id)`, keyed by qualname
        so only this exact layer is remapped, never a same-named projection
        elsewhere (e.g. an unfused MoE expert's `gate_proj`)."""
        merged = maybe_prefix(prefix, self.merged_name)
        return {
            maybe_prefix(prefix, name): (merged, shard) for name, shard in self.shards
        }

    @property
    def packed_modules_mapping(self) -> dict[str, list[str]]:
        """`{merged_name: [projection names]}` so quantization can unpack the
        fused layer into its per-shard configs."""
        return {self.merged_name: [name for name, _ in self.shards]}

merged_cls class-attribute

Name of the vLLM class the merged projection becomes (for logging).

merged_name class-attribute

Attribute name of the merged module created by update_attrs.

packed_modules_mapping property

{merged_name: [projection names]} so quantization can unpack the fused layer into its per-shard configs.

shards abstractmethod property

Each projection's original name and its shard id in the merged module.

Source for both orig_to_new_stacked and packed_modules_mapping.

orig_to_new_stacked(prefix)

WeightsMapper.orig_to_new_stacked entries for one fused instance.

Maps each checkpoint name to (merged_name, shard_id), keyed by qualname so only this exact layer is remapped, never a same-named projection elsewhere (e.g. an unfused MoE expert's gate_proj).

Source code in vllm/model_executor/models/transformers/fusers/base.py
def orig_to_new_stacked(self, prefix: str) -> dict[str, tuple[str, ShardId]]:
    """`WeightsMapper.orig_to_new_stacked` entries for one fused instance.

    Maps each checkpoint name to `(merged_name, shard_id)`, keyed by qualname
    so only this exact layer is remapped, never a same-named projection
    elsewhere (e.g. an unfused MoE expert's `gate_proj`)."""
    merged = maybe_prefix(prefix, self.merged_name)
    return {
        maybe_prefix(prefix, name): (merged, shard) for name, shard in self.shards
    }

local_output_sizes(merged_name)

Source for the per-rank widths of the merged linear self.<merged_name>.

Source code in vllm/model_executor/models/transformers/fusers/base.py
def local_output_sizes(merged_name: str) -> str:
    """Source for the per-rank widths of the merged linear `self.<merged_name>`."""
    merged = f"self.{merged_name}"
    return f"[s // {merged}.tp_size for s in {merged}.output_sizes]"