Skip to content

vllm.distributed.weight_transfer.m2n_source

Trainer-side weight sources for the NCCL M2N backend.

m2n plans a transfer from both sides' layouts, so the trainer has to say how it holds each parameter — which the base WeightSource / ParamMeta pair does not express. This module supplies the m2n flavor of both, and the DTensor-backed implementation that covers the common trainer.

A source declares its mesh once (mesh()) and a placement per parameter, since one topology describes every tensor on the side.

Classes:

Functions:

DTensorModuleSource

Bases: M2NWeightSource

M2NWeightSource over module.named_parameters().

Covers both the FSDP/DTensor trainer (placement read off each parameter, local shard yielded via to_local()) and the plain replicated trainer, with no special casing. Trainers with a custom producer (a Megatron export, MoE re-fusing) subclass M2NWeightSource instead.

Methods:

  • __init__ –

    Wrap module; num_trainer_ranks sizes the mesh for plain tensors.

  • __iter__ –

    Yield each parameter's local shard (to_local()), or the tensor itself.

  • mesh –

    The mesh every sharded parameter agrees on.

  • metadata –

    Read global shape/dtype and placements without gathering shards.

Source code in vllm/distributed/weight_transfer/m2n_source.py
class DTensorModuleSource(M2NWeightSource):
    """`M2NWeightSource` over `module.named_parameters()`.

    Covers both the FSDP/DTensor trainer (placement read off each parameter,
    local shard yielded via `to_local()`) and the plain replicated trainer, with
    no special casing. Trainers with a custom producer (a Megatron export, MoE
    re-fusing) subclass `M2NWeightSource` instead.
    """

    def __init__(self, module: torch.nn.Module, num_trainer_ranks: int) -> None:
        """Wrap `module`; `num_trainer_ranks` sizes the mesh for plain tensors."""
        self._module = module
        self._num_trainer_ranks = num_trainer_ranks

    def mesh(self) -> M2NMesh:
        """The mesh every sharded parameter agrees on.

        Replicated parameters span the trainer without a device mesh of their
        own, so they never decide this; a model whose *sharded* parameters
        disagree cannot be described by one side mesh and is rejected.
        """
        meshes = {
            mesh_from_tensor(p, self._num_trainer_ranks)
            for _, p in self._module.named_parameters()
            if placements_from_tensor(p) is not REPLICATED
        }
        if len(meshes) > 1:
            raise ValueError(
                "nccl_m2n needs one mesh per side, but this module's sharded "
                f"parameters span several: {sorted(m.dims for m in meshes)}"
            )
        return meshes.pop() if meshes else M2NMesh((self._num_trainer_ranks, 1), 0)

    def metadata(self) -> list[ParamMeta]:
        """Read global shape/dtype and placements without gathering shards."""
        return [
            M2NParamMeta(name, p.dtype, tuple(p.shape), placements_from_tensor(p))
            for name, p in self._module.named_parameters()
        ]

    def __iter__(self) -> Iterator[tuple[str, torch.Tensor]]:
        """Yield each parameter's local shard (`to_local()`), or the tensor itself."""
        for name, param in self._module.named_parameters():
            to_local = getattr(param, "to_local", None)
            yield name, (to_local() if callable(to_local) else param)

__init__(module, num_trainer_ranks)

Wrap module; num_trainer_ranks sizes the mesh for plain tensors.

Source code in vllm/distributed/weight_transfer/m2n_source.py
def __init__(self, module: torch.nn.Module, num_trainer_ranks: int) -> None:
    """Wrap `module`; `num_trainer_ranks` sizes the mesh for plain tensors."""
    self._module = module
    self._num_trainer_ranks = num_trainer_ranks

__iter__()

Yield each parameter's local shard (to_local()), or the tensor itself.

Source code in vllm/distributed/weight_transfer/m2n_source.py
def __iter__(self) -> Iterator[tuple[str, torch.Tensor]]:
    """Yield each parameter's local shard (`to_local()`), or the tensor itself."""
    for name, param in self._module.named_parameters():
        to_local = getattr(param, "to_local", None)
        yield name, (to_local() if callable(to_local) else param)

mesh()

The mesh every sharded parameter agrees on.

Replicated parameters span the trainer without a device mesh of their own, so they never decide this; a model whose sharded parameters disagree cannot be described by one side mesh and is rejected.

Source code in vllm/distributed/weight_transfer/m2n_source.py
def mesh(self) -> M2NMesh:
    """The mesh every sharded parameter agrees on.

    Replicated parameters span the trainer without a device mesh of their
    own, so they never decide this; a model whose *sharded* parameters
    disagree cannot be described by one side mesh and is rejected.
    """
    meshes = {
        mesh_from_tensor(p, self._num_trainer_ranks)
        for _, p in self._module.named_parameters()
        if placements_from_tensor(p) is not REPLICATED
    }
    if len(meshes) > 1:
        raise ValueError(
            "nccl_m2n needs one mesh per side, but this module's sharded "
            f"parameters span several: {sorted(m.dims for m in meshes)}"
        )
    return meshes.pop() if meshes else M2NMesh((self._num_trainer_ranks, 1), 0)

metadata()

Read global shape/dtype and placements without gathering shards.

