Base Classes와 커스텀 엔진
Base Classes와 커스텀 엔진 (Base Classes and Custom Engines)
가중치 전송 시스템의 각 조각은 서로 독립적으로 교체할 수 있게 설계됐어요. 이 문서는 그 네 가지 추상화(abstraction)가 무엇이고, 직접 커스텀 엔진이나 소스를 만들 때 어떤 불변식을 지켜야 하는지 설명해요.
네 가지 추상화
가중치 전송 시스템은 네 가지 추상화로 만들어지며, 각각 독립적으로 교체 가능해요.
| 추상화 | 쪽 | 담당하는 질문 |
|---|---|---|
WeightSource |
트레이너 | 어떤 가중치를 보낼지 |
VLLMWeightSyncClient |
트레이너 | 추론 엔진에 어떻게 도달할지 — RL 스택의 자체 vLLM 래퍼용 어댑터 |
TrainerWeightTransferEngine |
트레이너 | 바이트를 어떻게 전송할지 |
WeightTransferEngine |
추론 | 어떻게 받아서 로드할지 |
두 엔진은 WeightTransferTrainerFactory와 WeightTransferEngineFactory라는 두 개의 별도 팩토리에 등록돼요. 백엔드 이름을 관례상 공유하지만, 트레이너 프로세스는 워커 엔진을, 워커는 트레이너 엔진을 인스턴스화하지 않으므로 레지스트리는 독립적으로 유지됩니다.
트레이너 쪽 (Trainer Side)
WeightSource
이것은 트레이너 가중치가 어떤 모양이든 그에 맞추는 어댑터예요. WeightSource는 특정 프레임워크에 맞춰 가중치를 어떻게 추출할지 정의합니다.
엔진이 받는 것은 항상 같아요. HF 형식 파라미터 이름과, 이미 전체(언샤딩) 모양으로 실체화된 텐서입니다. 그ather, re-fuse, dequantize, renaming 같은 작업은 모두 소스 안에서 일어나야 해요.
WeightSource는 재순회 가능(re-iterable) 하고 두 개의 필수 채널이 있어요.
metadata() -> list[ParamMeta]— 아무것도 전송하지 않고 모든 파라미터의 이름, wire dtype, 전체 모양을 반환. 모양을 로컬에서 알 수 있으면 저렴하고(FSDPDTensor는 전역 모양을 앎), 실체화해야 모양을 알 수 있는 생산자(Megatron-Bridge export)의 경우 첫 호출에서 비쌀 수 있으며 이때는 캐시해야 해요.- 순회 (iteration) — 완전히 실체화된
(name, tensor)쌍을 하나씩 산출.
@dataclass(frozen=True)
class ParamMeta:
name: str
dtype: torch.dtype
shape: tuple[int, ...]
두 채널은 요소별로 일치해야 해요. metadata()는 순회가 정확히 무엇을 산출할지 선언해야 합니다. 같은 파라미터, 같은 순서, 같은 dtype과 shape으로요. 이것은 특정 백엔드가 아니라 ABC의 불변식이에요. 두 채널 사이에서 파라미터를 재정렬하거나 생략하거나 재타입하는 소스는, 우연히 테스트하는 백엔드가 눈치채지 못하더라도 잘못된 것이죠. 백엔드는 두 채널을 모두 읽고 일치할 것이라고 신뢰할 자유가 있습니다. Dense NCCL은 그렇게 하며 강제합니다.
실체화는 보통 집단 연산이라, 모든 트레이너 랭크가 같은 소스를 같은 순서로, 함께(lockstep) 순회해야 해요. 안 그러면 랭크가 교착됩니다. metadata() 자체도 커스텀 생산자에겐 집단 연산일 수 있어서 모든 랭크에서 실행되고, 결과를 보내는 건 송신자뿐이에요.
iter(source)는 매 라운드 새로운 패스를 산출해야 합니다.
ModuleSource
ModuleSource(module)는 module.named_parameters() 위의 일반적인 경우예요. 일반 모듈과 FSDP 샤딩 모듈을 특수 취급 없이 다룹니다. 순회는 각 DTensor를 full_tensor()로 all-gather하고, metadata()는 전역 .shape / .dtype을 읽어 gather을 유발하지 않아요.
from vllm.distributed.weight_transfer import ModuleSource
source = ModuleSource(model)
커스텀 소스 (Custom sources)
가중치가 HF 형식이 되도록 작업이 필요할 때 — 프레임워크 특유의 export, 재융합 단계, dtype 캐스트 — WeightSource를 서브클래싱해요.
from vllm.distributed.weight_transfer import ParamMeta, WeightSource
class MegatronBridgeSource(WeightSource):
"""Megatron model -> HF names, via a bridge that gathers TP/PP/EP internally
and returns full tensors on every rank."""
def __init__(self, bridge, module, dtype):
self._bridge, self._module, self._dtype = bridge, module, dtype
self._meta: list[ParamMeta] | None = None
def _export(self):
return self._bridge.export_hf_weights(self._module)
def metadata(self) -> list[ParamMeta]:
# Cache: for producers that must materialize to learn shapes, this is
# the expensive channel. Runs on every rank (it may be a collective).
if self._meta is None:
self._meta = [
ParamMeta(name, self._dtype, tuple(t.shape))
for name, t in self._export()
]
return self._meta
def __iter__(self):
# Must yield exactly what metadata() declared, in the same order.
for name, tensor in self._export():
yield name, tensor.to(self._dtype).detach().contiguous()
held_names(): 부분 소유권
기본적으로 모든 랭크가 모든 파라미터를 생산할 수 있다고 가정해요 — 위 소스가 모든 병렬에 걸쳐 gather한 뒤 산출하기 때문에 그렇죠. 단순하고 항상 올바르지만, gather 비용을 모든 랭크가 전부 부담한다는 뜻이기도 해요.
랭크가 나뉘어 각자 모델의 일부만 가질 때는 선택적으로 held_names()를 오버라이드할 수 있어요. 이 랭크가 가진 파라미터 이름을 반환합니다(기본값 None은 전부를 의미).
def held_names(self):
# This pipeline stage's layers, and within them only this EP rank's experts.
return self._my_stage_names - self._foreign_expert_names
이것은 다양한 트레이너 레이아웃을 다룹니다 — 파이프라인 스테이지(랭크가 일부 레이어를 가짐), 전문가 병렬(랭크가 일부 전문가를 가짐), 둘 다, 또는 어느 쪽에도 안 맞는 모양까지요. 파라미터별로 라우팅할 수 있는 백엔드(예: sharded RDT)는 각 이름을 실제로 가진 랭크에서 끌어옵니다.
오버라이드할 때 따라오는 세 가지 요구사항이 있어요.
metadata()는 모든 랭크에서 여전히 전체 모델을 설명해야 해요. 송신자의 메타데이터만 추론 쪽에 닿으므로, 자기 몫만 보고한 랭크는 나머지를 조용히 전송하지 못하게 만들 거예요. Sharded RDT는 init에서 랭크 간에 이를 교차 확인합니다.- 모든 이름을 최소 한 랭크가 가져야 해요. 안 그러면 절대 서빙될 수 없죠. 엔진은 init에서 첫 번째 고아(orphan) 이름을 명시해 예외를 던집니다.
- 순회는 이 랭크가 가지지 않은 이름에 대해
None을 산출해요. 이름은 여전히 메타데이터 순서로 나타나 순서 검사가 랭크 간에 정렬된 채 유지되고, 데이터만 없습니다. 이름을 주장한 뒤 그에 대해None을 산출하는 것은 엔진이 이름으로 보고하는 오류예요.
부분 소유권은 sharded RDT에서만 동작해요. 파라미터별로 라우팅하는 백엔드만 held_names()를 존중합니다. Broadcast 백엔드는 이를 무시하고 모든 랭크에서 모든 이름을 보내므로, 거기서 부분 소유권을 선언해도 달라지는 게 없습니다.
Gather groups
일부 백엔드는 모델 단위가 아니라 레이어 단위로 전송해서, metadata()를 gather groups로 분할해요. layerwise_groups는 각 이름을 그것이 가진 가장 바깥 인덱스 세그먼트로 키잉하므로, 한 그룹이 한 디코더 레이어가 돼요. 인덱스 없는 이름들(embedding, 최종 norm, lm_head)이 나타나는 곳에서 각자 별도 그룹을 이룹니다.
group 0 model.embed_tokens.weight
group 1 model.layers.0.* <- one decoder layer
group 2 model.layers.1.*
...
group N+1 model.norm.weight, lm_head.weight
리터럴 접두사가 아니라 인덱스에 키잉하면 아키텍처별 테이블이 필요 없어요. model.layers.0., model.language_model.layers.0.(최근 Qwen 텍스트 체크포인트), transformer.h.0.(GPT-2, Falcon), backbone.layers.0.(Mamba), 비전 타워의 visual.blocks.0. 모두 같은 방식으로 분할됩니다. 취하는 인덱스는 가장 바깥 것이어서 MoE 레이어를 통째로 유지해요. model.layers.3.mlp.experts.7.w1 같은 전문가별 이름은 전문가가 아니라 레이어에 키잉되죠.
그룹 인덱스 _g_는 모든 랭크와 모든 소비자에서 같은 레이어를 의미합니다. 모든 쪽이 한 랭크의 metadata()에서 그걸 파생하기 때문이에요. 이 일치 덕분에 백엔드가 버퍼를 한 레이어로 묶고, 모두가 끝내면 레이어를 해제할 수 있어요.
한 리프 모듈의 소스는 모두 한 그룹에 있어야 해요. sharded-RDT 엔진은 마지막 청크가 도착하자마자 그룹을 해제하므로, 그룹을 가로지르는 모듈은 stall watchdog이 발동할 때까지 pull을 방치하게 됩니다. 기본 파티션이 이를 보장하고, groups() 오버라이드는 이를 유지해야 해요.
다음 두 훅이 이어지며 둘 다 동작하는 기본값을 가져요.
groups()— 메타데이터 순서대로 이 랭크의 그룹. 기본값은 하나 이상의 held 이름을 가진 그룹으로 제한된layerwise_groups(metadata())예요. 여기서 아무것도 가지지 않은 그룹은 완전히 건너뜁니다.iter_groups()— 한 번에 한 그룹씩 배치된 같은 스트림. 기본값은__iter__를 구동하고 그 출력을 배치하며, 이름이 메타데이터 순서로 오는지 검사해요. 프레임워크가 한 번에 전체 그룹을 생산할 수 있으면 오버라이드하세요. 실체화는 보통 집단 연산이라, 텐서별로 구동하는 대신 그룹별로 구동하면 전문가별 MoE 모델에서 약 37k번의 generator 재개가 약 95번으로 줄어듭니다.
VLLMWeightSyncClient
이것은 RL 스택이 vLLM에 도달하는 방식이 어떻든 그에 맞추는 어댑터예요. 많은 RL 프레임워크가 추론 엔진을 자기 추상화로 감싸고, 각자 나름대로 vLLM에 닿습니다. VLLMWeightSyncClient는 그 고유한 모양이 적응되는 단일 이음새라서, 가중치 동기화 엔진은 컨트롤 플레인에 구애받지 않는 상태로 유지돼요.
계약은 이것뿐입니다. 래퍼가 어떤 모양이든, 결국 같은 네 호출에 수렴해야 한다는 것 — 설정 시 init_weight_transfer_engine 한 번, 그다음 라운드마다 start_weight_update → 한 번 이상의 update_weights → finish_weight_update. 트레이너 엔진이 추론 쪽에서 필요로 하는 모든 것이 이 호출들을 통해 갑니다.
class VLLMWeightSyncClient(Protocol):
def init_weight_transfer_engine(self, init_info: dict[str, Any]) -> None: ...
def start_weight_update(self) -> None: ...
def update_weights(self, update_info: dict[str, Any]) -> None: ...
def finish_weight_update(self, weight_version: str | None = None) -> None: ...
이것은 @runtime_checkable 구조적 Protocol(PEP 544)이라 적응을 저렴하게 만들죠. 그 네 메서드를 가진 어떤 객체든 이미 이를 만족합니다. 프레임워크의 기존 래퍼는 보통 네 개의 전달 메서드를 추가하는 것만으로 클라이언트가 될 수 있어요.
vLLM과 함께 제공되는 두 구현이 있어요.
| 클라이언트 | 통화 대상 |
|---|---|
RayVLLMWeightSyncClient(handle) |
하나 이상의 AsyncLLM/LLM Ray actor. 리스트를 받아 모든 핸들로 각 호출을 확장(fan out)하고 전부에서 블록해서, 멀티 액터(예: 멀티-DP) 배포를 하나의 단위로 구동 |
HTTPVLLMWeightSyncClient(base_url, timeout=300) |
RLHF HTTP 라우트를 통한 vLLM 서버 |
커스텀 weight sync 클라이언트는 다음과 같이 구현할 수 있어요.
class MyFrameworkWeightSyncClient:
"""Adapts one RL framework's rollout pool to the four weight-sync calls."""
def __init__(self, rollout_pool):
self.pool = rollout_pool # whatever your stack already has
def init_weight_transfer_engine(self, init_info):
# Fan out to every replica and block: all of them receive weights.
self.pool.broadcast_rpc("init_weight_transfer_engine", init_info=init_info)
def start_weight_update(self):
self.pool.broadcast_rpc("start_weight_update")
def update_weights(self, update_info):
self.pool.broadcast_rpc("update_weights", update_info=update_info)
def finish_weight_update(self, weight_version=None):
self.pool.broadcast_rpc("finish_weight_update")
if weight_version is not None:
self.pool.broadcast_rpc("update_weight_version", weight_version)
어떤 어댑터에서든 두 가지를 제대로 해야 해요.
- 모든 레플리카에 닿고, 전부 끝날 때까지 블록. 가중치 갱신은 로드 밸런싱된 요청이 아니에요. 모델 사본을 가진 모든 워커가 받아야 합니다. 전부 끝나기 전에 반환하면 트레이너가 아직 로드 중인 워커를 앞질러 가죠. (두 내장 클라이언트 모두 이렇게 합니다. Ray는 핸들 위로 확장해서, HTTP는 서버의 DP 클라이언트가 내부적으로 broadcast하기 때문이에요.)
- 실패 시 예외 발생. 트레이너 엔진은 예외에 의존해 추론 쪽 오류를 표면화해요. 예외를 삼키는 클라이언트는 실패한 동기화를 조용히 오래된 가중치로 만들거나, 전송이 워커와 rendezvous 하는 백엔드라면 교착으로 만듭니다.
참고: HTTP는 원시 CUDA IPC 핸들을 실을 수 없어서, HTTPVLLMWeightSyncClient는 이를 pickle하고 base64로 인코딩해 ipc_handles_pickled 필드에 넣어요. 워커는 VLLM_ALLOW_INSECURE_SERIALIZATION=1일 때만 역직렬화합니다. JSON 네이티브 페이로드(NCCL)를 쓰는 백엔드는 그대로 통과합니다.
TrainerWeightTransferEngine
트레이너 쪽 엔진이에요. 전송 상태(NCCL communicator, IPC 디바이스 정보, 전송 계획)를 보유하고, WeightSource에서 가중치를 끌어오며, VLLMWeightSyncClient로 추론 쪽을 구동해요. init info 타입에 제네릭이고, trainer_init 클래스메서드 팩토리로 만들어지며, send_weights()로 구동됩니다.
| 메서드 | 설명 |
|---|---|
trainer_init(init_info, *, client, source=None) |
클래스메서드. 추론 쪽과 rendezvous 하고 준비된 인스턴스 반환 |
send_weights() |
가중치를 밀어 넣고 전체 갱신 라운드트립 구동 |
shutdown() |
communicator / 프로세스 그룹 정리. 기본 no-op |
trainer_init과 send_weights 둘 다 모든 트레이너 랭크에서 호출돼요. is_sender는 trainer_init에서 init_info.rank로 한 번 결정됩니다. 각 엔진은 모든 랭크에서 실제 클라이언트를 갖지만, 컨트롤 플레인 RPC와 전송을 self.is_sender로 가드해서 송신자만 와이어에 닿아요. 하지만 송신자가 아닌 랭크도 모든 집단 연산을 실행해 그룹이 정렬된 채 유지됩니다.
트레이너 쪽은 WeightTransferConfig를 받지 않아요. 백엔드는 init info의 backend ClassVar에서 오고, wire params도 init info를 타고 갑니다.
TrainerInitInfo
trainer_init에 넘기는 init_info예요. 호출자가 전송을 구성하는 방법으로, 백엔드를 선택하고 이 프로세스가 어떤 랭크인지 말하며 wire params를 담아요. 각 백엔드가 이것을 서브클래싱하고, 베이스 클래스는 모든 백엔드가 필요로 하는 필드 하나를 가집니다.
@dataclass
class TrainerInitInfo:
backend: ClassVar[str] # factory dispatch key
rank: int = field(kw_only=True)
@property
def is_sender(self) -> bool:
return self.rank == 0
rank— 이 트레이너 프로세스의 랭크, 명시적으로 제공. 엔진은 전역 프로세스 그룹에서 읽지 않아요. FSDP / TP / PP / EP처럼 여러 그룹이 존재하면 모호하니까요. 랭크 0은 항상 송신자이며,trainer_init이is_sender로 해석합니다. 키워드 전용이라 백엔드 서브클래스가 위치 필드를 자유롭게 추가할 수 있어요.backend—__init__필드가 아니라ClassVar. 팩토리가 디스패치하기 위해 읽는 백엔드 고정 상수라서, 호출자가backend=인자를 절대 넘기지 않는 이유예요. 모든 서브클래스가 설정해야 하며,__init_subclass__가 그렇지 않으면 예외를 던집니다.
서브클래스는 또한 전송의 wire params(packed, 버퍼 크기)를 담아요. 송신자가 trainer_init 안에서 워커로 전파하므로 양쪽이 어긋날 수 없습니다. 구체적인 필드는 NCCLTrainerInitInfo와 IPCTrainerInitInfo를 참고하세요.
Full-Resync vs. Delta 백엔드
source는 선택이라서 백엔드를 두 가지 모양으로 나눠요.
- Full resync (NCCL, IPC) — 안정적인
WeightSource가trainer_init에서 고정되고 매 라운드 재순회됨;send_weights()는 인자를 받지 않음. 이 백엔드들은source가 null이 아닌지 스스로 검증해요. - Delta (sparse NCCL) — 페이로드가 매 라운드 달라져서 안정적인 소스가 없음. 엔진은
source를 받지 않고, 각 라운드 페이로드가 바로send_weights(patches)로 전달돼요.
커스텀 트레이너 엔진 구현
from dataclasses import dataclass
from typing import ClassVar
from typing_extensions import Self
from vllm.distributed.weight_transfer.base import (
TrainerInitInfo,
TrainerWeightTransferEngine,
VLLMWeightSyncClient,
WeightSource,
)
@dataclass
class MyTrainerInitInfo(TrainerInitInfo):
backend: ClassVar[str] = "my_backend"
endpoint: str
chunk_size_bytes: int = 256 * 1024 * 1024 # a wire param: shipped to the worker
class MyTrainerWeightTransferEngine(TrainerWeightTransferEngine[MyTrainerInitInfo]):
init_info_cls = MyTrainerInitInfo
def __init__(self, *, client, source, is_sender=True, chunk_size_bytes=0):
super().__init__(client=client, source=source, is_sender=is_sender)
self.chunk_size_bytes = chunk_size_bytes
@classmethod
def trainer_init(
cls,
init_info: MyTrainerInitInfo,
*,
client: VLLMWeightSyncClient,
source: WeightSource | None = None,
) -> Self:
if source is None:
raise ValueError("my_backend requires a WeightSource.")
engine = cls(
client=client,
source=source,
is_sender=init_info.is_sender,
chunk_size_bytes=init_info.chunk_size_bytes,
)
if engine.is_sender:
# Ship the must-agree wire params so the worker decodes exactly as
# this trainer encodes, then open the trainer-side endpoint.
engine.client.init_weight_transfer_engine(
{"chunk_size_bytes": init_info.chunk_size_bytes}
)
return engine
def send_weights(self) -> None:
assert self.source is not None
meta = self.source.metadata() # every rank: may be a collective
if not self.is_sender:
for _ in self.source: # stay in the trainer-side collective
pass
return
self.client.start_weight_update()
self.client.update_weights(
{
"names": [m.name for m in meta],
"dtype_names": [str(m.dtype).split(".")[-1] for m in meta],
"shapes": [list(m.shape) for m in meta],
}
)
for name, tensor in self.source:
... # transmit
self.client.finish_weight_update()
두 가지를 제대로 해야 하는데, 둘 다 내장 백엔드를 물렸던 부분들이에요.
- 반환 전에 소진(drain).
send_weights는 전송이 진행 중인 채로 반환하면 안 돼요. 전송 버퍼를 살려 두는 그 무엇이든 프레임과 함께 죽고, 추론 쪽의finish_weight_update사후 처리가 아직 도착하지 않은 가중치를 확정할 수 있어요. - 오류 경로에서 컨트롤 플레인 스레드에 참여(join)하지 말 것. 전송과 동시에 사이드 스레드에서
update_weights를 실행하고(NCCL처럼) 전송이 예외를 던지면, 워커는 여전히 대응하는 집단 연산에서 블록되어 절대 반환하지 않아요. 기다리지 않고 executor를 종료해서 실제 예외가 교착 대신 표면화되게 하세요.
WeightTransferTrainerFactory
from vllm.distributed.weight_transfer import WeightTransferTrainerFactory
# Lazy loading (recommended): the module is imported only when the backend is used
WeightTransferTrainerFactory.register_engine(
"my_backend",
"my_package.my_module",
"MyTrainerWeightTransferEngine",
)
# Or register the class directly
WeightTransferTrainerFactory.register_engine("my_backend", MyTrainerWeightTransferEngine)
engine = WeightTransferTrainerFactory.trainer_init(
init_info=MyTrainerInitInfo(rank=0, endpoint="..."), # `backend` selects the engine
client=client,
source=source,
)
추론 쪽 (Inference Side)
WeightTransferEngine
두 개의 dataclass 타입으로 파라미터화된 제네릭 추상 클래스예요.
TInitInfo(extendsWeightTransferInitInfo): 백엔드별 초기화 파라미터.TUpdateInfo(extendsWeightTransferUpdateInfo): 백엔드별 가중치 갱신 메타데이터.
서브클래스는 다섯 개 메서드를 구현해야 해요.
| 메서드 | 설명 |
|---|---|
init_transfer_engine(init_info) |
각 추론 워커에서 통신 채널 초기화, 트레이너가 준 wire params 기록 |
start_weight_update() |
갱신 준비 (예: layerwise reload 시작); in-place 엔진은 no-op |
finish_weight_update() |
갱신 종료 (예: layerwise reload 종료); in-place 엔진은 no-op |
receive_weights(update_info) |
가중치를 받아 self.model에 로드 |
shutdown() |
리소스 정리 |
베이스 클래스가 제공하는 것:
__init__—config(WeightTransferConfig),vllm_config(VllmConfig),device(torch.device),model(nn.Module)를 받음.update_weights(update_info_dict)—receive_weights의 얇은 래퍼. 딕셔너리를 타입 dataclass로 파싱하고receive_weights를 호출하며 디바이스를 동기화. 단, 엔진이 아래의defers_processing을 설정하면 예외.parse_init_info/parse_update_info— API 레벨 딕셔너리를 타입 dataclass로 변환하고 잘못된 페이로드에ValueError를 던짐.set_weight_update_target/reset_weight_update_target— 갱신을 추측 드래프트 모델에 재타깃하는 데 사용.
wire params는 페이로드가 아니라 핸드셰이크에서 읽어요. 양쪽이 동의해야 하는 것 — packed, 버퍼 지오메트리 — 은 init info로 오고 init_transfer_engine에서 self에 저장한 뒤 receive_weights에서 self로 읽어야 합니다. 라운드별 update info는 라운드별 메타데이터만 담아요. 이것이 트레이너/워커 불일치를 표현 불가능하게 만드는 이유죠.
defers_processing: 반환된 갱신이 적용이 아니라 대기열에 들어간 것을 의미할 때
GPU 사후 처리를 백그라운드 스레드로 파이프라인하는 엔진은 update_weights가 디바이스를 동기화하도록 둘 수 없어요. 그 스레드들에서 블록하고 파이프라인을 직렬화할 테니까요. 그런 엔진은 클래스 속성 defers_processing = True를 설정하고, 갱신별 동기화를 생략하며, finish_weight_update에서 완료를 보장합니다.
finish_weight_update를 거치는 호출자는 아무것도 할 필요 없어요. 엔진이 거기서 소진(drain)하니까요. 꼬리를 스스로 구동하는 호출자 — 자체 finalize_layerwise_reload를 실행하는 경우 — 는 플래그를 확인하고 먼저 drain_pending()을 호출해야 해요. 이 플래그가 설정되면 반환된 update_weights는 적용됨이 아니라 대기열에 들어감을 뜻하니까요. drain_pending()은 멱등이고, 동기적으로 처리하는 엔진에선 no-op이라 항상 호출해도 안전합니다.
Sharded RDT가 이 플래그를 설정하는 내장 엔진이에요. 백그라운드 스레드에서 자체 CUDA 스트림으로 scatter·quantize를 수행하므로, finalize_layerwise_reload가 실행되기 전에 drain_pending()이 두 큐를 모두 조인하고 두 스트림을 동기화합니다.
요청 클래스 (Request Classes)
API 레벨 요청 클래스는 일반 딕셔너리를 사용해 백엔드에 구애받지 않는 직렬화를 제공해요.
from vllm.distributed.weight_transfer.base import (
WeightTransferInitRequest,
WeightTransferUpdateRequest,
)
# Init request (dict is converted to backend-specific TInitInfo)
init_request = WeightTransferInitRequest(
init_info={"master_address": "10.0.0.1", "master_port": 29500, ...}
)
# Update request (dict is converted to backend-specific TUpdateInfo)
update_request = WeightTransferUpdateRequest(
update_info={"names": [...], "dtype_names": [...], "shapes": [...]}
)
내장 클라이언트를 쓰면 이것들을 손으로 만들 일이 없어요. RayVLLMWeightSyncClient가 딕셔너리를 감싸주고, HTTPVLLMWeightSyncClient가 JSON으로 게시합니다.
LLM/API 레이어에서 추측 드래프트 모델을 타깃으로 하려면 start_weight_update() 대신 start_draft_weight_update()를 호출해요. update_weights / finish_weight_update는 변경되지 않습니다. 이를 지원할 수 없는 엔진은 supports_draft_weight_update = False를 설정해요.
커스텀 엔진 구현
1. Info dataclass 정의
from dataclasses import dataclass
from vllm.distributed.weight_transfer.base import (
WeightTransferEngine,
WeightTransferInitInfo,
WeightTransferUpdateInfo,
)
@dataclass
class MyInitInfo(WeightTransferInitInfo):
endpoint: str
chunk_size_bytes: int = 256 * 1024 * 1024 # must-agree wire param
@dataclass
class MyUpdateInfo(WeightTransferUpdateInfo):
names: list[str]
dtype_names: list[str]
shapes: list[list[int]]
# Per-round metadata only.
2. 엔진 구현
class MyWeightTransferEngine(WeightTransferEngine[MyInitInfo, MyUpdateInfo]):
init_info_cls = MyInitInfo
update_info_cls = MyUpdateInfo
def init_transfer_engine(self, init_info: MyInitInfo) -> None:
# Record the trainer's wire params, then set up the connection.
self.chunk_size_bytes = init_info.chunk_size_bytes
...
def start_weight_update(self) -> None:
# Checkpoint-format engines: run initialize_layerwise_reload(self.model).
# In-place engines: no-op
...
def finish_weight_update(self) -> None:
# Checkpoint-format engines: run finalize_layerwise_reload(...).
# In-place engines: no-op
...
def receive_weights(self, update_info: MyUpdateInfo) -> None:
weights = []
for name, dtype_name, shape in zip(
update_info.names, update_info.dtype_names, update_info.shapes
):
dtype = getattr(torch, dtype_name)
weight = self._fetch_weight(name, shape, dtype)
weights.append((name, weight))
self.model.load_weights(weights)
def shutdown(self) -> None:
# Clean up resources
...
3. 팩토리에 등록
from vllm.distributed.weight_transfer import WeightTransferEngineFactory
# Option 1: Lazy loading (recommended for built-in engines)
WeightTransferEngineFactory.register_engine(
"my_backend",
"my_package.my_module",
"MyWeightTransferEngine",
)
# Option 2: Direct class registration
WeightTransferEngineFactory.register_engine(
"my_backend",
MyWeightTransferEngine,
)
등록 후 사용자는 WeightTransferConfig(backend="my_backend")로 백엔드를 선택할 수 있어요.
WeightTransferEngineFactory
팩토리는 lazy loading이 있는 레지스트리 패턴을 사용해요. 내장 엔진(nccl, ipc, sparse_nccl, sharded_rdt)은 import 시점에 등록되지만, 모듈은 백엔드가 실제로 요청될 때만 로드됩니다. 이렇게 하면 필요 없을 때 무거운 의존성(NCCL communicator 등) import를 피할 수 있죠.
from vllm.distributed.weight_transfer import WeightTransferEngineFactory
# Create an engine from config
engine = WeightTransferEngineFactory.create_engine(
config=weight_transfer_config,
vllm_config=vllm_config,
device=device,
model=model,
)
vLLM은 워커 시작 중 이것을 자동 호출해요. 자체 워커에 엔진을 임베드할 때만 직접 호출하면 됩니다.