From db64e71508651b535883eb5e43371642336071f2 Mon Sep 17 00:00:00 2001 From: weikaiwen <34648228+kevssim@users.noreply.github.com> Date: Fri, 28 Aug 2026 11:22:09 +0800 Subject: [PATCH 1/3] wip --- .../Checkpoint Engine/CheckpointEngine.md | 19 +- .../CheckpointEngine.md" | 18 +- src/twinkle/checkpoint_engine/__init__.py | 9 +- src/twinkle/checkpoint_engine/base.py | 6 +- src/twinkle/checkpoint_engine/manager.py | 186 +++++++--- src/twinkle/model/megatron/megatron.py | 27 +- .../model/transformers/transformers.py | 28 +- .../sampler/sglang_sampler/sglang_sampler.py | 22 +- .../sampler/vllm_sampler/vllm_sampler.py | 22 +- tests/checkpoint_engine/__init__.py | 1 + tests/checkpoint_engine/test_manager_naive.py | 340 ++++++++++++++++++ tests/sampler/test_ipc_checkpoint_engine.py | 2 +- 12 files changed, 585 insertions(+), 95 deletions(-) create mode 100644 tests/checkpoint_engine/__init__.py create mode 100644 tests/checkpoint_engine/test_manager_naive.py diff --git a/docs/source_en/Components/Checkpoint Engine/CheckpointEngine.md b/docs/source_en/Components/Checkpoint Engine/CheckpointEngine.md index 1a7c39bfa..cd4f371e6 100644 --- a/docs/source_en/Components/Checkpoint Engine/CheckpointEngine.md +++ b/docs/source_en/Components/Checkpoint Engine/CheckpointEngine.md @@ -2,6 +2,15 @@ CheckpointEngine is a component used to synchronize model weights between trainer and inference processes, primarily used in RLHF training to synchronize weights between Actor models and Rollout samplers. +`CheckpointEngineManager` exposes four modes: + +- `auto`: local objects use `naive`; Ray actor handlers use `standalone`. +- `naive`: stream the model's weight generator directly into a local sampler without creating a checkpoint engine. +- `colocate`: synchronize Ray actors sharing GPUs through CUDA IPC. +- `standalone`: synchronize disaggregated Ray actors through NCCL on GPU or HCCL on NPU. + +`auto` never infers `colocate`, because actor placement cannot be determined reliably from the driver. + ## Basic Interface ```python @@ -39,7 +48,7 @@ class CheckpointEngine(ABC): ## Available Checkpoint Engines -Twinkle provides two checkpoint engine implementations: +Twinkle provides three cross-process checkpoint engine implementations. `naive` mode bypasses them. ### NCCLCheckpointEngine @@ -61,10 +70,18 @@ A checkpoint engine that uses HCCL for weight transfer between Ascend NPUs. See: [HCCLCheckpointEngine](HCCLCheckpointEngine.md) +### IPCCheckpointEngine + +A CUDA IPC engine for model and sampler Ray actors placed on the same physical GPUs. NCCL cannot be +used for this topology because it rejects multiple ranks bound to one GPU. Weight buckets are mapped +between the actor processes rather than broadcast between devices. + ## How to Choose - **NCCLCheckpointEngine**: Suitable for GPU environments, provides the highest transfer performance - **HCCLCheckpointEngine**: Suitable for Ascend NPU environments +- **IPCCheckpointEngine**: Required for colocated Ray actors sharing physical GPUs +- **No engine (`naive`)**: Local model and sampler objects in the same process > Checkpoint engine is a key component of RLHF training infrastructure, ensuring that trainers and samplers use consistent model weights. > Currently, synchronization is divided into two cases based on merge_and_sync=True/False. When set to True, the LoRA is merged into the base model and then synchronized. diff --git "a/docs/source_zh/\347\273\204\344\273\266/\346\243\200\346\237\245\347\202\271\345\274\225\346\223\216/CheckpointEngine.md" "b/docs/source_zh/\347\273\204\344\273\266/\346\243\200\346\237\245\347\202\271\345\274\225\346\223\216/CheckpointEngine.md" index 338be10db..24e90947b 100644 --- "a/docs/source_zh/\347\273\204\344\273\266/\346\243\200\346\237\245\347\202\271\345\274\225\346\223\216/CheckpointEngine.md" +++ "b/docs/source_zh/\347\273\204\344\273\266/\346\243\200\346\237\245\347\202\271\345\274\225\346\223\216/CheckpointEngine.md" @@ -2,6 +2,15 @@ CheckpointEngine (检查点引擎) 是用于在训练器和推理进程之间同步模型权重的组件,主要用于 RLHF 训练中 Actor 模型和 Rollout 采样器之间的权重同步。 +`CheckpointEngineManager` 提供四种模式: + +- `auto`:本地对象使用 `naive`;Ray actor handler 使用 `standalone`。 +- `naive`:模型的权重生成器直接流式传入本地 sampler,不创建 CheckpointEngine。 +- `colocate`:共享 GPU 的 Ray actors 通过 CUDA IPC 同步。 +- `standalone`:分离部署的 Ray actors 在 GPU 上使用 NCCL,在 NPU 上使用 HCCL。 + +`auto` 不会推断 `colocate`,因为 driver 无法可靠判断 actor 的实际设备放置。 + ## 基本接口 ```python @@ -39,7 +48,7 @@ class CheckpointEngine(ABC): ## 可用的检查点引擎 -Twinkle 提供了两种检查点引擎实现: +Twinkle 提供三种跨进程检查点引擎实现;`naive` 模式会绕过这些引擎。 ### NCCLCheckpointEngine @@ -61,10 +70,17 @@ Twinkle 提供了两种检查点引擎实现: 详见: [HCCLCheckpointEngine](HCCLCheckpointEngine.md) +### IPCCheckpointEngine + +适用于模型和 sampler Ray actors 被放置在同一组物理 GPU 上的 CUDA IPC 引擎。NCCL 会拒绝多个 +rank 绑定同一张 GPU,因此该拓扑必须通过 CUDA IPC 在进程间映射权重 bucket,而不是跨设备广播。 + ## 如何选择 - **NCCLCheckpointEngine**: 适用于 GPU 环境,提供最高的传输性能 - **HCCLCheckpointEngine**: 适用于昇腾 NPU 环境 +- **IPCCheckpointEngine**: 适用于共享物理 GPU 的 colocated Ray actors +- **不创建引擎 (`naive`)**: 适用于同一进程内的本地 model 和 sampler > 检查点引擎是 RLHF 训练基础设施的关键组件,确保训练器和采样器使用一致的模型权重。 > 目前的同步分为merge_and_sync=True/False两种情况,为True时将lora合并仅基模并同步,为False时仅同步lora权重。另外,多租户直接附加lora文件到vLLM上,在merge_and_sync=False,或使用多租户时, diff --git a/src/twinkle/checkpoint_engine/__init__.py b/src/twinkle/checkpoint_engine/__init__.py index 37fba79ba..1b2cbd247 100644 --- a/src/twinkle/checkpoint_engine/__init__.py +++ b/src/twinkle/checkpoint_engine/__init__.py @@ -1,8 +1,10 @@ # Copyright (c) ModelScope Contributors. All rights reserved. """Checkpoint Engine for weight synchronization between trainer and rollout. -Provides NCCL/HCCL-based weight broadcast from training model workers to -inference sampler workers in STANDALONE (disaggregated) deployment mode. +``CheckpointEngineManager`` supports three synchronization modes: direct +generator streaming for local objects (``naive``), CUDA IPC for colocated Ray +actors (``colocate``), and NCCL/HCCL for disaggregated Ray actors +(``standalone``). Reference: https://github.com/volcengine/verl/tree/main/verl/checkpoint_engine @@ -16,7 +18,7 @@ from .base import CheckpointEngine, TensorMeta from .hccl_checkpoint_engine import HCCLCheckpointEngine from .ipc_checkpoint_engine import IPCCheckpointEngine -from .manager import CheckpointEngineManager +from .manager import CheckpointEngineManager, CheckpointEngineMode from .mixin import CheckpointEngineMixin # Import backend implementations to register them from .nccl_checkpoint_engine import NCCLCheckpointEngine @@ -25,6 +27,7 @@ 'CheckpointEngine', 'CheckpointEngineMixin', 'CheckpointEngineManager', + 'CheckpointEngineMode', 'NCCLCheckpointEngine', 'HCCLCheckpointEngine', 'IPCCheckpointEngine', diff --git a/src/twinkle/checkpoint_engine/base.py b/src/twinkle/checkpoint_engine/base.py index f3a1d8918..e5f00c7f0 100644 --- a/src/twinkle/checkpoint_engine/base.py +++ b/src/twinkle/checkpoint_engine/base.py @@ -16,10 +16,12 @@ class TensorMeta(TypedDict): class CheckpointEngine(ABC): - """Abstract base class for checkpoint engines. + """Abstract base class for cross-process checkpoint engines. A checkpoint engine handles weight synchronization between trainer and rollout - processes. The typical workflow is: + processes. Local ``naive`` synchronization bypasses this interface and streams + the model's weight generator directly into the sampler. The typical cross-process + workflow is: In trainer process (rank 0): >>> engine = CheckpointEngineRegistry.new('nccl', bucket_size=512<<20) diff --git a/src/twinkle/checkpoint_engine/manager.py b/src/twinkle/checkpoint_engine/manager.py index 355a9b58b..9f7a24666 100644 --- a/src/twinkle/checkpoint_engine/manager.py +++ b/src/twinkle/checkpoint_engine/manager.py @@ -1,6 +1,6 @@ # Copyright (c) ModelScope Contributors. All rights reserved. # Adapted from https://github.com/volcengine/verl/blob/main/verl/checkpoint_engine/base.py -from typing import List, Optional +from typing import List, Literal, Optional from twinkle import Platform, get_logger from .base import CheckpointEngine @@ -8,14 +8,21 @@ logger = get_logger() +CheckpointEngineMode = Literal['auto', 'naive', 'colocate', 'standalone'] +_VALID_MODES = {'auto', 'naive', 'colocate', 'standalone'} + class CheckpointEngineManager: - """Weight synchronization manager for Twinkle. + """Weight synchronization manager for local and Ray deployments. + + ``mode`` selects one of three synchronization paths: + + * ``naive`` streams a local model's weight generator directly into a local sampler. + * ``colocate`` connects Ray model and sampler actors sharing GPUs through CUDA IPC. + * ``standalone`` connects disaggregated Ray actors through NCCL/HCCL. - Coordinates weight synchronization between training model and inference sampler, either when they - reside on **different GPUs** (disaggregated / standalone deployment, the default) or when they - **share** one (``colocate=True``). Colocation replaces the NCCL broadcast drawn below with a CUDA - IPC handover per GPU -- not as an optimisation, but because NCCL refuses two ranks on one device. + ``auto`` resolves local objects to ``naive`` and Ray actor handlers to ``standalone``. It never + guesses ``colocate`` because actor placement cannot be inferred reliably from the driver. Architecture (following verl's CheckpointEngineManager): @@ -25,7 +32,7 @@ class CheckpointEngineManager: │ (Ray actors) │ │ (Ray actors) │ │ │ │ │ │ │ │ ▼ │ │ ▼ │ - │ CheckpointEngine │ NCCL broadcast │ CheckpointEngine │ + │ CheckpointEngine │ NCCL/HCCL/CUDA IPC │ CheckpointEngine │ │ send_weights() │ ─────────────────► │ receive_weights()│ │ │ │ │ │ │ │ │ ▼ │ @@ -42,11 +49,11 @@ class CheckpointEngineManager: >>> manager = CheckpointEngineManager(model=model, sampler=sampler) >>> manager.sync_weights() # Call after each training step - Colocated, the caller also owns the memory schedule, because only it knows where in the loop the - device is free. The sampler must have its weights resident to be written into -- ``sleep(1)`` puts - them on the host -- and the trainer has to step aside before a rollout: + With colocated Ray actors, the caller also owns the memory schedule, because only it knows where + in the loop the device is free. The sampler must have its weights resident to be written into -- + ``sleep(1)`` puts them on the host -- and the trainer has to step aside before a rollout: - >>> manager = CheckpointEngineManager(model=model, sampler=sampler, colocate=True) + >>> manager = CheckpointEngineManager(model=model, sampler=sampler, mode='colocate') >>> sampler.wake_up(tags=['weights']) # able to receive, still without a KV cache >>> manager.sync_weights() >>> model.offload_to_cpu() # the trainer's turn is over @@ -64,20 +71,15 @@ def __init__( model: 'CheckpointEngineMixin', sampler: 'CheckpointEngineMixin', platform: str = 'GPU', - colocate: bool = False, + mode: CheckpointEngineMode = 'auto', ) -> None: self.model = model self.sampler = sampler - self.colocate = colocate - self.backend_cls = self.decide_backend_engine(platform, colocate) + self.requested_mode = mode + self.mode = self._resolve_mode(mode, model, sampler) + self.backend_cls = self.decide_backend_engine(platform, self.mode) - # Validate Ray actors - assert hasattr(model, '_actors') and model._actors, \ - 'CheckpointEngineManager requires model to be deployed as Ray actors' - assert hasattr(sampler, '_actors') and sampler._actors, \ - 'CheckpointEngineManager requires sampler to be deployed as Ray actors' - - if colocate: + if self.mode == 'colocate': # Each side builds its own engine inside its worker, so both have to be told which one. self.model.set_checkpoint_engine_backend('ipc') self.sampler.set_checkpoint_engine_backend('ipc') @@ -91,14 +93,50 @@ def __init__( self._model_keys: Optional[List[str]] = None @staticmethod - def decide_backend_engine(platform: Optional[str] = None, colocate: bool = False) -> 'CheckpointEngine': - if colocate: + def _resolve_mode( + mode: CheckpointEngineMode, + model: 'CheckpointEngineMixin', + sampler: 'CheckpointEngineMixin', + ) -> Literal['naive', 'colocate', 'standalone']: + if mode not in _VALID_MODES: + valid = ', '.join(sorted(_VALID_MODES)) + raise ValueError(f'Unknown checkpoint engine mode {mode!r}; expected one of: {valid}.') + + model_has_actors = bool(getattr(model, '_actors', None)) + sampler_has_actors = bool(getattr(sampler, '_actors', None)) + if model_has_actors != sampler_has_actors: + raise ValueError( + 'CheckpointEngineManager requires model and sampler to use the same deployment shape: ' + 'both must be local objects or both must be Ray actor handlers.') + + if mode == 'auto': + return 'standalone' if model_has_actors else 'naive' + if mode == 'naive' and model_has_actors: + raise ValueError("mode='naive' requires local model and sampler objects without Ray actors.") + if mode in ('colocate', 'standalone') and not model_has_actors: + raise ValueError(f"mode={mode!r} requires model and sampler to be Ray actor handlers.") + return mode + + @staticmethod + def decide_backend_engine( + platform: Optional[str] = None, + mode: Literal['naive', 'colocate', 'standalone'] = 'standalone', + ) -> Optional['CheckpointEngine']: + if mode == 'naive': + return None + + platform_name = Platform.get_platform(platform).__name__ + if mode == 'colocate': + if platform_name != 'GPU': + raise NotImplementedError("mode='colocate' currently requires the GPU platform.") from twinkle.checkpoint_engine import IPCCheckpointEngine return IPCCheckpointEngine - if Platform.get_platform(platform).__name__ == 'GPU': + if mode != 'standalone': + raise ValueError(f'Cannot select a backend for unresolved mode {mode!r}.') + if platform_name == 'GPU': from twinkle.checkpoint_engine import NCCLCheckpointEngine return NCCLCheckpointEngine - elif Platform.get_platform(platform).__name__ == 'NPU': + elif platform_name == 'NPU': from twinkle.checkpoint_engine import HCCLCheckpointEngine return HCCLCheckpointEngine else: @@ -124,8 +162,12 @@ def sync_weights(self, merge_and_sync=True): Returns: None """ - model_metadata = self.model.prepare_checkpoint_engine([True] - + [False] * (self.model.device_mesh.world_size - 1)) + if self.mode == 'naive': + self._sync_weights_naive(merge_and_sync) + return + + is_master = [True] + [False] * (self.model.device_mesh.world_size - 1) + model_metadata = self.model.prepare_checkpoint_engine(is_master) self.sampler.prepare_checkpoint_engine(False) model_kwargs, sampler_kwargs = self.backend_cls.build_topology( self.model.device_mesh.world_size, @@ -146,36 +188,7 @@ def sync_weights(self, merge_and_sync=True): self._peft_config = self.model.get_peft_config_dict() peft_config = self._peft_config - if self._model_keys is None: - if hasattr(self.sampler, 'get_state_keys'): - self._model_keys = self.sampler.get_state_keys() - - if self._model_keys is None: - self._model_keys = [] - - # vLLM may have grouped params - use word boundaries to avoid substring matches - import re - _STACKED_MAPPINGS = [ - (re.compile(r'\bqkv_proj\b'), ('q_proj', 'k_proj', 'v_proj', 'q', 'k', 'v')), - (re.compile(r'\bgate_up_proj\b'), ('gate_proj', 'up_proj')), - (re.compile(r'\bin_proj_ba\b'), ('in_proj_b', 'in_proj_a')), - (re.compile(r'\blanguage_model\.model\b'), ('model.language_model', )), - (re.compile(r'^visual\.'), ('model.visual.', )), - ] - - def _expand_keys(keys): - result = set(keys) - for key in keys: - for pattern, individuals in _STACKED_MAPPINGS: - if pattern.search(key): - for ind in individuals: - result.add(pattern.sub(ind, key)) - return result - - # Two passes for chain expansion (e.g., language_model.model + qkv_proj) - expanded = _expand_keys(self._model_keys) - expanded = _expand_keys(expanded) - self._model_keys = list(expanded) + self._ensure_model_keys() model_result = self.model.send_weights( base_sync_done=self.base_sync_done, merge_and_sync=merge_and_sync, model_keys=self._model_keys) @@ -190,3 +203,62 @@ def _expand_keys(keys): self.base_sync_done = True if not merge_and_sync: logger.info('Base model sync completed, subsequent syncs will be LoRA-only') + + def _ensure_model_keys(self): + if self._model_keys is not None: + return + + if hasattr(self.sampler, 'get_state_keys'): + self._model_keys = self.sampler.get_state_keys() + + if self._model_keys is None: + self._model_keys = [] + + # vLLM may have grouped params - use word boundaries to avoid substring matches + import re + _STACKED_MAPPINGS = [ + (re.compile(r'\bqkv_proj\b'), ('q_proj', 'k_proj', 'v_proj', 'q', 'k', 'v')), + (re.compile(r'\bgate_up_proj\b'), ('gate_proj', 'up_proj')), + (re.compile(r'\bin_proj_ba\b'), ('in_proj_b', 'in_proj_a')), + (re.compile(r'\blanguage_model\.model\b'), ('model.language_model', )), + (re.compile(r'^visual\.'), ('model.visual.', )), + ] + + def _expand_keys(keys): + result = set(keys) + for key in keys: + for pattern, individuals in _STACKED_MAPPINGS: + if pattern.search(key): + for ind in individuals: + result.add(pattern.sub(ind, key)) + return result + + # Two passes for chain expansion (e.g., language_model.model + qkv_proj) + expanded = _expand_keys(self._model_keys) + expanded = _expand_keys(expanded) + self._model_keys = list(expanded) + + def _sync_weights_naive(self, merge_and_sync): + """Stream model weights directly into a local sampler.""" + peft_config = None + if self.base_sync_done and not merge_and_sync: + if self._peft_config is None: + self._peft_config = self.model.get_peft_config_dict() + peft_config = self._peft_config + + self._ensure_model_keys() + weights = self.model._get_weight_generator( + base_sync_done=self.base_sync_done, + merge_and_sync=merge_and_sync, + model_keys=self._model_keys, + ) + self.sampler.receive_weights( + weights=weights, + base_sync_done=self.base_sync_done, + peft_config=peft_config, + ) + + if not self.base_sync_done: + self.base_sync_done = True + if not merge_and_sync: + logger.info('Base model sync completed, subsequent syncs will be LoRA-only') diff --git a/src/twinkle/model/megatron/megatron.py b/src/twinkle/model/megatron/megatron.py index 37a3ee22e..214e99be0 100644 --- a/src/twinkle/model/megatron/megatron.py +++ b/src/twinkle/model/megatron/megatron.py @@ -1784,6 +1784,7 @@ def get_train_configs(self, **kwargs): # prepare_checkpoint_engine, init_checkpoint_process_group, and # finalize_checkpoint_engine are inherited from CheckpointEngineMixin. # + # The weight generator is shared by direct and checkpoint-engine sync. # Key difference from TransformersModel: Megatron uses TP/PP, so # get_hf_state_dict() internally performs TP allgather and handles PP # layer distribution. All model ranks MUST execute the weight generator @@ -1791,8 +1792,7 @@ def get_train_configs(self, **kwargs): # model_actor[0] (rank=0 in the checkpoint engine) actually broadcasts # via NCCL; others consume the generator silently (rank=-1). - @remote_function(dispatch='all', lazy_collect=True) - def send_weights( + def _get_weight_generator( self, adapter_name: str = None, base_sync_done: bool = False, @@ -1801,7 +1801,6 @@ def send_weights( ): if adapter_name is None: adapter_name = self._get_default_group() - engine = self._get_or_create_checkpoint_engine() @contextmanager def merge_lora(): @@ -1893,15 +1892,33 @@ def weight_generator(): else: yield from _raw_weights(False) + return weight_generator() + + @remote_function(dispatch='all', lazy_collect=True) + def send_weights( + self, + adapter_name: str = None, + base_sync_done: bool = False, + merge_and_sync: bool = False, + model_keys: List[str] = None, + ): + engine = self._get_or_create_checkpoint_engine() + weight_generator = self._get_weight_generator( + adapter_name=adapter_name, + base_sync_done=base_sync_done, + merge_and_sync=merge_and_sync, + model_keys=model_keys, + ) + is_sender = (engine.rank is not None and engine.rank == 0) if not is_sender: - for _name, _tensor in weight_generator(): + for _name, _tensor in weight_generator: pass return async def _send(): - await engine.send_weights(weight_generator()) + await engine.send_weights(weight_generator) result_container = {'error': None} diff --git a/src/twinkle/model/transformers/transformers.py b/src/twinkle/model/transformers/transformers.py index 5b4a2eef1..715ae5649 100644 --- a/src/twinkle/model/transformers/transformers.py +++ b/src/twinkle/model/transformers/transformers.py @@ -1576,10 +1576,9 @@ def get_train_configs(self, **kwargs) -> str: # ========================================================================= # prepare_checkpoint_engine, init_checkpoint_process_group, and # finalize_checkpoint_engine are inherited from CheckpointEngineMixin. - # Only send_weights_via_checkpoint_engine is model-specific. + # The weight generator is shared by direct and checkpoint-engine sync. - @remote_function(dispatch='all', lazy_collect=True) - def send_weights( + def _get_weight_generator( self, adapter_name: str = None, base_sync_done: bool = False, @@ -1589,7 +1588,6 @@ def send_weights( ): if adapter_name is None: adapter_name = self._get_default_group() - engine = self._get_or_create_checkpoint_engine() # Get state dict from unwrapped model model = self.strategy.unwrap_model(self.model) @@ -1663,11 +1661,31 @@ def weight_generator(): yield name, tensor _print_weight_example(names) + return weight_generator() + + @remote_function(dispatch='all', lazy_collect=True) + def send_weights( + self, + adapter_name: str = None, + base_sync_done: bool = False, + merge_and_sync: bool = False, + model_keys: List[str] = None, + **kwargs, + ): + engine = self._get_or_create_checkpoint_engine() + weight_generator = self._get_weight_generator( + adapter_name=adapter_name, + base_sync_done=base_sync_done, + merge_and_sync=merge_and_sync, + model_keys=model_keys, + **kwargs, + ) + # Run async send_weights in a dedicated event loop thread. # We cannot use the Ray worker's event loop because it may already # be occupied, and send_weights uses run_in_executor internally. async def _send(): - await engine.send_weights(weight_generator()) + await engine.send_weights(weight_generator) result_container = {'error': None} diff --git a/src/twinkle/sampler/sglang_sampler/sglang_sampler.py b/src/twinkle/sampler/sglang_sampler/sglang_sampler.py index 7903bd42a..230bda272 100644 --- a/src/twinkle/sampler/sglang_sampler/sglang_sampler.py +++ b/src/twinkle/sampler/sglang_sampler/sglang_sampler.py @@ -395,30 +395,32 @@ def receive_weights( self, base_sync_done: bool = False, peft_config: dict = None, + weights=None, ): """Receive weights from the trainer and stream them into sglang. - Which transport delivers them is the checkpoint engine's business, not this method's: NCCL - broadcast when the trainer is on other GPUs, CUDA IPC when it shares this one, where NCCL - cannot be used at all. - - The checkpoint engine's ``receive_weights()`` async generator is handed straight to - :meth:`SGLangEngine.update_weights`, which consumes it one tensor at a time into a bucket and - forwards each full bucket to sglang. Peak extra memory is one bucket rather than a second copy - of the model, the same reason the vLLM path streams. + With no ``weights`` argument, the checkpoint engine supplies an async generator from its + NCCL/HCCL/CUDA IPC transport. A local naive caller can instead provide the model's synchronous + generator directly. Either iterator is handed straight to :meth:`SGLangEngine.update_weights`, + which consumes it one tensor at a time into a bucket and forwards each full bucket to sglang. + Peak extra memory is one bucket rather than a second copy of the model. Args: base_sync_done: If True, this would be a LoRA-only sync. peft_config: PEFT config dict for LoRA adapter loading. + weights: Optional synchronous/asynchronous weight iterator. If + omitted, weights are received from the checkpoint engine. Raises: NotImplementedError: For a LoRA-only sync; see the class docstring. """ - engine = self._get_or_create_checkpoint_engine() + if weights is None: + engine = self._get_or_create_checkpoint_engine() + weights = engine.receive_weights() async def _receive_and_load(): await self.engine.update_weights( - engine.receive_weights(), # async generator — not materialised + weights, # async/sync generator — not materialised peft_config=peft_config, base_sync_done=base_sync_done, ) diff --git a/src/twinkle/sampler/vllm_sampler/vllm_sampler.py b/src/twinkle/sampler/vllm_sampler/vllm_sampler.py index d470a4b0c..d2c9520a1 100644 --- a/src/twinkle/sampler/vllm_sampler/vllm_sampler.py +++ b/src/twinkle/sampler/vllm_sampler/vllm_sampler.py @@ -467,19 +467,17 @@ def receive_weights( self, base_sync_done: bool = False, peft_config: dict = None, + weights=None, ): """Receive weights from the trainer and stream them into vLLM. - Which transport delivers them is the checkpoint engine's business, not this method's: NCCL - broadcast when the trainer is on other GPUs, CUDA IPC when it shares this one, where NCCL - cannot be used at all. Either way what arrives here is the same async generator. - Uses a **streaming pipeline** to avoid accumulating a full model-weight copy on GPU: - 1. ``CheckpointEngine.receive_weights()`` yields tensors from - the engine's buckets (async generator, GPU tensors). - 2. The async generator is passed **directly** to + 1. With no ``weights`` argument, ``CheckpointEngine.receive_weights()`` + yields tensors from NCCL/HCCL/CUDA IPC buckets. A local naive + caller can instead provide the model's synchronous generator. + 2. The weight iterator is passed **directly** to ``VLLMEngine.update_weights()`` which consumes it one tensor at a time, copying each into a GPU IPC bucket and flushing to the vLLM worker subprocess when the bucket is full. @@ -490,18 +488,22 @@ def receive_weights( Args: base_sync_done: If True, this is a LoRA-only sync. peft_config: PEFT config dict for LoRA adapter loading. + weights: Optional synchronous/asynchronous weight iterator. If + omitted, weights are received from the checkpoint engine. Returns: Number of weights loaded (approximate, from engine log). """ - engine = self._get_or_create_checkpoint_engine() + if weights is None: + engine = self._get_or_create_checkpoint_engine() + weights = engine.receive_weights() async def _receive_and_load(): - # Stream the received tensors directly into vLLM via IPC. + # Stream model/checkpoint-engine tensors directly into vLLM via IPC. # VLLMEngine.update_weights accepts an async generator and # handles bucket packing + ZMQ transfer internally. await self.engine.update_weights( - engine.receive_weights(), # async generator — not materialised + weights, # async/sync generator — not materialised peft_config=peft_config, base_sync_done=base_sync_done, ) diff --git a/tests/checkpoint_engine/__init__.py b/tests/checkpoint_engine/__init__.py new file mode 100644 index 000000000..85b3e739d --- /dev/null +++ b/tests/checkpoint_engine/__init__.py @@ -0,0 +1 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. diff --git a/tests/checkpoint_engine/test_manager_naive.py b/tests/checkpoint_engine/test_manager_naive.py new file mode 100644 index 000000000..eb2c041b1 --- /dev/null +++ b/tests/checkpoint_engine/test_manager_naive.py @@ -0,0 +1,340 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""CPU-only tests for CheckpointEngineManager mode selection and direct sync.""" + +import asyncio + +import pytest + +from twinkle.checkpoint_engine.manager import CheckpointEngineManager + + +class _Mesh: + world_size = 1 + data_world_size = 1 + + +class _Model: + + def __init__(self, weights): + self.device_mesh = _Mesh() + self._weights = weights + self.generator_calls = [] + self.peft_config_calls = 0 + self._checkpoint_engine = None + + def _get_weight_generator(self, **kwargs): + self.generator_calls.append(kwargs) + weights = self._weights(kwargs) if callable(self._weights) else self._weights + + def _weights(): + yield from weights + + return _weights() + + def get_peft_config_dict(self): + self.peft_config_calls += 1 + return {'r': 8, 'target_modules': ['q_proj']} + + +class _Sampler: + + def __init__(self, fail=None): + self.device_mesh = _Mesh() + self.calls = [] + self.loaded = [] + self.fail = fail + self._checkpoint_engine = None + + def get_state_keys(self): + return ['q_proj.weight'] + + def receive_weights(self, **kwargs): + self.calls.append(kwargs) + if self.fail is not None: + raise self.fail + self.loaded.append(list(kwargs['weights'])) + + +@pytest.mark.parametrize('requested_mode', ['auto', 'naive']) +def test_local_sync_streams_generator_without_checkpoint_engine(requested_mode): + model = _Model([('q_proj.weight', 'base')]) + sampler = _Sampler() + + manager = CheckpointEngineManager(model, sampler, platform='CPU', mode=requested_mode) + manager.sync_weights(merge_and_sync=False) + + assert manager.requested_mode == requested_mode + assert manager.mode == 'naive' + assert manager.backend_cls is None + assert manager.base_sync_done is True + assert sampler.loaded == [[('q_proj.weight', 'base')]] + assert model._checkpoint_engine is None + assert sampler._checkpoint_engine is None + assert model.generator_calls == [{ + 'base_sync_done': False, + 'merge_and_sync': False, + 'model_keys': ['q_proj.weight'], + }] + + +def test_local_lora_sync_reuses_peft_config_and_sends_only_incremental_weights(): + model = _Model(lambda call: ([('q_proj.lora_A', 'adapter')] + if call['base_sync_done'] else [('q_proj.weight', 'base')])) + sampler = _Sampler() + manager = CheckpointEngineManager(model, sampler, platform='CPU') + + manager.sync_weights(merge_and_sync=False) + manager.sync_weights(merge_and_sync=False) + + assert sampler.loaded == [[('q_proj.weight', 'base')], [('q_proj.lora_A', 'adapter')]] + assert model.peft_config_calls == 1 + assert sampler.calls[0]['peft_config'] is None + assert sampler.calls[1]['peft_config'] == {'r': 8, 'target_modules': ['q_proj']} + assert sampler.calls[0]['base_sync_done'] is False + assert sampler.calls[1]['base_sync_done'] is True + assert model.generator_calls[1]['base_sync_done'] is True + + +def test_local_merge_sync_generates_a_full_weight_set_each_time(): + model = _Model([('q_proj.weight', 'merged')]) + sampler = _Sampler() + manager = CheckpointEngineManager(model, sampler, platform='CPU') + + manager.sync_weights(merge_and_sync=True) + manager.sync_weights(merge_and_sync=True) + + assert sampler.loaded == [[('q_proj.weight', 'merged')], [('q_proj.weight', 'merged')]] + assert all(call['merge_and_sync'] is True for call in model.generator_calls) + assert [call['base_sync_done'] for call in model.generator_calls] == [False, True] + + +def test_local_failure_does_not_mark_base_sync_done(): + model = _Model([('q_proj.weight', 'base')]) + error = RuntimeError('sampler failed') + sampler = _Sampler(fail=error) + manager = CheckpointEngineManager(model, sampler, platform='CPU') + + with pytest.raises(RuntimeError, match='sampler failed') as exc_info: + manager.sync_weights(merge_and_sync=False) + + assert exc_info.value is error + assert manager.base_sync_done is False + + sampler.fail = None + manager.sync_weights(merge_and_sync=False) + assert model.generator_calls[-1]['base_sync_done'] is False + + +def test_local_weight_generator_failure_does_not_mark_base_sync_done(): + error = ValueError('weight generation failed') + + def broken_weights(): + yield 'q_proj.weight', 'base' + raise error + + model = _Model(broken_weights()) + sampler = _Sampler() + manager = CheckpointEngineManager(model, sampler, platform='CPU') + + with pytest.raises(ValueError, match='weight generation failed') as exc_info: + manager.sync_weights(merge_and_sync=False) + + assert exc_info.value is error + assert manager.base_sync_done is False + + +def test_mixed_deployment_shape_fails_at_initialization(): + model = _Model([]) + sampler = _Sampler() + sampler._actors = [object()] + + with pytest.raises(ValueError, match='same deployment shape'): + CheckpointEngineManager(model, sampler, platform='CPU') + + +@pytest.mark.parametrize( + ('mode', 'use_actors', 'match'), + [ + ('naive', True, "mode='naive' requires local"), + ('colocate', False, "mode='colocate' requires"), + ('standalone', False, "mode='standalone' requires"), + ('unknown', False, 'Unknown checkpoint engine mode'), + ], +) +def test_explicit_mode_validates_deployment_shape(mode, use_actors, match): + model = _Model([]) + sampler = _Sampler() + if use_actors: + model._actors = [object()] + sampler._actors = [object()] + + with pytest.raises(ValueError, match=match): + CheckpointEngineManager(model, sampler, platform='CPU', mode=mode) + + +def test_backend_selection_uses_resolved_mode(): + from twinkle.checkpoint_engine import IPCCheckpointEngine, NCCLCheckpointEngine + + assert CheckpointEngineManager.decide_backend_engine('GPU', mode='naive') is None + assert CheckpointEngineManager.decide_backend_engine('GPU', mode='colocate') is IPCCheckpointEngine + assert CheckpointEngineManager.decide_backend_engine('GPU', mode='standalone') is NCCLCheckpointEngine + + +@pytest.mark.parametrize('sampler_backend', ['vllm', 'sglang']) +@pytest.mark.parametrize('provide_weights', [True, False]) +def test_sampler_receive_weights_selects_direct_or_checkpoint_stream(sampler_backend, provide_weights): + if sampler_backend == 'vllm': + from twinkle.sampler.vllm_sampler.vllm_sampler import vLLMSampler as sampler_cls + else: + from twinkle.sampler.sglang_sampler.sglang_sampler import SGLangSampler as sampler_cls + + class _InferenceEngine: + + def __init__(self): + self.loaded = None + self.invalidated = False + + async def update_weights(self, weights, **kwargs): + self.loaded = (list(weights), kwargs) + + def invalidate_synced_lora(self): + self.invalidated = True + + class _CheckpointEngine: + + def receive_weights(self): + return iter([('checkpoint.weight', 'checkpoint')]) + + sampler = object.__new__(sampler_cls) + sampler.engine = _InferenceEngine() + sampler._run_in_loop = asyncio.run + checkpoint_engine = _CheckpointEngine() + checkpoint_engine_calls = [] + + def get_checkpoint_engine(): + checkpoint_engine_calls.append(True) + return checkpoint_engine + + sampler._get_or_create_checkpoint_engine = get_checkpoint_engine + direct_weights = iter([('direct.weight', 'direct')]) if provide_weights else None + + sampler_cls.receive_weights.__wrapped__(sampler, weights=direct_weights) + + expected_weights = [('direct.weight', 'direct')] if provide_weights else [('checkpoint.weight', 'checkpoint')] + assert sampler.engine.loaded == (expected_weights, {'peft_config': None, 'base_sync_done': False}) + assert len(checkpoint_engine_calls) == (0 if provide_weights else 1) + assert sampler.engine.invalidated is (sampler_backend == 'vllm') + + +def test_auto_ray_actor_sync_keeps_standalone_checkpoint_engine_lifecycle(monkeypatch): + events = [] + + class _Backend: + + @classmethod + def build_topology(cls, trainer_world_size, rollout_world_size, metadata): + events.append(('build_topology', trainer_world_size, rollout_world_size)) + return ({'rank': [0], 'world_size': [2], 'master_metadata': [metadata[0]]}, + {'rank': [1], 'world_size': [2], 'master_metadata': [metadata[0]]}) + + class _ActorModel(_Model): + + def __init__(self): + super().__init__([('q_proj.weight', 'base')]) + self._actors = [object()] + + def prepare_checkpoint_engine(self, is_master): + events.append(('model_prepare', is_master)) + return {'zmq_ip': '127.0.0.1', 'zmq_port': 1} + + def init_checkpoint_process_group(self, **kwargs): + events.append(('model_init_submitted', kwargs)) + return lambda: events.append(('model_init_waited', kwargs)) + + def send_weights(self, **kwargs): + events.append(('send_submitted', kwargs)) + return lambda: events.append(('send_waited', kwargs)) + + def finalize_checkpoint_engine(self): + events.append('model_finalize') + + class _ActorSampler(_Sampler): + + def __init__(self): + super().__init__() + self._actors = [object()] + + def prepare_checkpoint_engine(self, is_master): + events.append(('sampler_prepare', is_master)) + + def init_checkpoint_process_group(self, **kwargs): + events.append(('sampler_init_submitted', kwargs)) + return lambda: events.append(('sampler_init_waited', kwargs)) + + def receive_weights(self, **kwargs): + events.append(('receive_submitted', kwargs)) + return lambda: events.append(('receive_waited', kwargs)) + + def finalize_checkpoint_engine(self): + events.append('sampler_finalize') + + model = _ActorModel() + sampler = _ActorSampler() + + def decide_backend(platform=None, mode='standalone'): + assert mode == 'standalone' + return _Backend + + monkeypatch.setattr(CheckpointEngineManager, 'decide_backend_engine', staticmethod(decide_backend)) + + manager = CheckpointEngineManager(model, sampler, platform='GPU') + manager.sync_weights() + + assert manager.requested_mode == 'auto' + assert manager.mode == 'standalone' + assert [event if isinstance(event, str) else event[0] for event in events] == [ + 'model_prepare', + 'sampler_prepare', + 'build_topology', + 'model_init_submitted', + 'sampler_init_submitted', + 'model_init_waited', + 'sampler_init_waited', + 'send_submitted', + 'receive_submitted', + 'send_waited', + 'receive_waited', + 'model_finalize', + 'sampler_finalize', + ] + receive_kwargs = next( + event[1] + for event in events + if isinstance(event, tuple) and event[0] == 'receive_submitted' + ) + assert 'weights' not in receive_kwargs + assert manager.base_sync_done is True + + +def test_colocate_mode_configures_actor_backends(monkeypatch): + configured = [] + backend = object() + + class _ActorRole: + _actors = [object()] + + def set_checkpoint_engine_backend(self, name): + configured.append(name) + + def decide_backend(platform=None, mode='standalone'): + assert platform == 'GPU' + assert mode == 'colocate' + return backend + + monkeypatch.setattr(CheckpointEngineManager, 'decide_backend_engine', staticmethod(decide_backend)) + + manager = CheckpointEngineManager(_ActorRole(), _ActorRole(), platform='GPU', mode='colocate') + + assert manager.mode == 'colocate' + assert manager.backend_cls is backend + assert configured == ['ipc', 'ipc'] diff --git a/tests/sampler/test_ipc_checkpoint_engine.py b/tests/sampler/test_ipc_checkpoint_engine.py index 0e7532b26..70158a8dc 100644 --- a/tests/sampler/test_ipc_checkpoint_engine.py +++ b/tests/sampler/test_ipc_checkpoint_engine.py @@ -234,7 +234,7 @@ def weights(): def test_the_manager_picks_this_engine_only_when_colocating(): from twinkle.checkpoint_engine import CheckpointEngineManager, IPCCheckpointEngine, NCCLCheckpointEngine - assert CheckpointEngineManager.decide_backend_engine('GPU', colocate=True) is IPCCheckpointEngine + assert CheckpointEngineManager.decide_backend_engine('GPU', mode='colocate') is IPCCheckpointEngine assert CheckpointEngineManager.decide_backend_engine('GPU') is NCCLCheckpointEngine From d6312ca5fd5fe5f5d9d7c91a97c54d17e0b6ba20 Mon Sep 17 00:00:00 2001 From: weikaiwen <34648228+kevssim@users.noreply.github.com> Date: Fri, 28 Aug 2026 15:15:24 +0800 Subject: [PATCH 2/3] wip --- src/twinkle/checkpoint_engine/manager.py | 6 ++---- 1 file changed, 2 insertions(+), 4 deletions(-) diff --git a/src/twinkle/checkpoint_engine/manager.py b/src/twinkle/checkpoint_engine/manager.py index 9f7a24666..cdf30a0d8 100644 --- a/src/twinkle/checkpoint_engine/manager.py +++ b/src/twinkle/checkpoint_engine/manager.py @@ -127,8 +127,6 @@ def decide_backend_engine( platform_name = Platform.get_platform(platform).__name__ if mode == 'colocate': - if platform_name != 'GPU': - raise NotImplementedError("mode='colocate' currently requires the GPU platform.") from twinkle.checkpoint_engine import IPCCheckpointEngine return IPCCheckpointEngine if mode != 'standalone': @@ -166,8 +164,8 @@ def sync_weights(self, merge_and_sync=True): self._sync_weights_naive(merge_and_sync) return - is_master = [True] + [False] * (self.model.device_mesh.world_size - 1) - model_metadata = self.model.prepare_checkpoint_engine(is_master) + model_metadata = self.model.prepare_checkpoint_engine([True] + + [False] * (self.model.device_mesh.world_size - 1)) self.sampler.prepare_checkpoint_engine(False) model_kwargs, sampler_kwargs = self.backend_cls.build_topology( self.model.device_mesh.world_size, From 6a62babd0e08c6a0118abb70f69de7780f45cbbe Mon Sep 17 00:00:00 2001 From: weikaiwen <34648228+kevssim@users.noreply.github.com> Date: Mon, 31 Aug 2026 09:42:31 +0800 Subject: [PATCH 3/3] wip --- .../ipc_checkpoint_engine.py | 72 +++++++-- .../sampler/vllm_sampler/vllm_engine.py | 33 ++-- .../vllm_sampler/vllm_worker_extension.py | 32 ++-- src/twinkle/utils/platforms/npu.py | 59 ++++++- ...st_ipc_checkpoint_engine_device_neutral.py | 149 ++++++++++++++++++ tests/utils/test_npu_ipc_support.py | 39 +++++ 6 files changed, 339 insertions(+), 45 deletions(-) create mode 100644 tests/checkpoint_engine/test_ipc_checkpoint_engine_device_neutral.py create mode 100644 tests/utils/test_npu_ipc_support.py diff --git a/src/twinkle/checkpoint_engine/ipc_checkpoint_engine.py b/src/twinkle/checkpoint_engine/ipc_checkpoint_engine.py index b008f15e8..2353e0e4b 100644 --- a/src/twinkle/checkpoint_engine/ipc_checkpoint_engine.py +++ b/src/twinkle/checkpoint_engine/ipc_checkpoint_engine.py @@ -31,7 +31,8 @@ import zmq from typing import Any, AsyncGenerator, Generator -from twinkle import get_logger +from twinkle import Platform, get_logger +from twinkle.utils.framework import Torch from .base import CheckpointEngine logger = get_logger() @@ -49,7 +50,7 @@ class IPCCheckpointEngine(CheckpointEngine): - """Hand weights to a sampler on the same GPU by mapping memory instead of copying it.""" + """Hand weights to a sampler on the same device by mapping memory instead of copying it.""" def __init__(self, bucket_size: int = 512 << 20, **kwargs) -> None: # Smaller default than the NCCL engine's 3 GB: a bigger bucket buys nothing when the transfer @@ -65,21 +66,22 @@ def __init__(self, bucket_size: int = 512 << 20, **kwargs) -> None: self.socket = None self._context = None self._handle = None + self._shm = None # Receiver side: the mapping of the sender's buffer, kept across buckets. Re-mapping per # bucket is what makes device memory appear to grow during a sync. self._mapped: torch.Tensor | None = None + self._mapped_shms = [] self._mapped_signature = None # ── rendezvous ─────────────────────────────────────────────────────── @staticmethod def endpoint() -> str: - """The socket both peers derive independently from the GPU they share. + """The socket both peers derive independently from the device they share. - The device's UUID rather than its index: under Ray each role sees its own GPU as index 0, so - indices collide across ranks while UUIDs do not. + The platform helper obtains the physical device UUID for the current local device. """ - uuid = str(torch.cuda.get_device_properties(torch.cuda.current_device()).uuid) + uuid = str(Platform.get_vllm_device_uuid(Torch.get_current_device())) return f'ipc:///tmp/twinkle-colocate-{uuid}.sock' def prepare(self) -> dict[str, Any]: @@ -165,9 +167,17 @@ def finalize(self): path = self.endpoint().removeprefix('ipc://') if os.path.exists(path): os.unlink(path) + if self._shm is not None: + self.send_buf = None + self._shm.close() + self._shm.unlink() + self._shm = None self.send_buf = None self._handle = None self._mapped = None + for shm in self._mapped_shms: + shm.close() + self._mapped_shms.clear() self._mapped_signature = None self.rank = None @@ -178,9 +188,29 @@ def _ensure_buffer(self, min_size: int) -> None: if self.send_buf is not None and self.send_buf.numel() >= min_size: return size = max(self.bucket_size, min_size) - self.send_buf = torch.empty(size, dtype=torch.uint8, device=torch.cuda.current_device()) + platform = Platform.get_platform() + if platform.device_prefix() == 'npu' and not platform.is_ipc_supported(): + from multiprocessing import shared_memory + + if self._shm is not None: + self.send_buf = None + self._shm.close() + self._shm.unlink() + self._shm = None + self._shm = shared_memory.SharedMemory(create=True, size=size) + self.send_buf = torch.frombuffer(self._shm.buf, dtype=torch.uint8, count=size) + self._handle = {'name': self._shm.name, 'size': size} + return + + self.send_buf = torch.empty( + size, + dtype=torch.uint8, + device=f'{platform.device_prefix()}:{Torch.get_current_device()}', + ) # One handle per buffer, reused for every bucket: the buffer is refilled, not reallocated, so # the mapping stays valid and the receiver can keep it. + if platform.device_prefix() == 'npu': + import torch_npu # noqa: F401 from torch.multiprocessing.reductions import reduce_tensor self._handle = reduce_tensor(self.send_buf) @@ -223,7 +253,7 @@ async def send_weights(self, weights: Generator[tuple[str, torch.Tensor], None, def _flush(self, bucket_meta: list[dict], is_last: bool) -> None: """Publish the filled part of the buffer and wait until the receiver is done with it.""" # The copies above are non_blocking; without this the receiver could map bytes not yet written. - torch.cuda.synchronize() + Torch.synchronize() self.socket.send(pickle.dumps({'handle': self._handle, 'bucket_meta': bucket_meta, 'is_last': is_last})) # The receiver copies out of this buffer, so it must say so before we overwrite it. self.socket.recv() @@ -245,7 +275,7 @@ async def receive_weights(self) -> AsyncGenerator[tuple[str, torch.Tensor], None yield meta['name'], buffer[start:start + nbytes].view(meta['dtype']).view(meta['shape']) # Consumers copy with non_blocking=True, so the acknowledgement has to wait for the copies # and not merely for the loop above. - torch.cuda.synchronize() + Torch.synchronize() self.socket.send(b'ack') if message['is_last']: break @@ -259,13 +289,25 @@ def _map(self, handle) -> torch.Tensor: signature = self._handle_signature(handle) if self._mapped is not None and signature == self._mapped_signature: return self._mapped - from torch.multiprocessing.reductions import rebuild_cuda_tensor + if isinstance(handle, dict): + from multiprocessing import shared_memory + + mapped_shm = shared_memory.SharedMemory(name=handle['name']) + self._mapped_shms.append(mapped_shm) + self._mapped = torch.frombuffer( + mapped_shm.buf, + dtype=torch.uint8, + count=handle['size'], + ) + self._mapped_signature = signature + return self._mapped + func, args = handle args = list(args) - # Both peers see the shared GPU as their own device 0, but be explicit rather than trust the - # index the sender happened to record. - args[6] = torch.cuda.current_device() - self._mapped = func(*args) if callable(func) else rebuild_cuda_tensor(*args) + if Platform.device_prefix() == 'npu': + import torch_npu # noqa: F401 + args[6] = Torch.get_current_device() + self._mapped = func(*args) self._mapped_signature = signature return self._mapped @@ -276,6 +318,8 @@ def _handle_signature(handle) -> tuple: Locally implemented rather than shared with the sampler's worker extension, which has the same helper: the sampler imports this package, so importing it back would be circular. """ + if isinstance(handle, dict): + return tuple(handle.items()) _, args = handle return tuple((type(v).__name__, bytes(v) if isinstance(v, (bytes, bytearray)) else v) for v in args if isinstance(v, (bytes, bytearray, int, float, bool, str)) or v is None) diff --git a/src/twinkle/sampler/vllm_sampler/vllm_engine.py b/src/twinkle/sampler/vllm_sampler/vllm_engine.py index b1e1790de..dc9825429 100644 --- a/src/twinkle/sampler/vllm_sampler/vllm_engine.py +++ b/src/twinkle/sampler/vllm_sampler/vllm_engine.py @@ -51,7 +51,7 @@ class VLLMEngine(BaseSamplerEngine): This engine uses vLLM v1's AsyncLLM and supports: - Tinker-compatible sample() API with logprobs - Multi-tenant LoRA adapters for client-server mode - - Weight synchronization via load_weights (colocated) or CUDA IPC + - Weight synchronization via load_weights (colocated) or device IPC - Sleep/wake_up for GPU memory management in colocated training Deployment scenarios: @@ -109,10 +109,10 @@ def __init__( # ``list_loras()`` per request. self._synced_lora_request: Optional[Any] = None - # Long-lived CUDA IPC bucket reused across all update_weights() + # Long-lived device IPC bucket reused across all update_weights() # calls. Allocating a new IPC buffer (and hence a new IPC handle) - # per sync forces every worker to create a new CUDA IPC mapping via - # ``rebuild_cuda_tensor`` because PyTorch's ``shared_cache`` cannot + # per sync forces every worker to create a new device IPC mapping via + # the reducer callable because PyTorch's ``shared_cache`` cannot # hit on unseen storage handles. The driver reclaims those mappings # lazily, which is the root cause of the slow GPU memory drift we # observed under frequent LoRA syncs. By pinning a single buffer @@ -578,14 +578,14 @@ async def update_weights( bucket_size_mb: int = 2048, **kwargs, ) -> None: - """Update model weights via ZMQ + CUDA IPC to worker extension. + """Update model weights via ZMQ + device IPC to worker extension. Accepts **either** a ``dict[str, Tensor]`` (legacy) **or** an async generator / sync generator of ``(name, tensor)`` pairs (streaming). The streaming path avoids accumulating a full model copy on GPU: tensors are consumed one-by-one from the generator, copied into a - GPU IPC bucket, and flushed to the vLLM worker subprocess when the + device IPC bucket, and flushed to the vLLM worker subprocess when the bucket is full. Args: @@ -621,15 +621,19 @@ async def _sync_iter(): weight_aiter = _sync_iter() - # Peek first tensor to detect device (GPU → IPC, CPU → SHM). + # Peek first tensor to detect device (supported accelerator → IPC, CPU → SHM). try: first_name, first_tensor = await weight_aiter.__anext__() except StopAsyncIteration: logger.warning('update_weights called with empty weights') return - use_gpu_ipc = first_tensor.is_cuda - use_shm = not use_gpu_ipc + use_device_ipc = first_tensor.is_cuda + if first_tensor.device.type == 'npu': + from twinkle.utils.platforms import NPU + + use_device_ipc = NPU.is_ipc_supported() + use_shm = not use_device_ipc # Use a per-sync unique IPC endpoint to avoid cross-actor collisions # when multiple sampler actors share the same device UUID. @@ -650,13 +654,16 @@ async def _sync_iter(): buffer = None shm = None - if use_gpu_ipc: + if use_device_ipc: + if first_tensor.device.type == 'npu': + # torch_npu registers the NPU reducer used by reduce_tensor. + import torch_npu # noqa: F401 from torch.multiprocessing.reductions import reduce_tensor # Reuse a long-lived IPC bucket whenever the requested size # fits. The handle is produced once and shipped to every # subsequent sync so each worker's ``shared_cache`` stays warm - # and no new CUDA IPC mapping is created per sync. + # and no new device IPC mapping is created per sync. need_realloc = ( self._ipc_buffer is None or self._ipc_buffer_size < bucket_size or self._ipc_buffer.device != first_tensor.device) @@ -714,7 +721,7 @@ def _zmq_send_recv(payload, where: str): )) # Send IPC/SHM handle, wait for worker ready (non-blocking) - handle_payload = ipc_handle if use_gpu_ipc else {'name': shm_name, 'size': bucket_size} + handle_payload = ipc_handle if use_device_ipc else {'name': shm_name, 'size': bucket_size} await loop.run_in_executor(None, _zmq_send_recv, handle_payload, 'handle handshake') # Stream weights into buckets and send to worker @@ -821,7 +828,7 @@ async def _flush_bucket(is_last: bool) -> None: elapsed = time.time() - start_time mode = 'LoRA' if base_sync_done and peft_config else 'base' logger.info(f'Updated {n_weights} {mode} weights via ' - f"{'IPC' if use_gpu_ipc else 'SHM'} in {elapsed:.2f}s") + f"{'IPC' if use_device_ipc else 'SHM'} in {elapsed:.2f}s") async def shutdown(self) -> None: """Shutdown the vLLM engine and release all resources. diff --git a/src/twinkle/sampler/vllm_sampler/vllm_worker_extension.py b/src/twinkle/sampler/vllm_sampler/vllm_worker_extension.py index 2ed5f5227..0a4ec3e56 100644 --- a/src/twinkle/sampler/vllm_sampler/vllm_worker_extension.py +++ b/src/twinkle/sampler/vllm_sampler/vllm_worker_extension.py @@ -45,26 +45,20 @@ def set_death_signal(): def _rebuild_ipc(handle, device_id: Optional[int] = None) -> torch.Tensor: - """Rebuild CUDA tensor from IPC handle.""" - from torch.multiprocessing.reductions import rebuild_cuda_tensor - + """Rebuild an accelerator tensor from an IPC reducer handle.""" func, args = handle list_args = list(args) if device_id is not None: list_args[6] = device_id - - if callable(func): - return func(*list_args) - else: - return rebuild_cuda_tensor(*list_args) + return func(*list_args) def _ipc_handle_signature(handle) -> Optional[tuple]: - """Derive a stable signature for a CUDA IPC handle. + """Derive a stable signature for an accelerator IPC handle. ``reduce_tensor`` returns ``(func, args)`` where ``args`` contains the - CUDA IPC storage handle bytes, storage size, ref-counter handle, etc. - Two handles are equivalent (i.e. map the same CUDA memory region) when + IPC storage handle bytes, storage size, ref-counter handle, etc. + Two handles are equivalent (i.e. map the same device memory region) when these inner fields match. We hash only the parts that are picklable and comparable to avoid accidental mismatches due to local objects. """ @@ -127,12 +121,12 @@ def update_weights_from_ipc( use_shm: bool = False, zmq_handle: Optional[str] = None, ) -> None: - """Receive and load weights via ZMQ + CUDA IPC/SHM. + """Receive and load weights via ZMQ + device IPC/SHM. Called via ``collective_rpc("update_weights_from_ipc", ...)`` from :meth:`VLLMEngine.update_weights`. The VLLMEngine sends weights - in buckets over a ZMQ REQ/REP channel backed by CUDA IPC (GPU - tensors) or shared memory (CPU tensors). + in buckets over a ZMQ REQ/REP channel backed by device IPC + (accelerator tensors) or shared memory (CPU tensors). For TP > 1, only TP rank 0 communicates with the VLLMEngine over ZMQ. It broadcasts the IPC handle and bucket metadata to other @@ -142,7 +136,7 @@ def update_weights_from_ipc( Args: peft_config: If provided with base_sync_done, loads as LoRA. base_sync_done: If True and peft_config, replaces existing LoRA. - use_shm: If True, use shared memory instead of CUDA IPC. + use_shm: If True, use shared memory instead of device IPC. zmq_handle: Optional ZMQ IPC endpoint. If None, uses _get_zmq_handle(). """ import torch.distributed as dist @@ -196,6 +190,10 @@ def _broadcast_obj(obj): # ── Step 2: Receive and broadcast IPC/SHM handle ── buffer, shm = None, None + if not use_shm and self.device.type == 'npu': + # Register the NPU reducer before recv_pyobj() unpickles its callable. + import torch_npu # noqa: F401 + if is_driver: try: comm_metadata = socket.recv_pyobj() @@ -210,9 +208,9 @@ def _broadcast_obj(obj): if not use_shm: handle = comm_metadata # All TP ranks rebuild the IPC buffer from the same handle. - # CUDA IPC allows any process on the same node to map the memory. + # Device IPC allows any process on the same node to map the memory. # Reuse a cached buffer across syncs when the sender reuses the - # same IPC handle: this avoids creating a fresh CUDA IPC mapping + # same IPC handle: this avoids creating a fresh device IPC mapping # per sync, which the driver releases lazily and is the root # cause of the apparent GPU memory growth under frequent syncs. handle_signature = _ipc_handle_signature(handle) diff --git a/src/twinkle/utils/platforms/npu.py b/src/twinkle/utils/platforms/npu.py index de15707f6..fb76f1293 100644 --- a/src/twinkle/utils/platforms/npu.py +++ b/src/twinkle/utils/platforms/npu.py @@ -1,11 +1,14 @@ # Copyright (c) ModelScope Contributors. All rights reserved. import hashlib import os +import platform import re import socket import subprocess from typing import Optional +from packaging import version + from .base import Platform # ref: https://www.hiascend.com/document/detail/zh/canncommercial/83RC1/maintenref/envvar/envref_07_0144.html @@ -16,7 +19,6 @@ # NPU-side socket port pool used by HCCL for device communication channels. _HCCL_NPU_SOCKET_PORT_RANGE_ENV = 'HCCL_NPU_SOCKET_PORT_RANGE' - def _derive_hccl_socket_env_defaults(master_port: int) -> dict: """Derive deterministic default HCCL socket env values from master_port.""" # Keep values stable per job and spread jobs across non-overlapping ranges. @@ -99,6 +101,61 @@ def ensure_npu_backend() -> None: class NPU(Platform): + @staticmethod + def is_ipc_supported() -> bool: + """Return whether HDK and CANN meet the NPU device-IPC requirement.""" + try: + result = subprocess.run( + ['npu-smi', 'info', '-t', 'board', '-i', '1'], capture_output=True, text=True, check=True) + except subprocess.CalledProcessError: + visible_devices = (os.environ.get('ASCEND_VISIBLE_DEVICES') + or os.environ.get('ASCEND_RT_VISIBLE_DEVICES')) + if not visible_devices: + raise + device_id = int(visible_devices.split(',')[0]) + try: + result = subprocess.run( + ['npu-smi', 'info', '-t', 'board', '-i', str(device_id)], + capture_output=True, + text=True, + check=True, + ) + except subprocess.CalledProcessError: + result = subprocess.run( + ['npu-smi', 'info', '-t', 'board', '-i', str(device_id // 2)], + capture_output=True, + text=True, + check=True, + ) + + software_version = next( + (line.split(':', 1)[1].strip().lower() + for line in result.stdout.splitlines() if 'Software Version' in line), + None, + ) + if software_version is None: + raise RuntimeError('Could not find Software Version in npu-smi output') + + ascend_home = os.environ.get('ASCEND_HOME_PATH', '/usr/local/Ascend/ascend-toolkit/latest') + info_file = os.path.join(ascend_home, f'{platform.machine()}-linux', 'ascend_toolkit_install.info') + with open(info_file) as info: + cann_version = next( + (line.split('=', 1)[1].strip().lower() for line in info if line.startswith('version=')), + None, + ) + if cann_version is None: + raise RuntimeError('Could not find version in CANN toolkit info file') + + pattern = r'(\d+\.\d+(?=\.t))|(\d+\.\d+(?:\.(?:rc\d+|\d+))?)' + software_match = re.match(pattern, software_version) + cann_match = re.match(pattern, cann_version) + if software_match is None or cann_match is None: + raise RuntimeError(f'Invalid NPU versions: HDK={software_version}, CANN={cann_version}') + software_base = software_match.group(1) or software_match.group(2) + cann_base = cann_match.group(1) or cann_match.group(2) + return (version.parse(software_base) >= version.parse('25.3.rc1') + and version.parse(cann_base) >= version.parse('8.3.rc1')) + @staticmethod def visible_device_env(): # Ascend runtime uses ASCEND_RT_VISIBLE_DEVICES. diff --git a/tests/checkpoint_engine/test_ipc_checkpoint_engine_device_neutral.py b/tests/checkpoint_engine/test_ipc_checkpoint_engine_device_neutral.py new file mode 100644 index 000000000..396cee822 --- /dev/null +++ b/tests/checkpoint_engine/test_ipc_checkpoint_engine_device_neutral.py @@ -0,0 +1,149 @@ +"""CPU-only coverage for the device-neutral parts of the IPC checkpoint engine.""" + +from unittest.mock import Mock + +import pytest +import torch + +import twinkle.checkpoint_engine.ipc_checkpoint_engine as ipc_module +from twinkle.checkpoint_engine import IPCCheckpointEngine + + +def test_endpoint_uses_platform_uuid_without_touching_cuda(monkeypatch): + """Endpoint construction must not query the current CUDA device.""" + monkeypatch.setattr(torch.cuda, 'current_device', Mock(side_effect=AssertionError('CUDA initialized'))) + monkeypatch.setattr(torch.cuda, 'get_device_properties', Mock(side_effect=AssertionError('CUDA initialized'))) + monkeypatch.setattr(ipc_module.Torch, 'get_current_device', lambda: 0) + monkeypatch.setattr( + ipc_module.Platform, + 'get_vllm_device_uuid', + staticmethod(lambda device_id=0, platform=None: f'uuid-{device_id}'), + ) + + assert IPCCheckpointEngine.endpoint() == 'ipc:///tmp/twinkle-colocate-uuid-0.sock' + + +def test_map_uses_reducer_callable_and_receiver_device(monkeypatch): + calls = [] + + def rebuild(*args): + calls.append(args) + return 'mapped' + + monkeypatch.setattr(ipc_module.Platform, 'device_prefix', staticmethod(lambda: 'cuda')) + monkeypatch.setattr(ipc_module.Torch, 'get_current_device', lambda: 3) + sender_args = [None, None, None, None, None, None, 17] + + engine = IPCCheckpointEngine() + assert engine._map((rebuild, sender_args)) == 'mapped' + assert calls[0][6] == 3 + assert sender_args[6] == 17 + + +def test_vllm_worker_uses_reducer_callable_and_receiver_device(): + from twinkle.sampler.vllm_sampler.vllm_worker_extension import _rebuild_ipc + + calls = [] + + def rebuild(*args): + calls.append(args) + return 'mapped' + + sender_args = [None, None, None, None, None, None, 17] + + assert _rebuild_ipc((rebuild, sender_args), device_id=3) == 'mapped' + assert calls[0][6] == 3 + assert sender_args[6] == 17 + + +def test_buffer_and_flush_use_framework_device_and_sync_helpers(monkeypatch): + class CPU: + @staticmethod + def device_prefix(): + return 'cpu' + + monkeypatch.setattr(ipc_module.Platform, 'get_platform', staticmethod(lambda platform=None: CPU)) + monkeypatch.setattr(ipc_module.Torch, 'get_current_device', lambda: 0) + + from torch.multiprocessing import reductions + + monkeypatch.setattr(reductions, 'reduce_tensor', lambda tensor: (None, [None])) + engine = IPCCheckpointEngine(bucket_size=8) + engine._ensure_buffer(8) + assert engine.send_buf.device.type == 'cpu' + + class Socket: + def send(self, payload): + self.payload = payload + + def recv(self): + return b'ack' + + engine.socket = Socket() + synchronize = Mock() + monkeypatch.setattr(ipc_module.Torch, 'synchronize', synchronize) + engine._flush([], is_last=True) + synchronize.assert_called_once_with() + + +def test_npu_without_device_ipc_uses_shared_memory_and_closes_it(monkeypatch): + class NPU: + @staticmethod + def device_prefix(): + return 'npu' + + @staticmethod + def is_ipc_supported(): + return False + + monkeypatch.setattr(ipc_module.Platform, 'get_platform', staticmethod(lambda platform=None: NPU)) + monkeypatch.setattr(ipc_module.Torch, 'get_current_device', lambda: 0) + + sender = IPCCheckpointEngine(bucket_size=8) + sender._ensure_buffer(8) + handle = sender._handle + assert set(handle) == {'name', 'size'} + assert sender.send_buf.device.type == 'cpu' + + receiver = IPCCheckpointEngine() + mapped = receiver._map(handle) + assert mapped.device.type == 'cpu' + assert mapped.numel() == 8 + + sender.finalize() + del mapped + receiver.finalize() + + from multiprocessing import shared_memory + + with pytest.raises(FileNotFoundError): + shared_memory.SharedMemory(name=handle['name']) + + +def test_shared_memory_growth_keeps_the_previous_receiver_mapping_alive(monkeypatch): + class NPU: + @staticmethod + def device_prefix(): + return 'npu' + + @staticmethod + def is_ipc_supported(): + return False + + monkeypatch.setattr(ipc_module.Platform, 'get_platform', staticmethod(lambda platform=None: NPU)) + + sender = IPCCheckpointEngine(bucket_size=8) + receiver = IPCCheckpointEngine() + sender._ensure_buffer(8) + first = receiver._map(sender._handle) + first[0] = 7 + + sender._ensure_buffer(16) + second = receiver._map(sender._handle) + + assert first[0].item() == 7 + assert second.numel() == 16 + + del first, second + sender.finalize() + receiver.finalize() diff --git a/tests/utils/test_npu_ipc_support.py b/tests/utils/test_npu_ipc_support.py new file mode 100644 index 000000000..f61d1594a --- /dev/null +++ b/tests/utils/test_npu_ipc_support.py @@ -0,0 +1,39 @@ +from types import SimpleNamespace + +import pytest + +from twinkle.utils.platforms import npu + + +@pytest.mark.parametrize( + ('software_version', 'cann_version', 'expected'), + [ + ('25.3.rc1.2', '8.3.rc1', True), + ('25.5.t3.b001', '8.3.0', True), + ('25.2.0', '8.3.rc1', False), + ('25.3.rc1', '8.2.0', False), + ], +) +def test_npu_ipc_version_gate(monkeypatch, tmp_path, software_version, cann_version, expected): + monkeypatch.setattr( + npu.subprocess, + 'run', + lambda *args, **kwargs: SimpleNamespace(stdout=f'Software Version : {software_version}\n'), + ) + monkeypatch.setattr(npu.platform, 'machine', lambda: 'x86_64') + cann_dir = tmp_path / 'x86_64-linux' + cann_dir.mkdir() + (cann_dir / 'ascend_toolkit_install.info').write_text(f'version={cann_version}\n') + monkeypatch.setenv('ASCEND_HOME_PATH', str(tmp_path)) + + assert npu.NPU.is_ipc_supported() is expected + + +def test_npu_ipc_detection_errors_are_not_hidden(monkeypatch): + def fail(*args, **kwargs): + raise RuntimeError('npu-smi failed') + + monkeypatch.setattr(npu.subprocess, 'run', fail) + + with pytest.raises(RuntimeError, match='npu-smi failed'): + npu.NPU.is_ipc_supported()