Source code in vllm/distributed/weight_transfer/m2n_source.py
def metadata(self) -> list[ParamMeta]:
    """Read global shape/dtype and placements without gathering shards."""
    return [
        M2NParamMeta(name, p.dtype, tuple(p.shape), placements_from_tensor(p))
        for name, p in self._module.named_parameters()
    ]

M2NWeightSource

Bases: WeightSource

A WeightSource that also describes how the trainer holds its weights.

Unlike ModuleSource, iteration yields each rank's local shard, not a materialized full tensor: gathering is exactly the cost m2n removes.

Methods:

  • __iter__ –

    Yield (name, local shard) pairs in the same order as metadata().

  • mesh –

    The trainer's rank topology, shared by every parameter.

  • metadata –

    Name, dtype, full shape, and trainer placement for each parameter.

Source code in vllm/distributed/weight_transfer/m2n_source.py
class M2NWeightSource(WeightSource):
    """A `WeightSource` that also describes how the trainer holds its weights.

    Unlike `ModuleSource`, iteration yields each rank's **local shard**, not a
    materialized full tensor: gathering is exactly the cost m2n removes.
    """

    def mesh(self) -> M2NMesh:
        """The trainer's rank topology, shared by every parameter."""
        raise NotImplementedError

    def metadata(self) -> list[ParamMeta]:
        """Name, dtype, full shape, and trainer placement for each parameter."""
        raise NotImplementedError

    def __iter__(self) -> Iterator[tuple[str, torch.Tensor]]:
        """Yield `(name, local shard)` pairs in the same order as `metadata()`."""
        raise NotImplementedError

__iter__()

Yield (name, local shard) pairs in the same order as metadata().

Source code in vllm/distributed/weight_transfer/m2n_source.py
def __iter__(self) -> Iterator[tuple[str, torch.Tensor]]:
    """Yield `(name, local shard)` pairs in the same order as `metadata()`."""
    raise NotImplementedError

mesh()

The trainer's rank topology, shared by every parameter.

Source code in vllm/distributed/weight_transfer/m2n_source.py
def mesh(self) -> M2NMesh:
    """The trainer's rank topology, shared by every parameter."""
    raise NotImplementedError

metadata()

Name, dtype, full shape, and trainer placement for each parameter.

Source code in vllm/distributed/weight_transfer/m2n_source.py
def metadata(self) -> list[ParamMeta]:
    """Name, dtype, full shape, and trainer placement for each parameter."""
    raise NotImplementedError

_placement_code(placement)

Map a torch.distributed placement onto an m2n placement code.

Source code in vllm/distributed/weight_transfer/m2n_source.py
def _placement_code(placement: Any) -> int:
    """Map a `torch.distributed` placement onto an m2n placement code."""
    name = type(placement).__name__
    if name == "Replicate":
        return REPLICATE
    if name == "Shard":
        return int(placement.dim)
    raise ValueError(
        f"nccl_m2n cannot express the {name} placement; only Replicate and "
        "Shard are supported"
    )

mesh_from_tensor(tensor, num_trainer_ranks)

The mesh a parameter lives on, as an M2NMesh.

A plain tensor is identical on every trainer rank, so it spans all of them.

Source code in vllm/distributed/weight_transfer/m2n_source.py
def mesh_from_tensor(tensor: torch.Tensor, num_trainer_ranks: int) -> M2NMesh:
    """The mesh a parameter lives on, as an `M2NMesh`.

    A plain tensor is identical on every trainer rank, so it spans all of them.
    """
    device_mesh = getattr(tensor, "device_mesh", None)
    if device_mesh is None:
        return M2NMesh((num_trainer_ranks, 1), 0)

    grid = device_mesh.mesh
    ranks = grid.flatten().tolist()
    if ranks != list(range(len(ranks))):
        raise ValueError(
            "nccl_m2n requires the trainer's device mesh to cover the "
            f"contiguous rank interval [0, {len(ranks)}); got {ranks}"
        )
    if grid.ndim > 2:
        raise ValueError(
            f"nccl_m2n supports 1-D and 2-D device meshes, got {grid.ndim}-D"
        )
    dims = tuple(grid.shape)
    return M2NMesh(dims if len(dims) == 2 else (dims[0], 1), 0)

placements_from_tensor(tensor)

How a parameter is placed over its mesh, or REPLICATED.

REPLICATED is returned rather than a (REPLICATE, REPLICATE) pair: m2n cannot express that, and the size-1-axis encoding it needs instead is applied later by resolve_layout, which owns that workaround.

Source code in vllm/distributed/weight_transfer/m2n_source.py
def placements_from_tensor(tensor: torch.Tensor) -> Placements | None:
    """How a parameter is placed over its mesh, or `REPLICATED`.

    `REPLICATED` is returned rather than a `(REPLICATE, REPLICATE)` pair: m2n
    cannot express that, and the size-1-axis encoding it needs instead is
    applied later by `resolve_layout`, which owns that workaround.
    """
    placements = getattr(tensor, "placements", None)
    if placements is None:
        return REPLICATED
    codes = [_placement_code(p) for p in placements]
    if all(code == REPLICATE for code in codes):
        return REPLICATED
    if len(codes) == 1:
        return (codes[0], REPLICATE)
    return (codes[0], codes[1])