Bases: ABC, Generic[TInitInfo, TUpdateInfo]
Base class for weight transfer engines that handle transport of model weights from a trainer to inference workers.
This abstraction separates weight transfer transport logic from the worker implementation, allowing different backends (NCCL, CUDA IPC, RDMA[TODO]) to be plugged in.
Each engine owns its full weight-update lifecycle: start_weight_update, update_weights, and finish_weight_update. Layerwise reloading (used by checkpoint-format engines) is opted into per engine by running it inside start_weight_update/finish_weight_update. Engines that apply weights in place (e.g. sparse patches) leave those methods as no-ops.
Session lifecycle state (whether an update is active) is tracked by the worker, not the engine, so subclasses do not need to chain to super() in their lifecycle methods.
Subclasses should define
init_info_cls: Type of backend-specific initialization info update_info_cls: Type of backend-specific update info
Source code in vllm/distributed/weight_transfer/base.py
| class WeightTransferEngine(ABC, Generic[TInitInfo, TUpdateInfo]):
"""
Base class for weight transfer engines that handle transport of model weights
from a trainer to inference workers.
This abstraction separates weight transfer transport logic from the worker
implementation, allowing different backends (NCCL, CUDA IPC, RDMA[TODO]) to be
plugged in.
Each engine owns its full weight-update lifecycle: `start_weight_update`,
`update_weights`, and `finish_weight_update`. Layerwise reloading (used by
checkpoint-format engines) is opted into per engine by running it inside
`start_weight_update`/`finish_weight_update`. Engines that apply weights in
place (e.g. sparse patches) leave those methods as no-ops.
Session lifecycle state (whether an update is active) is tracked by the
worker, not the engine, so subclasses do not need to chain to `super()` in
their lifecycle methods.
Subclasses should define:
init_info_cls: Type of backend-specific initialization info
update_info_cls: Type of backend-specific update info
"""
# Subclasses should override these class attributes
init_info_cls: type[TInitInfo]
update_info_cls: type[TUpdateInfo]
def __init__(
self,
config: WeightTransferConfig,
vllm_config: "VllmConfig",
device: torch.device,
model: torch.nn.Module,
) -> None:
"""
Initialize the weight transfer engine.
Args:
config: The configuration for the weight transfer engine
vllm_config: The full vLLM config (provides parallel/model config and
is used to set the current config when running layerwise reload)
device: The device this worker's model lives on
model: The local model instance which will receive the weights
"""
self.config = config
self.vllm_config = vllm_config
self.parallel_config: ParallelConfig = vllm_config.parallel_config
self.model_config = vllm_config.model_config
self.device = device
self.model = model
def parse_init_info(self, init_dict: dict[str, Any]) -> TInitInfo:
"""
Construct typed init info from dict with validation.
Args:
init_dict: Dictionary containing backend-specific initialization parameters
Returns:
Typed backend-specific init info dataclass
Raises:
ValueError: If init_dict is invalid for this backend
"""
try:
return self.init_info_cls(**init_dict)
except TypeError as e:
raise ValueError(
f"Invalid init_info for {self.__class__.__name__}: {e}"
) from e
def parse_update_info(self, update_dict: dict[str, Any]) -> TUpdateInfo:
"""
Construct typed update info from dict with validation.
Args:
update_dict: Dictionary containing backend-specific update parameters
Returns:
Typed backend-specific update info dataclass
Raises:
ValueError: If update_dict is invalid for this backend
"""
try:
return self.update_info_cls(**update_dict)
except TypeError as e:
raise ValueError(
f"Invalid update_info for {self.__class__.__name__}: {e}"
) from e
@abstractmethod
def init_transfer_engine(self, init_info: TInitInfo) -> None:
"""
Initialize the weight transfer mechanism.
This is called once at the beginning of training.
Args:
init_info: Backend-specific initialization info
"""
raise NotImplementedError
@abstractmethod
def start_weight_update(self) -> None:
"""
Prepare the engine for a new weight update.
Checkpoint-format engines initialize layerwise reloading here; engines
that apply weights in place leave this as a no-op. Must not chain to
`super()`.
"""
raise NotImplementedError
@abstractmethod
def finish_weight_update(self) -> None:
"""
Finalize the current weight update.
Checkpoint-format engines finalize layerwise reloading here; engines
that apply weights in place leave this as a no-op. Must not chain to
`super()`.
"""
raise NotImplementedError
def update_weights(self, update_info: dict[str, Any]) -> None:
"""
Receive one weight update chunk and load it into the model.
This is stateless orchestration: parse the backend-specific update info,
receive the weights, then synchronize so the new weights are visible to
the next forward pass. Session-lifecycle bookkeeping is handled by the
worker.
Args:
update_info: Dictionary containing backend-specific update info
"""
typed_update_info = self.parse_update_info(update_info)
self.receive_weights(typed_update_info)
# NCCL broadcast / IPC paths may be asynchronous. Synchronize here so the
# next step uses the new weights.
torch.accelerator.synchronize()
@abstractmethod
def receive_weights(self, update_info: TUpdateInfo) -> None:
"""
Receive weights from the trainer and load them into the model.
Implementations should load weights incrementally (one or a few at a
time) into `self.model` to avoid OOM.
Args:
update_info: Backend-specific update info containing parameter metadata
and any backend-specific data
"""
raise NotImplementedError
@abstractmethod
def shutdown(self) -> None:
"""
Shutdown the weight transfer engine.
This should be called when the worker is shutting down.
"""
raise NotImplementedError
@staticmethod
@abstractmethod
def trainer_send_weights(
iterator: Iterator[Any],
trainer_args: dict[str, Any] | Any,
) -> None:
"""
Send weights from trainer to inference workers.
This is a static method that can be called from the trainer process
to send weights to all inference workers.
Args:
iterator: Iterator of backend-specific items to send. Dense engines
iterate (name, tensor) tuples; sparse engines iterate
patch objects. Tensors should be on the appropriate device.
trainer_args: Dictionary containing backend-specific arguments needed
to send weights. The structure depends on the backend:
- NCCL: Contains 'group', 'src', 'packed', etc.
- IPC: Contains 'mode' ('http' or 'ray'),
'llm_handle' (for Ray), 'url' (for HTTP), etc.
Example:
>>> param_iter = ((n, p) for n, p in model.named_parameters())
>>> engine.trainer_send_weights(param_iter, trainer_args)
"""
raise NotImplementedError
|
__init__
Initialize the weight transfer engine.
Parameters:
| Name | Type | Description | Default |
config | WeightTransferConfig | The configuration for the weight transfer engine | required |
vllm_config | VllmConfig | The full vLLM config (provides parallel/model config and is used to set the current config when running layerwise reload) | required |
device | device | The device this worker's model lives on | required |
model | Module | The local model instance which will receive the weights | required |
Source code in vllm/distributed/weight_transfer/base.py
| def __init__(
self,
config: WeightTransferConfig,
vllm_config: "VllmConfig",
device: torch.device,
model: torch.nn.Module,
) -> None:
"""
Initialize the weight transfer engine.
Args:
config: The configuration for the weight transfer engine
vllm_config: The full vLLM config (provides parallel/model config and
is used to set the current config when running layerwise reload)
device: The device this worker's model lives on
model: The local model instance which will receive the weights
"""
self.config = config
self.vllm_config = vllm_config
self.parallel_config: ParallelConfig = vllm_config.parallel_config
self.model_config = vllm_config.model_config
self.device = device
self.model = model
|
finish_weight_update abstractmethod
finish_weight_update() -> None
Finalize the current weight update.
Checkpoint-format engines finalize layerwise reloading here; engines that apply weights in place leave this as a no-op. Must not chain to super().
Source code in vllm/distributed/weight_transfer/base.py
| @abstractmethod
def finish_weight_update(self) -> None:
"""
Finalize the current weight update.
Checkpoint-format engines finalize layerwise reloading here; engines
that apply weights in place leave this as a no-op. Must not chain to
`super()`.
"""
raise NotImplementedError
|
init_transfer_engine abstractmethod
init_transfer_engine(init_info: TInitInfo) -> None
Initialize the weight transfer mechanism. This is called once at the beginning of training.
Parameters:
| Name | Type | Description | Default |
init_info | TInitInfo | Backend-specific initialization info | required |
Source code in vllm/distributed/weight_transfer/base.py
| @abstractmethod
def init_transfer_engine(self, init_info: TInitInfo) -> None:
"""
Initialize the weight transfer mechanism.
This is called once at the beginning of training.
Args:
init_info: Backend-specific initialization info
"""
raise NotImplementedError
|
parse_init_info
parse_init_info(init_dict: dict[str, Any]) -> TInitInfo
Construct typed init info from dict with validation.
Parameters:
| Name | Type | Description | Default |
init_dict | dict[str, Any] | Dictionary containing backend-specific initialization parameters | required |
Returns:
| Type | Description |
TInitInfo | Typed backend-specific init info dataclass |
Raises:
| Type | Description |
ValueError | If init_dict is invalid for this backend |
Source code in vllm/distributed/weight_transfer/base.py
| def parse_init_info(self, init_dict: dict[str, Any]) -> TInitInfo:
"""
Construct typed init info from dict with validation.
Args:
init_dict: Dictionary containing backend-specific initialization parameters
Returns:
Typed backend-specific init info dataclass
Raises:
ValueError: If init_dict is invalid for this backend
"""
try:
return self.init_info_cls(**init_dict)
except TypeError as e:
raise ValueError(
f"Invalid init_info for {self.__class__.__name__}: {e}"
) from e
|
parse_update_info
parse_update_info(
update_dict: dict[str, Any],
) -> TUpdateInfo
Construct typed update info from dict with validation.
Parameters:
| Name | Type | Description | Default |
update_dict | dict[str, Any] | Dictionary containing backend-specific update parameters | required |
Returns:
| Type | Description |
TUpdateInfo | Typed backend-specific update info dataclass |
Raises:
| Type | Description |
ValueError | If update_dict is invalid for this backend |
Source code in vllm/distributed/weight_transfer/base.py
| def parse_update_info(self, update_dict: dict[str, Any]) -> TUpdateInfo:
"""
Construct typed update info from dict with validation.
Args:
update_dict: Dictionary containing backend-specific update parameters
Returns:
Typed backend-specific update info dataclass
Raises:
ValueError: If update_dict is invalid for this backend
"""
try:
return self.update_info_cls(**update_dict)
except TypeError as e:
raise ValueError(
f"Invalid update_info for {self.__class__.__name__}: {e}"
) from e
|
receive_weights abstractmethod
receive_weights(update_info: TUpdateInfo) -> None
Receive weights from the trainer and load them into the model.
Implementations should load weights incrementally (one or a few at a time) into self.model to avoid OOM.
Parameters:
| Name | Type | Description | Default |
update_info | TUpdateInfo | Backend-specific update info containing parameter metadata and any backend-specific data | required |
Source code in vllm/distributed/weight_transfer/base.py
| @abstractmethod
def receive_weights(self, update_info: TUpdateInfo) -> None:
"""
Receive weights from the trainer and load them into the model.
Implementations should load weights incrementally (one or a few at a
time) into `self.model` to avoid OOM.
Args:
update_info: Backend-specific update info containing parameter metadata
and any backend-specific data
"""
raise NotImplementedError
|
shutdown abstractmethod
Shutdown the weight transfer engine. This should be called when the worker is shutting down.
Source code in vllm/distributed/weight_transfer/base.py
| @abstractmethod
def shutdown(self) -> None:
"""
Shutdown the weight transfer engine.
This should be called when the worker is shutting down.
"""
raise NotImplementedError
|
start_weight_update abstractmethod
start_weight_update() -> None
Prepare the engine for a new weight update.
Checkpoint-format engines initialize layerwise reloading here; engines that apply weights in place leave this as a no-op. Must not chain to super().
Source code in vllm/distributed/weight_transfer/base.py
| @abstractmethod
def start_weight_update(self) -> None:
"""
Prepare the engine for a new weight update.
Checkpoint-format engines initialize layerwise reloading here; engines
that apply weights in place leave this as a no-op. Must not chain to
`super()`.
"""
raise NotImplementedError
|
trainer_send_weights abstractmethod staticmethod
Send weights from trainer to inference workers.
This is a static method that can be called from the trainer process to send weights to all inference workers.
Parameters:
| Name | Type | Description | Default |
iterator | Iterator[Any] | Iterator of backend-specific items to send. Dense engines iterate (name, tensor) tuples; sparse engines iterate patch objects. Tensors should be on the appropriate device. | required |
trainer_args | dict[str, Any] | Any | Dictionary containing backend-specific arguments needed to send weights. The structure depends on the backend: - NCCL: Contains 'group', 'src', 'packed', etc. - IPC: Contains 'mode' ('http' or 'ray'), 'llm_handle' (for Ray), 'url' (for HTTP), etc. | required |
Example
param_iter = ((n, p) for n, p in model.named_parameters()) engine.trainer_send_weights(param_iter, trainer_args)
Source code in vllm/distributed/weight_transfer/base.py
| @staticmethod
@abstractmethod
def trainer_send_weights(
iterator: Iterator[Any],
trainer_args: dict[str, Any] | Any,
) -> None:
"""
Send weights from trainer to inference workers.
This is a static method that can be called from the trainer process
to send weights to all inference workers.
Args:
iterator: Iterator of backend-specific items to send. Dense engines
iterate (name, tensor) tuples; sparse engines iterate
patch objects. Tensors should be on the appropriate device.
trainer_args: Dictionary containing backend-specific arguments needed
to send weights. The structure depends on the backend:
- NCCL: Contains 'group', 'src', 'packed', etc.
- IPC: Contains 'mode' ('http' or 'ray'),
'llm_handle' (for Ray), 'url' (for HTTP), etc.
Example:
>>> param_iter = ((n, p) for n, p in model.named_parameters())
>>> engine.trainer_send_weights(param_iter, trainer_args)
"""
raise NotImplementedError
|
update_weights
update_weights(update_info: dict[str, Any]) -> None
Receive one weight update chunk and load it into the model.
This is stateless orchestration: parse the backend-specific update info, receive the weights, then synchronize so the new weights are visible to the next forward pass. Session-lifecycle bookkeeping is handled by the worker.
Parameters:
| Name | Type | Description | Default |
update_info | dict[str, Any] | Dictionary containing backend-specific update info | required |
Source code in vllm/distributed/weight_transfer/base.py
| def update_weights(self, update_info: dict[str, Any]) -> None:
"""
Receive one weight update chunk and load it into the model.
This is stateless orchestration: parse the backend-specific update info,
receive the weights, then synchronize so the new weights are visible to
the next forward pass. Session-lifecycle bookkeeping is handled by the
worker.
Args:
update_info: Dictionary containing backend-specific update info
"""
typed_update_info = self.parse_update_info(update_info)
self.receive_weights(typed_update_info)
# NCCL broadcast / IPC paths may be asynchronous. Synchronize here so the
# next step uses the new weights.
torch.accelerator.synchronize()
|