Skip to content

[Perf][PersonaPlex] Hoist RoPE tables and RingKV position computation out of the per-layer loop - #7850

Merged
linyueqian merged 1 commit into
vllm-project:mainfrom
Sendoh-code:main
Oct 8, 2026
Merged

linyueqian merged 1 commit into
vllm-project:mainfrom
Sendoh-code:main

Conversation

@Sendoh-code

@Sendoh-code Sendoh-code commented Sep 19, 2026 •

Copy link
Copy Markdown
Contributor

[Perf][PersonaPlex] Hoist RoPE tables and RingKV position computation out of the per-layer loop

Summary

As noted in #7389, PersonaPlexTemporalStreaming.step() (32 Helium layers) and
_MimiStreamingTransformer.step() (8 Mimi layers) rebuilt two tables in every
layer on every frame, even though the tables depend only on the stack's offset and
the ring capacity:

  • the RoPE cos/sin rotation tables (inside _apply_rope)
  • the RingKV write indexes and pre-mask position table (inside _RingKV.complete)

All layers of a stack advance the same per-row [B] offset together, so both
tables are the same for every layer within one step(). This PR builds them once
per step() and passes them to every layer.

This revision is based on main @ 3bc3f1a7d (after #7670 and #8192) and keeps
the per-row offset / active row-mask API.

Changes

personaplex_temporal.py

  • _rope_tables(offset, seq_len, head_dim, max_period=10_000.0) takes the
    per-row [B] offset and returns (rotr, roti), each of shape [B, 1, T, D/2].
  • _apply_rope(q, k, rotr, roti) now only applies the rotation, using the
    precomputed tables.
  • _ringkv_positions(offset, seq_len, capacity, active) returns the write
    indexes ([B, T]) and the pre-mask positions ([B, capacity]), computed
    from each row's post-write end offset offset + T * active. The moshi
    delta <= 0 convention comment moved here together with that line.
  • _RingKV.complete(k, v, active, indexes=None, positions=None) still guards
    the active rows (gather → where → scatter_) and still applies each row's
    start_offset mask (elastic recycle) on every call. The start_offset mask is
    per-ring state, so it is never hoisted. indexes and positions are optional.
    When they are omitted, complete() builds them from the ring's own
    end_offset, so the standalone complete(k, v, active) calls (e.g. in
    test_ring_per_row_offsets.py) work unchanged.
  • _TemporalLayer.forward(x, kv, offset, context, active, rope, ring).
  • PersonaPlexTemporalStreaming.step(frame_embedding, active) keeps the same
    signature as main. It builds rope (passing self.max_period) and ring
    once per step.

personaplex_mimi.py

  • _MimiStreamingTransformer.step(x, active) keeps the same signature as
    main. Each instance builds rope and ring once per step from its own [B]
    _offset, and passes them to _MimiTransformerLayer.forward(..., active, rope, ring).
    The encoder and decoder instances share no state.
  • The active handling in the codec entry points is the same as on main.

There is no global or class-level cache. An earlier version of this PR cached
the tables in a module-level dict. That cache thrashed: the Mimi encoder and
decoder have the same head_dim and capacity but different offsets, so they
kept evicting each other's entry. A cache-hit branch would also add overhead and
a graph break. Keeping the tables as locals of step() avoids all of this.

vllm_omni/deploy/personaplex.yaml is unchanged from main.

Performance

Setup: NVIDIA A100-SXM4-40GB, torch 2.13.0+cu132, bf16, B=4, all rows
active, random weights. Each run does 40 warm-up steps, then measures 30 step()
calls. Kernel count is the number of CUDA kernel events from torch.profiler
divided by 30. Wall time is eager and measured without the profiler. Both
revisions ran the same script; only the source tree on PYTHONPATH changed.

Stack (config) kernels/step (base → PR) eager ms/step (base → PR)
Temporal: 32 layers, dim=4096, context=3000, T=1 3149 → 2408 (-23.5%) 62.6 → 49.7
Mimi: 8 layers, dim=512, context=250, T=2 707 → 542 (-23.3%) 15.0 → 11.2
Mimi: same, T=10 (codec_chunk_frames: 5) 691 → 526 (-23.9%) 15.2 → 11.1

Correctness

This is a pure refactor and should not change any numbers. The new test file
tests/model_executor/models/personaplex/test_temporal_streaming_hoist.py
contains a frozen oracle: a verbatim copy of the per-row _apply_rope,
_RingKV.complete, layer forward and stack step from main @ 3bc3f1a7d.
Every run compares the live code against this oracle with torch.equal. There
are 14 test cases, all on CPU:

  • RoPE (test_apply_rope_matches_legacy_per_row, 5 cases). Covers per-row
    offsets such as (1, 250, 3000), with T ∈ {1, 2, 4, 10} and
    head_dim ∈ {64, 128}. It also checks that the tables have shape [B, 1, T, D/2].
  • Table reuse (test_rope_tables_reused_across_layers_within_step, 1
    case). Applies one table pair to several q/k pairs, the way the layers of one
    step use it.
  • RingKV (test_ringkv_complete_matches_legacy_across_wrap, 4 cases).
    Parametrized over (T, capacity) ∈ {(1, 5), (2, 7), (5, 7), (5, 12)}, with
    2,000 steps per case. Each step uses a random active mask and random
    reset_slot / bump_slot_start / reset_row calls. Every step compares three
    rings: the oracle, a ring given the hoisted tables, and a standalone ring
    given no tables. The test also asserts that some write actually wraps around
    the ring end in the middle of a chunk.
  • Temporal end-to-end (test_temporal_streaming_step_matches_legacy_end_to_end,
    1 case). Runs 2,000 steps with context=16 and random active masks, then
    compares the outputs, the stack _offset, and the end_offset /
    start_offset of every ring.
  • Mimi end-to-end (test_mimi_streaming_step_matches_legacy_end_to_end, 3
    cases). Covers T ∈ {1, 2, 10} with 1,000 steps each, including recycling
    rows after the ring has wrapped.

A mutation check confirms that the tests protect the moshi convention: changing
delta <= 0 to delta < 0 in _ringkv_positions makes 8 tests fail.

Test plan

All runs below use 706e91f4f.

# New suite for this change
.venv/bin/pytest tests/model_executor/models/personaplex/test_temporal_streaming_hoist.py -v
# 14 passed

# PersonaPlex model tests (includes #7670's test_ring_per_row_offsets.py, whose CUDA cases ran on an A100),
# stage input processors, and the deploy-config factory
.venv/bin/pytest tests/model_executor/models/personaplex/ \
  tests/model_executor/stage_input_processors/test_personaplex.py \
  tests/config/test_config_factory.py
# 555 passed

pre-commit run --files \
  vllm_omni/model_executor/models/personaplex/personaplex_temporal.py \
  vllm_omni/model_executor/models/personaplex/personaplex_mimi.py \
  tests/model_executor/models/personaplex/test_temporal_streaming_hoist.py
# all hooks passed
  • test_temporal_streaming_hoist.py: 14/14
  • Full PersonaPlex model tests + stage input processor tests + config factory tests: 555/555
  • Mutation check (<= → <): 8 failures, as expected
  • pre-commit passes on all changed files

@vllm-omni-review-bot

Copy link
Copy Markdown

This PR appears to belong to: docs/design/module/model_integration.md, docs/design/module/ar_runtime.md.

Module owners: @tzhouam @fake0fan @Gaohan123

Routing: @tzhouam via module of the changed files, CODEOWNERS; @fake0fan via module of the changed files; @Gaohan123 via module of the changed files

@Sendoh-code, please review your own changes and leave a short self-review comment describing what you checked. PRs without author self-review may not be assigned a reviewer.

Please take a look when you have a chance. If you would like an automated review, mention @vllm-omni-review-bot in a comment.

@Sendoh-code

Sendoh-code commented Sep 20, 2026 •

Copy link
Copy Markdown
Contributor Author

Self-review: The change is minior, the only change is to add a one-time calculation for rope tables and positions in the ringKV class at the start of each step and reuse for all the layers. Plus this the signature of forward function is changed to pass those cache. Have added relevant tests and also run original tests about personaplex and all of them passed. Also the number of lauched kernel before and after the change can suggest a speed up. This PR is ready for review.

@hsliuustc0106 hsliuustc0106 added the enhancement New feature or request label Sep 20, 2026
@Sendoh-code

Copy link
Copy Markdown
Contributor Author

@linyueqian Could you take a look at this PR and check whether this is what you want from issue #7389? Thanks!

@linyueqian linyueqian left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

Thanks for taking item 5 of #7389. The hoist has the right shape: one RoPE table and one ring position table per step(), threaded through forward() as locals with no shared cache, which is what the RFC asked for after the module-level dict thrashed between the Mimi encoder and decoder.

I could only do a static read at 9bf1377c: the branch is 183 commits behind main (merge-base 65055774), no ready lane has run, and a trial merge with main conflicts in all three touched production paths. Two of the conflicts are the substance of this round and are inline, together with the deploy YAML deletion.

DCO is action_required on this PR because none of the four commits carry the sign-off trailer; git commit -s on each commit during the rebase fixes it.

Validation: git diff 65055774..9bf1377c read in full; git merge-tree origin/main 9bf1377c for the conflict set; main signatures checked at 825ab20b (#7670). The kernel-launch numbers in the description are plausible for a 32-layer stack but unverified here; we will run a base-vs-head step() A/B on H100-class hardware once the branch is rebased on #7670, since the per-row offsets change what the tables have to look like.

# 8 active codebooks (cb 0..7) -> 1920 PCM samples @ 24 kHz / 12.5 Hz.
# * Stage 0 runs cudagraph by default; Stage 1 (Mimi) is eager because the
# external decoder performs host-to-device copies that cannot be captured.
pipeline: personaplex

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

[blocking] The last commit (9bf1377c, "Delete change in vllm_omni/deploy/personaplex.yaml") removes the whole file instead of reverting your edits to it. PersonaPlexPipeline resolves default_deploy_config_name="personaplex.yaml" (vllm_omni/model_executor/models/personaplex/pipeline.py:33), so after this PR a plain vllm serve nvidia/personaplex-7b-v1 --omni has no deploy config, and the README, docs/serving/*.md and the duplex unit tests that read this path all break. main also changed the file since your branch point (#7454, #7522), which is the modify/delete conflict GitHub reports. Please restore it from main and drop this commit from the branch. It also fails CI deterministically: tests/config/test_config_factory.py::test_registered_default_deploy_config_exists asserts that every registered pipeline's default deploy config is a file.

B, H, T, D = q.shape
ds = torch.arange(D // 2, device=q.device, dtype=torch.float32)

def _rope_tables(D, x, offset: torch.Tensor, max_period: float = 10_000.0):

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

[blocking] main moved under this since you branched: #7670 (825ab20b) made _offset a per-row [B] tensor, builds the RoPE phase as ts.view(-1, 1, T, 1) from offset.float().view(-1, 1), and gave _RingKV.complete() an active row mask. This helper assumes one offset for the whole batch (ts.view(1, -1, 1) at line 53) and _ringkv_positions returns a single [1, capacity] table, so rebasing as-is would drop the per-row offsets that multi-session batches now rely on, and the trial merge conflicts here and in personaplex_mimi.py. The hoist still works per row: compute rotr/roti once per step as [B, 1, T, D/2] and the pre-mask positions as [B, capacity] from the per-row end_offset, keep active flowing into complete(), and thread both tables down exactly as you do now. The frozen oracle in the new test then needs to copy main's per-row functions rather than the pre-#7670 ones, otherwise it pins the wrong reference. The concrete casualty is tests/model_executor/models/personaplex/test_ring_per_row_offsets.py from #7670, which calls the three-argument complete(k, v, active) and the two-argument step(x, active) on both stacks and stops collecting under these signatures. One rebase trap to plan for: with [B] offsets indexes becomes [B, T], and index_copy_ in complete only takes a 1-D index, which is why main writes the ring with scatter_.


def step(self, x: torch.Tensor) -> torch.Tensor:
"""``x`` is ``[B, T, dim]`` (T = positions this frame, typically 2)."""
rotr, roti = _rope_tables(self.layers[0].head_dim, x, self._offset)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

[blocking] Same root cause as the temporal stack, anchored where the hoist lands: on main this transformer's _offset is torch.zeros(batch_size, ...) and step() takes an active row mask (#7670), so an inactive row's offset does not advance; this branch still has the pre-#7670 [1] offset and a step(x) with no mask, and the single table built here is shared by every row. The codec entrypoints this PR leaves untouched (encode_frame, decode_frame, decode_frames, _run_stages, the streaming conv wrappers) all thread active on main, so the rebase has to re-add it through them as well as through this hoist. The Mimi encoder and decoder instances each keep their own [B] offsets, so the per-step table cost stays one build per instance.

kr, ki = k[..., 0].float(), k[..., 1].float()
rotr = torch.cos(freqs * ts)
roti = torch.sin(freqs * ts)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

[important] Formatting will fail pre-commit: this whitespace-only line, the trailing spaces after self.end_offset.add_(k.shape[2]) in complete(), the unspaced annotations in _ringkv_positions (offset:torch.Tensor, dtype = torch.long) and the single blank line before class _RingKV are all things ruff-format rewrites. pre-commit run --all-files from the repo root and committing the result clears it.

return qo.view(*dims, D), ko.view(*dims, D)


def _ringkv_positions(offset:torch.Tensor, T:int, capacity:int, device):

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

[suggestion] The moshi bit-parity warning that lived in complete() ("delta <= 0, not < 0, is moshi's exact convention ... do NOT change this") was dropped when the position math moved here; carry it over, since this is now the line a future editor would be tempted to change. While here, _rope_tables gained a max_period parameter but step() never passes self.max_period, so the constructor argument stays dead; wire it or drop it.

offset = torch.zeros(1, dtype=torch.long)
num_steps = 4000 # capacity=5 -> hundreds of wraps
for step in range(num_steps):
T = 1

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

[suggestion] test_ringkv_complete_matches_legacy_across_wrap hardcodes T = 1, but the deploy config streams codec_chunk_frames: 5, so frames with T > 1 are the production case for the Mimi stack; add a T > 1 case that crosses a chunk boundary so the wrap arithmetic is covered where it matters.

@Sendoh-code

Copy link
Copy Markdown
Contributor Author

Thank you for the review! I'll fix them and solve the conflict soon.

@vllm-omni-review-bot

Copy link
Copy Markdown

Omni ReviewBot: no human activity for 7 days

@Sendoh-code this pull request has had no human commit, comment or review since 2026-09-23. Please confirm the current plan and next step. The author or a maintainer decides whether to change the PR state.

To keep it moving, any one of these is enough: push an update, reply to the open blocker, or post the current plan and timeline.

@Sendoh-code

Copy link
Copy Markdown
Contributor Author

Omni ReviewBot: no human activity for 7 days

@Sendoh-code this pull request has had no human commit, comment or review since 2026-09-23. Please confirm the current plan and next step. The author or a maintainer decides whether to change the PR state.

To keep it moving, any one of these is enough: push an update, reply to the open blocker, or post the current plan and timeline.

Thanks for the reminder! I am still working on this PR. The next step is to fix the requested changes, and I plan to push an update in a week or so.

… out of the per-layer loop

Build the per-row RoPE cos/sin tables ([B, 1, T, D/2]) and the RingKV write
indexes ([B, T]) / pre-mask position table ([B, capacity]) once per step() in
PersonaPlexTemporalStreaming and _MimiStreamingTransformer, and thread them
through every layer instead of rebuilding them per layer. The per-row
start_offset mask and the active-row write guard stay inside
_RingKV.complete(); complete() still builds its own tables when called
standalone.

Rebased onto the per-row offsets / active-mask API from vllm-project#7670.

Signed-off-by: Yifan Tian <Yifan.Tian@colorado.edu>
@Sendoh-code

Sendoh-code commented Oct 5, 2026 •

Copy link
Copy Markdown
Contributor Author

@linyueqian Thanks for the detailed review! I rebuilt the change on top of current main (3bc3f1a7d, after #7670/#8192) as a single signed-off commit instead of replaying the old ones:

  1. Deploy YAML: the branch now starts from main. Both the deletion commit and the stray gpu_memory_utilization edit are gone, and personaplex.yaml is byte-identical to main.
  2. Per-row offsets (temporal): _rope_tables(offset[B], T, head_dim, max_period) returns [B, 1, T, D/2] tables, and _ringkv_positions(offset[B], T, capacity, active) returns indexes [B, T] and pre-mask positions [B, capacity] from the per-row post-write end offset. complete() keeps active, the gather/where/scatter_ guard and the fresh per-row start_offset mask. The tables are optional args, so standalone complete(k, v, active) (and test_ring_per_row_offsets.py) is unchanged. step(x, active) keeps main's signature.
  3. Mimi: same hoist in _MimiStreamingTransformer.step(x, active), once per instance per step. The codec entrypoints keep main's active threading untouched.
  4. Formatting: pre-commit run passes on all touched files.
  5. The moshi delta <= 0 warning moved with the line into _ringkv_positions. step() now passes self.max_period.
  6. The RingKV test is parametrized over (T, capacity) ∈ {(1,5), (2,7), (5,7), (5,12)}, and the test asserts that a chunk actually straddles the ring end. Mimi e2e is parametrized over T ∈ {1, 2, 10}. All runs include random active masks plus reset_slot/bump_slot_start/reset_row.
  7. The frozen oracle is now a verbatim copy of main's per-row functions (pinned to 3bc3f1a7d).
  8. DCO: the commit is signed off.

Local: 555 passed across tests/model_executor/models/personaplex/, stage_input_processors/test_personaplex.py and tests/config/test_config_factory.py (CUDA cases on A100). Refreshed kernel counts per step() on A100, bf16, B=4: temporal 3149 → 2408 (-23.5%), Mimi 707 → 542 (-23.3%).

@linyueqian linyueqian added the ready label to trigger buildkite CI label Oct 5, 2026

@linyueqian linyueqian left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

Reviewed rebuilt 706e91f. The deploy configuration is restored and the hoist now preserves per-row offsets, active masks and per-ring reset boundaries. The updated wrap tests address previous findings and static review surfaced no material issues.

[suggestion] Please refresh the PR description for this revision. The helper signatures, 54-test transcript, 9/9 hoist-suite count and 4,000 RingKV iterations describe older code, whereas the current tests define 14 parametrized cases and run 2,000 iterations per RingKV case. Identify the revisions and configurations behind the performance numbers and label or replace the historical results.

This was a static review only. I did not run the PR code, tests or benchmarks; the reported 555 passes and A100 measurements remain author-reported evidence.

@Sendoh-code

Copy link
Copy Markdown
Contributor Author

The PR description is updated, this should be clean to merge now.

@vllm-omni-review-bot

Copy link
Copy Markdown
Omni ReviewBot routing record

Assigned Strict on cursor (cursor-grok-4.6-high) under experiment fleet-strict-cursor-grok46-zcode-glm53flash-5050-c5-z10-20261002.

@linyueqian
linyueqian merged commit 6667316 into vllm-project:main Oct 8, 2026
7 of 9 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

enhancement New feature or request ready label to trigger buildkite CI

Projects

None yet

Development

Successfully merging this pull request may close these issues.

4 participants