Skip to content

[Feature][Ascend] Support sparse HCCL weight synchronization from Megatron training workers - #447

Draft
wangx700 wants to merge 6 commits into
vllm-project:ascendfrom
wangx700:sparse-hccl-ascend
Draft

wangx700 wants to merge 6 commits into
vllm-project:ascendfrom
wangx700:sparse-hccl-ascend

Conversation

@wangx700

@wangx700 wangx700 commented Sep 22, 2026 •

Copy link
Copy Markdown

Summary

Port of the sparse HCCL weight-sync feature (wangx700/vime#1) onto the ascend branch. After an initial dense synchronization, subsequent updates transfer only the changed weight elements between Megatron training workers and vLLM-Ascend rollout workers.

Initial update: Megatron -> dense HF weight seed -> rollout
Steady updates: Megatron local shard delta
             -> final HF global sparse coordinates
             -> HCCL positions + values transfer
             -> rollout runtime/TP-shard mapping and in-place update

What this PR adds

  • New delta-sync transport sparse_hccl for --update-weight-mode=delta (alongside disk): 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.
  • New modules under vime/backends/megatron_utils/update_weight/: delta_spec.py, delta_sync/ (encode, sparse_gather), megatron_delta_export.py, update_weight_from_sparse_hccl.py.
  • CLI: --update-weight-transport gains the sparse_hccl choice; new --update-weight-delta-batch-gather / --update-weight-delta-verify-every options and delta-mode validation.
  • Docs (EN/ZH) and unit tests (tests/test_sparse_hccl_delta_sync.py, tests/test_update_weight_factory.py).

Adaptations for the ascend branch

The ascend branch does not yet carry the delta weight-sync foundation that exists on main, so a few minimal anchors were added here to make the feature self-contained (kept as close to upstream main as possible):

  • update_weight/__init__.py and update_weight/common.py: bring in create_weight_updater and VimeRayWeightSyncClient (the latter identical to the version on main; the former is main's factory plus the new sparse_hccl branch).
  • backends/megatron_utils/actor.py: route delta + sparse_hccl to UpdateWeightFromSparseHCCL in this branch's inline updater selection.
  • backends/vllm_utils/vllm_engine.py: select the sparse_hccl weight-transfer backend in --weight-transfer-config when 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, named HfWeightIteratorSparseBridge because this branch already ships a different HfWeightIteratorBridge for the HfWeightIteratorBase registry (both classes coexist).
  • backends/megatron_utils/actor.py: keep weights live across wake_up()/sleep() when delta + sparse_hccl is configured (avoids a reload around each weight update).
  • The VLLM_USE_V2_MODEL_RUNNER env-default tweak from the source branch has no anchor on this branch and was skipped.

Dependencies / notes

  • Requires vllm_ascend.distributed.weight_transfer.sparse_hccl_engine and vllm_ascend.distributed.weight_transfer.sparse_weight_patch on the vLLM-Ascend side (separate PR).
  • All feature files (exporter, updater internals, delta-sync modules, scripts, tests, docs) are identical to the source branch's latest state; the remaining files are the minimal ascend-branch integration points listed above.
  • The machine-specific example script run-qwen3-4B-delta-sparse-hccl.sh from the source branch is intentionally omitted.
  • Synced with the source branch through ix time consume for sparse hccl. Opened as a draft for initial feedback on the integration points.

@gemini-code-assist gemini-code-assist Bot left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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())

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

high

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.

Suggested change
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()

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

high

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()

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

high

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.

Suggested change
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):

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

high

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.

Suggested change
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):

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

high

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):

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

high

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.

Suggested change
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(),

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

medium

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.

Suggested change
ray._private.services.get_node_ip_address(),
ray.util.get_node_ip_address(),

Comment on lines +259 to +261
def disconnect_rollout_engines(self) -> None:
self._group = None
self._client = None

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

medium

Clean up the persistent thread pool executor when disconnecting rollout engines to prevent resource leaks.

Suggested change
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

Comment on lines +326 to +342
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)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

medium

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.

Suggested change
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

Signed-off-by: wangx700 <wangxin700@huawei.com>
Signed-off-by: wangx700 <wangxin700@huawei.com>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant