Skip to content

fix(npu): restore RoPE cos/sin cache after colocate weight update - #446

Open
Bobfan wants to merge 3 commits into
vllm-project:ascendfrom
Bobfan:fix/colocate-rope-cache
Open

Bobfan wants to merge 3 commits into
vllm-project:ascendfrom
Bobfan:fix/colocate-rope-cache

Conversation

@Bobfan

@Bobfan Bobfan commented Sep 19, 2026

Copy link
Copy Markdown

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 --colocate from 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 (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 -> 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 (the weight-update endpoint, after finalize_layerwise_reload):

  1. 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) and write it back into the buffer.
  2. Force-re-record the vllm-ascend globals _cos_cache / _sin_cache / _cos_sin_cache (reset to None first 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 (default 1; set 0 to disable, e.g. for A/B comparison).

Evidence: first rollout after weight update, before vs after

Identical colocate setup (same script, GLM-4.7-Flash on Ascend 910B, TP workers colocated with training); the only difference is this fix.

Without the fix — output is incoherent token soup, reward 0:

prompt: "You have 2020 piles of coins in front of you. The first pile contains 1 coin, ...
         Find the minimum number of turns you need to take away all of these coins."
response (excerpt):
  <think>".\n\n").\n\n".";\n\n".\n\n`.\n\n".\n\n".\n\n". remotely".\n\n".\n\n".\n\n".\n\n
  0".\n\n0".").\n\n".\n\n0".0".\n\n0".\n\n0".\n\n0".0".\n\n0".\n\n0".\n\n0".0".0".
  ... S".S".S".T".T".\n\n"I + "S".S".S".S<|endoftext|>
label: 11, reward: 0

With the fix — every TP worker restores the cache, output is coherent and correct, reward 1:

INFO vime.backends.megatron_utils.update_weight.update_weight_from_tensor:
  [VIME-ROPE-RESTORE] recomputed cos_sin_cache for 1 rotary module(s)   (all TP workers)

prompt: "The number 24 can be made by multiplying together four prime numbers: 2, 2, 2 and 3.
         How many primes must be multiplied to make 2400?"
response (excerpt):
  <think>We are given: ... Compute prime factorization of 2400.
  2400 = 24 * 100 = (2^3 * 3) * (2^2 * 5^2) = 2^5 * 3 * 5^2.
  Thus 2400 consists of five 2's, one 3, and two 5's.
  The total number of prime factors is 5 + 1 + 2 = 8.
  Answer: \boxed{8}</think> ... Answer: \boxed{8}
label: 8, 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

  • End-to-end on Ascend 910B, colocate training of GLM-4.7-Flash
  • POST /finish_weight_update returns 200 OK
  • Training proceeds past the first rollout into backward

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>

@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 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.

Comment on lines +358 to +381
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)

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

To improve robustness, correctness, and efficiency, we should:

  1. Add defensive checks to ensure worker, worker.model_runner, and worker.model_runner.model are not None before accessing them.
  2. Reset the asc_rope globals (_cos_cache, _sin_cache, _cos_sin_cache) to None before 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.
  3. Safely check if cos_sin_cache exists and is not None on each rotary module.
  4. Use with torch.no_grad(): buf.copy_(cache) instead of buf.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)

Bobfan and others added 2 commits September 20, 2026 09:10
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>
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