Repository navigation
[Perf][PersonaPlex] Hoist RoPE tables and RingKV position computation out of the per-layer loop - #7850
Conversation
|
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. |
|
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. |
|
@linyueqian Could you take a look at this PR and check whether this is what you want from issue #7389? Thanks! |
linyueqian
left a comment
There was a problem hiding this comment.
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 |
There was a problem hiding this comment.
[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): |
There was a problem hiding this comment.
[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) |
There was a problem hiding this comment.
[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) | ||
|
|
There was a problem hiding this comment.
[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): |
There was a problem hiding this comment.
[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 |
There was a problem hiding this comment.
[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.
|
Thank you for the review! I'll fix them and solve the conflict soon. |
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>
|
@linyueqian Thanks for the detailed review! I rebuilt the change on top of current
Local: 555 passed across |
linyueqian
left a comment
There was a problem hiding this comment.
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.
|
The PR description is updated, this should be clean to merge now. |
Omni ReviewBot routing recordAssigned Strict on cursor (cursor-grok-4.6-high) under experiment |
[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 everylayer on every frame, even though the tables depend only on the stack's offset and
the ring capacity:
cos/sinrotation tables (inside_apply_rope)_RingKV.complete)All layers of a stack advance the same per-row
[B]offset together, so bothtables are the same for every layer within one
step(). This PR builds them onceper
step()and passes them to every layer.This revision is based on
main@3bc3f1a7d(after #7670 and #8192) and keepsthe per-row offset /
activerow-mask API.Changes
personaplex_temporal.py_rope_tables(offset, seq_len, head_dim, max_period=10_000.0)takes theper-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 theprecomputed tables.
_ringkv_positions(offset, seq_len, capacity, active)returns the writeindexes([B, T]) and the pre-maskpositions([B, capacity]), computedfrom each row's post-write end offset
offset + T * active. The moshidelta <= 0convention comment moved here together with that line._RingKV.complete(k, v, active, indexes=None, positions=None)still guardsthe
activerows (gather → where →scatter_) and still applies each row'sstart_offsetmask (elastic recycle) on every call. Thestart_offsetmask isper-ring state, so it is never hoisted.
indexesandpositionsare optional.When they are omitted,
complete()builds them from the ring's ownend_offset, so the standalonecomplete(k, v, active)calls (e.g. intest_ring_per_row_offsets.py) work unchanged._TemporalLayer.forward(x, kv, offset, context, active, rope, ring).PersonaPlexTemporalStreaming.step(frame_embedding, active)keeps the samesignature as
main. It buildsrope(passingself.max_period) andringonce per step.
personaplex_mimi.py_MimiStreamingTransformer.step(x, active)keeps the same signature asmain. Each instance buildsropeandringonce per step from its own[B]_offset, and passes them to_MimiTransformerLayer.forward(..., active, rope, ring).The encoder and decoder instances share no state.
activehandling in the codec entry points is the same as onmain.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_dimandcapacitybut different offsets, so theykept 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.yamlis unchanged frommain.Performance
Setup: NVIDIA A100-SXM4-40GB, torch 2.13.0+cu132, bf16,
B=4, all rowsactive, 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.profilerdivided by 30. Wall time is eager and measured without the profiler. Both
revisions ran the same script; only the source tree on
PYTHONPATHchanged.dim=4096,context=3000,T=1dim=512,context=250,T=2T=10(codec_chunk_frames: 5)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.pycontains a frozen oracle: a verbatim copy of the per-row
_apply_rope,_RingKV.complete, layerforwardand stackstepfrommain@3bc3f1a7d.Every run compares the live code against this oracle with
torch.equal. Thereare 14 test cases, all on CPU:
test_apply_rope_matches_legacy_per_row, 5 cases). Covers per-rowoffsets such as
(1, 250, 3000), withT ∈ {1, 2, 4, 10}andhead_dim ∈ {64, 128}. It also checks that the tables have shape[B, 1, T, D/2].test_rope_tables_reused_across_layers_within_step, 1case). Applies one table pair to several q/k pairs, the way the layers of one
step use it.
test_ringkv_complete_matches_legacy_across_wrap, 4 cases).Parametrized over
(T, capacity) ∈ {(1, 5), (2, 7), (5, 7), (5, 12)}, with2,000 steps per case. Each step uses a random
activemask and randomreset_slot/bump_slot_start/reset_rowcalls. Every step compares threerings: 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.
test_temporal_streaming_step_matches_legacy_end_to_end,1 case). Runs 2,000 steps with
context=16and randomactivemasks, thencompares the outputs, the stack
_offset, and theend_offset/start_offsetof every ring.test_mimi_streaming_step_matches_legacy_end_to_end, 3cases). Covers
T ∈ {1, 2, 10}with 1,000 steps each, including recyclingrows after the ring has wrapped.
A mutation check confirms that the tests protect the moshi convention: changing
delta <= 0todelta < 0in_ringkv_positionsmakes 8 tests fail.Test plan
All runs below use
706e91f4f.test_temporal_streaming_hoist.py: 14/14<=→<): 8 failures, as expectedpre-commitpasses on all changed files