Repository navigation
Conversation
In colocate mode the first rollout after weight update produced garbled output when training GLM-4.7-Flash (MLA + MoE) on Ascend NPU. Root cause: start_weight_update -> layerwise IPC reload (vllm model_loader/reload/meta.py capture_layer_to_meta / restore_layer_on_meta) moves the rotary cos_sin_cache buffer (persistent=False) to meta device. The layerwise reload only refills *parameters*, and finish_weight_update -> finalize_layerwise_reload / materialize_layer re-creates the buffer storage *empty*. The MLA rope path reads the vllm-ascend module globals _cos_cache/_sin_cache (get_cos_sin_mla -> _cos_cache[positions]), which end up zeroed, so the first rollout's npu_interleave_rope applies all-zero cos/sin, zeroes q, and attention collapses into garbled output. Fix: hook finish_weight_update so that after the original finalize we recompute cos_sin_cache deterministically (RotaryEmbedding ._compute_cos_sin_cache from base / rotary_dim / max_position_embeddings -- no reads of the clobbered buffer/inv_freq), write it back into the buffer, and force-re-record the vllm-ascend globals (resetting _cos_cache/_sin_cache/_cos_sin_cache to None first to bypass the record-once guard) so the MLA rope path sees real cos/sin on the first post-update rollout. Gated by VIME_RESTORE_ROPE (default 1, set 0 for A/B comparison). Verified end-to-end on Ascend 910B colocate training of GLM-4.7-Flash: every TP worker logs "[VIME-ROPE-RESTORE] recomputed cos_sin_cache for 1 rotary module(s)", and the first rollout after weight update is correct (math sample answered correctly, reward=1, no garbling). vllm-project#403 had to drop --colocate from the GLM-4.7 NPU e2e test due to colocate weight-sync breakage; this restores correct first-rollout behavior in colocate mode. Signed-off-by: Bobfan <fanbinbin@zju.edu.cn> Co-Authored-By: Claude Code <noreply@anthropic.com>
There was a problem hiding this comment.
Code Review
This pull request addresses an issue where the first rollout after a colocate weight-update produces garbled output. It introduces a _restore_rope_cache method to recompute and re-record the RoPE cos/sin cache, and patches finish_weight_update to trigger this restoration. The review feedback recommends enhancing robustness and efficiency by adding defensive checks for the worker and model runner, resetting global caches before the loop to prevent redundant operations, verifying the existence of cos_sin_cache, and using torch.no_grad() for PyTorch best practices.
| model = worker.model_runner.model | ||
| inner = getattr(model, "model", None) or model | ||
| count = 0 | ||
| for mod in inner.modules(): | ||
| if not isinstance(mod, RotaryEmbedding): | ||
| continue | ||
| try: | ||
| cache = mod._compute_cos_sin_cache() | ||
| buf = mod.cos_sin_cache | ||
| cache = cache.to(device=buf.device, dtype=buf.dtype) | ||
| buf.data.copy_(cache) | ||
| # MLA reads the globals _cos_cache/_sin_cache (NOT the buffer); | ||
| # _record_cos_and_sin_cache_interleaved is a "record once" no-op once | ||
| # they are set, so reset to None first to force re-record. | ||
| asc_rope._cos_cache = None | ||
| asc_rope._sin_cache = None | ||
| asc_rope._cos_sin_cache = None | ||
| if hasattr(asc_rope, "_record_cos_sin_cache"): | ||
| asc_rope._record_cos_sin_cache(buf) | ||
| if hasattr(asc_rope, "_record_cos_and_sin_cache_interleaved"): | ||
| asc_rope._record_cos_and_sin_cache_interleaved(buf) | ||
| count += 1 | ||
| except Exception as e: # noqa: BLE001 | ||
| _log.warning("[VIME-ROPE-RESTORE] failed for a rotary module", exc_info=True) |
There was a problem hiding this comment.
To improve robustness, correctness, and efficiency, we should:
- Add defensive checks to ensure
worker,worker.model_runner, andworker.model_runner.modelare notNonebefore accessing them. - Reset the
asc_ropeglobals (_cos_cache,_sin_cache,_cos_sin_cache) toNonebefore the loop instead of inside it. This avoids redundant resets and ensures we record the first encountered rotary module's cache (which aligns with the 'record once' guard behavior), rather than repeatedly overwriting them with the last module's cache. - Safely check if
cos_sin_cacheexists and is notNoneon each rotary module. - Use
with torch.no_grad(): buf.copy_(cache)instead ofbuf.data.copy_(cache)to adhere to modern PyTorch best practices.
model_runner = getattr(worker, "model_runner", None)
if model_runner is None:
return
model = getattr(model_runner, "model", None)
if model is None:
return
inner = getattr(model, "model", None) or model
# MLA reads the globals _cos_cache/_sin_cache (NOT the buffer);
# _record_cos_and_sin_cache_interleaved is a "record once" no-op once
# they are set, so reset to None first to force re-record.
# Resetting before the loop ensures we record the first encountered
# rotary module's cache and avoid redundant resets/overwrites.
asc_rope._cos_cache = None
asc_rope._sin_cache = None
asc_rope._cos_sin_cache = None
count = 0
for mod in inner.modules():
if not isinstance(mod, RotaryEmbedding):
continue
try:
buf = getattr(mod, "cos_sin_cache", None)
if buf is None:
continue
cache = mod._compute_cos_sin_cache()
cache = cache.to(device=buf.device, dtype=buf.dtype)
with torch.no_grad():
buf.copy_(cache)
if hasattr(asc_rope, "_record_cos_sin_cache"):
asc_rope._record_cos_sin_cache(buf)
if hasattr(asc_rope, "_record_cos_and_sin_cache_interleaved"):
asc_rope._record_cos_and_sin_cache_interleaved(buf)
count += 1
except Exception as e: # noqa: BLE001
_log.warning("[VIME-ROPE-RESTORE] failed for a rotary module", exc_info=True)
Documentation build overview
48 files changed ·
|
Address the review feedback on the rope restore path: - guard the worker.model_runner.model access: log a warning and skip instead of letting an exception escape into finish_weight_update - reset the vllm-ascend rope globals once before the module loop: the record-once guard then keeps the first module's cache (upstream first-writer semantics) and avoids per-module redundant resets - skip rotary modules without a cos_sin_cache buffer - use torch.no_grad() around the buffer copy instead of .data.copy_() Behavior is unchanged for the verified GLM-4.7-Flash colocate setup (exactly one rotary module). Signed-off-by: Bobfan <fanbinbin@zju.edu.cn> Co-Authored-By: Claude Code <noreply@anthropic.com>
- drop the unused `as e` from the except clause (ruff) - add a blank line after the inline `import logging` (black) Only the two lint hooks failed in vime-npu-ci vllm-project#288 / vime-ci #1271; all other files passed. Signed-off-by: Bobfan <fanbinbin@zju.edu.cn> Co-Authored-By: Claude Code <noreply@anthropic.com>
Problem
In colocate mode, the first rollout after weight update produced garbled output when training GLM-4.7-Flash (MLA + MoE) on Ascend NPU. #403 had to drop
--colocatefrom the GLM-4.7 NPU e2e test due to colocate weight-sync breakage (MTP loss 3.12, 0% speculative acceptance).Root cause
start_weight_update-> layerwise IPC reload (vllmmodel_loader/reload/meta.pycapture_layer_to_meta/restore_layer_on_meta) moves the rotarycos_sin_cachebuffer (persistent=False) to meta device. The layerwise reload only refills parameters, andfinish_weight_update->finalize_layerwise_reload/materialize_layerre-creates the buffer storage empty. The MLA rope path reads the vllm-ascend module globals_cos_cache/_sin_cache(get_cos_sin_mla->_cos_cache[positions]), which end up zeroed -> the first rollout'snpu_interleave_ropeapplies all-zero cos/sin, zeroes q, and attention collapses into garbled output.Fix
Hook
finish_weight_update(the weight-update endpoint, afterfinalize_layerwise_reload):cos_sin_cachedeterministically (RotaryEmbedding._compute_cos_sin_cache()frombase/rotary_dim/max_position_embeddings— no reads of the clobbered buffer/inv_freq) and write it back into the buffer._cos_cache/_sin_cache/_cos_sin_cache(reset toNonefirst to bypass the record-once guard in_record_cos_and_sin_cache_interleaved) so the MLA rope path sees real cos/sin on the first post-update rollout.Gated by
VIME_RESTORE_ROPE(default1; set0to disable, e.g. for A/B comparison).Evidence: first rollout after weight update, before vs after
Identical colocate setup (same script,
GLM-4.7-Flashon Ascend 910B, TP workers colocated with training); the only difference is this fix.Without the fix — output is incoherent token soup, reward 0:
With the fix — every TP worker restores the cache, output is coherent and correct, reward 1:
(The two runs sample different problems from the dataset, as expected; the contrast is reward 0 with collapsed output vs reward 1 with correct step-by-step reasoning, on the first rollout after weight update.)
Verification
POST /finish_weight_updatereturns200 OK