Repository navigation
Conversation
Signed-off-by: wangx700 <wangxin700@huawei.com>
Signed-off-by: wangx700 <wangxin700@huawei.com>
There was a problem hiding this comment.
Code Review
This pull request introduces Delta Weight Sync, a feature designed to optimize weight synchronization for non-colocated rollout engines by transferring only changed bytes instead of full checkpoints. It implements both disk and sparse_hccl transports, along with corresponding configuration options, validation rules, documentation, and unit tests. The review feedback focuses on improving backward compatibility and resource management. Specifically, it recommends replacing PyTorch 2.4+ torch.accelerator calls and Python 3.10+ zip(..., strict=True) usages with backward-compatible alternatives to support older environments. It also suggests using Ray's public API for retrieving node IP addresses and reusing a persistent, properly managed ThreadPoolExecutor to avoid thread creation overhead and prevent resource leaks.
| assert slots is not None | ||
| if record.param is None: | ||
| dtype = record.info.dtype | ||
| local_indices = torch.empty(0, dtype=torch.int64, device=torch.accelerator.current_device()) |
There was a problem hiding this comment.
torch.accelerator is only available in PyTorch 2.4+. Since many Ascend CANN environments run on older PyTorch versions (e.g., 2.1 or 2.2), using torch.accelerator.current_device() will cause runtime errors. A backward-compatible alternative using torch.npu or torch.cuda should be used instead.
| local_indices = torch.empty(0, dtype=torch.int64, device=torch.accelerator.current_device()) | |
| device = torch.device("npu", torch.npu.current_device()) if is_npu() else torch.device("cuda", torch.cuda.current_device()) | |
| local_indices = torch.empty(0, dtype=torch.int64, device=device) |
| return | ||
| assert self._client is not None | ||
| self._client.finish_weight_update(str(self.weight_version)) | ||
| torch.accelerator.synchronize() |
There was a problem hiding this comment.
torch.accelerator is only available in PyTorch 2.4+. To maintain backward compatibility with older PyTorch versions commonly used in Ascend CANN environments, use torch.npu.synchronize() or torch.cuda.synchronize() instead.
if hasattr(torch, "npu") and torch.npu.is_available():
torch.npu.synchronize()
else:
torch.cuda.synchronize()| self._ensure_export_index() | ||
| assert self._index is not None | ||
| prime_delta_snapshots(self._index, self._snapshots, pin=False) | ||
| torch.accelerator.synchronize() |
There was a problem hiding this comment.
torch.accelerator is only available in PyTorch 2.4+. To maintain backward compatibility with older PyTorch versions commonly used in Ascend CANN environments, use torch.npu.synchronize() or torch.cuda.synchronize() instead.
| torch.accelerator.synchronize() | |
| if hasattr(torch, "npu") and torch.npu.is_available(): | |
| torch.npu.synchronize() | |
| else: | |
| torch.cuda.synchronize() |
| local_offsets = [0] | ||
| for slot_index in range(slot_count): | ||
| local_offsets.append(local_offsets[-1] + counts_cpu[rank][slot_index]) | ||
| for lo, hi in zip(cuts[:-1], cuts[1:], strict=False): |
There was a problem hiding this comment.
The strict keyword argument in zip was introduced in Python 3.10. Since many Ascend CANN environments still run on Python 3.8 or 3.9, using strict=False will raise a TypeError. Since strict=False is the default behavior of zip, it can be safely omitted to ensure backward compatibility.
| for lo, hi in zip(cuts[:-1], cuts[1:], strict=False): | |
| for lo, hi in zip(cuts[:-1], cuts[1:]): |
| if self.is_wire_master: | ||
| for slots, counts, indices, values in batch: | ||
| offset = 0 | ||
| for (name, shape), count in zip(slots, counts.tolist(), strict=True): |
There was a problem hiding this comment.
The strict keyword argument in zip was introduced in Python 3.10. Since many Ascend CANN environments still run on Python 3.8 or 3.9, using strict=True will raise a TypeError. We can replace it with an explicit length assertion followed by a standard zip to ensure backward compatibility.
assert len(slots) == len(counts), "slots and counts length mismatch"
for (name, shape), count in zip(slots, counts.tolist()):| return | ||
| position_offset = value_offset = 0 | ||
| params = [] | ||
| for patch, shape in zip(self.patches, self.shapes, strict=True): |
There was a problem hiding this comment.
The strict keyword argument in zip was introduced in Python 3.10. Since many Ascend CANN environments still run on Python 3.8 or 3.9, using strict=True will raise a TypeError. We can replace it with an explicit length assertion followed by a standard zip to ensure backward compatibility.
| for patch, shape in zip(self.patches, self.shapes, strict=True): | |
| assert len(self.patches) == len(self.shapes), "patches and shapes length mismatch" | |
| for patch, shape in zip(self.patches, self.shapes): |
| with socket.socket() as sock: | ||
| sock.bind(("", 0)) | ||
| rendezvous[0] = ( | ||
| ray._private.services.get_node_ip_address(), |
There was a problem hiding this comment.
Relying on Ray's private APIs (such as ray._private.services.get_node_ip_address()) is discouraged as they can change or be removed in future Ray releases without notice. Use the public and stable ray.util.get_node_ip_address() instead.
| ray._private.services.get_node_ip_address(), | |
| ray.util.get_node_ip_address(), |
| def disconnect_rollout_engines(self) -> None: | ||
| self._group = None | ||
| self._client = None |
There was a problem hiding this comment.
Clean up the persistent thread pool executor when disconnecting rollout engines to prevent resource leaks.
| def disconnect_rollout_engines(self) -> None: | |
| self._group = None | |
| self._client = None | |
| def disconnect_rollout_engines(self) -> None: | |
| self._group = None | |
| self._client = None | |
| if getattr(self, "_executor", None) is not None: | |
| self._executor.shutdown(wait=True) | |
| self._executor = None |
| executor = ThreadPoolExecutor(max_workers=1) | ||
| try: | ||
| future = executor.submit(self._client.update_weights, asdict(update_info)) | ||
| # The HTTP control plane should enter receive_weights before the | ||
| # trainer starts the matching HCCL collectives. Surface an | ||
| # immediately rejected request instead of blocking forever in the | ||
| # first broadcast with no receiver. | ||
| time.sleep(2.0) | ||
| if future.done(): | ||
| future.result() | ||
| SparseHCCLWeightTransferEngine.trainer_send_weights( | ||
| iter(patches), | ||
| SparseHCCLTrainerSendWeightsArgs(group=self._group), | ||
| ) | ||
| future.result() | ||
| finally: | ||
| executor.shutdown(wait=False) |
There was a problem hiding this comment.
Creating and shutting down a new ThreadPoolExecutor on every call to _publish introduces significant thread creation overhead, especially since _publish is called for every chunk/flush. We should lazily initialize and reuse a single persistent executor instead.
| executor = ThreadPoolExecutor(max_workers=1) | |
| try: | |
| future = executor.submit(self._client.update_weights, asdict(update_info)) | |
| # The HTTP control plane should enter receive_weights before the | |
| # trainer starts the matching HCCL collectives. Surface an | |
| # immediately rejected request instead of blocking forever in the | |
| # first broadcast with no receiver. | |
| time.sleep(2.0) | |
| if future.done(): | |
| future.result() | |
| SparseHCCLWeightTransferEngine.trainer_send_weights( | |
| iter(patches), | |
| SparseHCCLTrainerSendWeightsArgs(group=self._group), | |
| ) | |
| future.result() | |
| finally: | |
| executor.shutdown(wait=False) | |
| if getattr(self, "_executor", None) is None: | |
| self._executor = ThreadPoolExecutor(max_workers=1) | |
| try: | |
| future = self._executor.submit(self._client.update_weights, asdict(update_info)) | |
| # The HTTP control plane should enter receive_weights before the | |
| # trainer starts the matching HCCL collectives. Surface an | |
| # immediately rejected request instead of blocking forever in the | |
| # first broadcast with no receiver. | |
| time.sleep(2.0) | |
| if future.done(): | |
| future.result() | |
| SparseHCCLWeightTransferEngine.trainer_send_weights( | |
| iter(patches), | |
| SparseHCCLTrainerSendWeightsArgs(group=self._group), | |
| ) | |
| future.result() | |
| except Exception: | |
| raise |
Documentation build overview
48 files changed ·
|
Signed-off-by: wangx700 <wangxin700@huawei.com>
Signed-off-by: wangx700 <wangxin700@huawei.com>
aae4cc6 to
7713ec7
Compare
Summary
Port of the sparse HCCL weight-sync feature (wangx700/vime#1) onto the
ascendbranch. After an initial dense synchronization, subsequent updates transfer only the changed weight elements between Megatron training workers and vLLM-Ascend rollout workers.What this PR adds
sparse_hcclfor--update-weight-mode=delta(alongsidedisk): a dense HF weight seed is sent once, then final-HF int32 positions and values are transferred via HCCL and patched in place on the rollout side.vime/backends/megatron_utils/update_weight/:delta_spec.py,delta_sync/(encode,sparse_gather),megatron_delta_export.py,update_weight_from_sparse_hccl.py.--update-weight-transportgains thesparse_hcclchoice; new--update-weight-delta-batch-gather/--update-weight-delta-verify-everyoptions and delta-mode validation.tests/test_sparse_hccl_delta_sync.py,tests/test_update_weight_factory.py).Adaptations for the
ascendbranchThe
ascendbranch does not yet carry the delta weight-sync foundation that exists onmain, so a few minimal anchors were added here to make the feature self-contained (kept as close to upstreammainas possible):update_weight/__init__.pyandupdate_weight/common.py: bring increate_weight_updaterandVimeRayWeightSyncClient(the latter identical to the version onmain; the former ismain's factory plus the newsparse_hcclbranch).backends/megatron_utils/actor.py: routedelta + sparse_hccltoUpdateWeightFromSparseHCCLin this branch's inline updater selection.backends/vllm_utils/vllm_engine.py: select thesparse_hcclweight-transfer backend in--weight-transfer-configwhen delta mode is configured.utils/arguments.py: introduce the delta-sync CLI arguments this feature extends.megatron_delta_export.py: reworked Bridge-based export (build_export_index(bridge, model)/trim_hf_vocab_padding/ sparse probe), carried verbatim from the source branch.hf_weight_iterator_sparse_bridge.py: the Bridge streaming iterator used by the sparse exporter, namedHfWeightIteratorSparseBridgebecause this branch already ships a differentHfWeightIteratorBridgefor theHfWeightIteratorBaseregistry (both classes coexist).backends/megatron_utils/actor.py: keep weights live acrosswake_up()/sleep()when delta +sparse_hcclis configured (avoids a reload around each weight update).VLLM_USE_V2_MODEL_RUNNERenv-default tweak from the source branch has no anchor on this branch and was skipped.Dependencies / notes
vllm_ascend.distributed.weight_transfer.sparse_hccl_engineandvllm_ascend.distributed.weight_transfer.sparse_weight_patchon the vLLM-Ascend side (separate PR).run-qwen3-4B-delta-sparse-hccl.shfrom the source branch is intentionally omitted.