Repository navigation
fix(hi3): support batched trainside AR with dynamic KV cache - #479
Conversation
Cached train-side HunyuanImage3 AR only worked at B=1. Fixes Tencent-Hunyuan#452. The reported crash is upstream's `assert len(cache_position) == 1` in `HunyuanStaticCache.update`, where `len()` on a 2-D cache position is the batch dim. A dynamic cache trims the returned KV to one scalar end position, which cannot represent per-row right-pad, so `_build_kv_cache` now only requests `dynamic=True` at B=1 and lets B>1 use the full static cache. That makes the assert unreachable rather than patching upstream. Three further defects blocked B=1 parity once the crash was gone: - Upstream deletes the attention mask during `gen_text` decode, which is only safe for a per-row trimmed cache. Over a static cache the decode query also attends right-pad KV and never-written slots, so `_cache_attention_mask` admits row i's keys [0, real_pos_i + step) and right-pads the prefill mask's key axis to the cache width. - Upstream re-reads the next input token with `input_ids.gather(1, position_ids)` at `real_pos`, so a sampled token must be scattered to `real_pos_i + step` (`_place_token`); appending at the padded tail fed short rows their own pad token. - `FusedMultimodalCondition.concat` padded `position_ids` with 0, and those are the cache write indices, so every pad position of a short row index_copy_'d its KV onto slot 0 and silently replaced that row's first real token. `_pad_positions` pads with continuing indices instead. B=1 keeps the dynamic cache and is bit-identical. Batched AR without `real_pos` now raises in `init_state` instead of failing later on an opaque SDPA shape mismatch.
|
Thanks for the thorough write-up. I checked the four fixes against upstream 1. Rebase onto #508 (merged). #508 removed
2. 3. Diffusion path needs a check. 4. Description accuracy. " 5. Scope note. Finished-sequence masking from #452's checklist isn't addressed: finished rows keep sampling and writing KV. That's harmless for other rows since the mask is per-row, and the packed |
…onflict in ar.py)
…lpers Review follow-ups for Tencent-Hunyuan#479 after merging Tencent-Hunyuan#508: - Require conditions.tokenizer_output and fused.prompt_lengths in init_state at every batch size. Upstream sets the decode position_ids from tokenizer_output.real_pos; without it the loop silently re-feeds the shifted prompt, at B=1 too. This makes the real_pos-is-None fallbacks (tail append, logits[:, -1] at prefill, early mask return) unreachable, so they are removed. - Resolve real_pos once into HunyuanImage3ARState from fused.prompt_lengths, which embed_for_ar already normalizes to [B], instead of re-deriving it from tokenizer_output every step. - Inline the single-caller mask and token-placement helpers into step, branching on step_idx instead of on the presence of attention_mask, and read past_key_values.dynamic directly now that the cache is never None. - Drop the pad_id fill: the appended column is always overwritten or never gathered before it is. - README: point the gotchas at the inlined code, note that padded B=1 prompts were also fixed, document the real_pos requirement, and fix the upstream attention class name.
…ched AR - Read real_pos from conditions.fused.prompt_lengths in step instead of shadowing it on HunyuanImage3ARState; step already receives the same conditions every call. - Fold the tokenizer_output / prompt_lengths requirement into the existing fused-field check in autoregress, the one validation site. - Drop the int() on max_cache_len (already int) and the clamp on the prefill logits index (real_pos is in [1, L] by construction). - Collapse _pad_positions' early returns; move its rationale pointer to the concat docstring and drop the measured numbers from the README.
- Remove the leading empty/single-item fast path: the equal-length check already covers both, and Batch.concat raises on empty itself. - Collect sequence lengths as a set from input_ids directly; every producer of this condition sets input_ids. - Replace the three function-local Batch imports and Batch.concat.__func__(cls, ...) with super().concat(...). - Drop _pad_attn: padding the mask's key and query axes is two _pad_seq calls, with identical zero fill. - Build _pad_seq's F.pad spec directly for the negative dims it is called with, and fold its None / already-long early returns. - Rebuild each shard with dataclasses.replace so unpadded fields (scatter indices, prompt_lengths) carry over instead of being re-listed by hand.
GPU validation on the real checkpoint (head
|
| run | tokens | mean |Δlogp| | max |Δlogp| |
|---|---|---|---|
| batched B=4 | 384 | 0.029 | 1.26 |
| per-row B=1 | 96 each | 0.016 – 0.039 | 0.15 – 0.93 |
The batched mean sits inside the B=1 range. The B=1 floor itself (mean ~0.03, max 0.93) is the model's existing bf16 rollout/replay mismatch and predates this PR.
Batched vs B=1 generations (common token prefix out of 96):
- t2t:
[0, 96, 47, 19]. One row is identical over all 96 tokens. - i2t:
[72, 32].
The divergence is batch-shape numerics, not padding. A prefill-logit comparison against B=1 at the last prompt position:
| batch | per-row KL(B=1 ‖ batched) | top-1 equal |
|---|---|---|
mixed lengths (pad [172, 145, 172, 0]) |
0.276, 0.037, 0.168, 0.105 | F, T, T, T |
| control: 4 copies of the 1410-token prompt, no padding | 0.067 each | T, T, T, T |
- The zero-padding control already shows KL 0.067, which is the bf16 / MoE-routing floor for changing the batch shape.
- KL does not track pad count: the row with 145 pad tokens (KL 0.037) is below the control. The unpadded row still shows KL 0.105.
- The only top-1 flip (row 0) is a near tie: its B=1 top-1/top-2 margin is 0.31.
A mask or token-placement bug would sit well above this floor. For comparison, on CPU, reverting _pad_positions to zero-padding moves logp by 0.94.
CPU harness (fp32)
This harness runs upstream's HunyuanStaticCache, prepare_inputs_for_generation and _update_model_kwargs_for_generation verbatim, around a toy attention model. It passes all of the following:
- batched (tokenizer-style right-pad) matches per-row B=1;
- batched (per-row +
concat) matches per-row B=1; - a right-padded B=1 prompt matches the unpadded B=1 prompt;
- the post-refactor code matches the original PR head bit for bit;
mainraises on batched input;- a missing
tokenizer_outputraises at B=1 and at B=4.
In the parity checks tokens are identical and logp differs by at most 5e-7.
Not covered
- Diffusion with a ragged
concat. In-tree, CFG cond/uncond are equal-length, so this path is not reached today; see fix(hi3): batched chat-template tokenization encodes only the first sample #520 before liftingforward_batch_size: 1. - FSDP-wrapped forwards: the model was loaded with a plain
from_pretraineddevice map. - Finished-sequence masking from fix(hi3): support batched trainside AR with dynamic KV cache #452's checklist: finished rows keep sampling. This is harmless to the other rows, but fix(hi3): support batched trainside AR with dynamic KV cache #452 should not be closed as fully resolved.
… concat - Replace the local _materialize duck-type check with the canonical unirl.distributed.tensor.ref.hydrate, which tests isinstance(TensorRef) and maps an empty ref to None instead of torch.empty(0). - Drop _pad_seq's value parameter: every caller passed a zero (0, False, 0.0), which is F.pad's default fill for every dtype here.
_pad_seq and _pad_positions were closures only to capture max_L. Take the target length as an argument instead, matching the module-level _pad_* helpers in types/conditions/text.py and janus_pro/conditions.py, so they are no longer redefined on every concat call. The negative-dim constraint and the README pointer move into their docstrings.
Follow-up: maintainer commits and fp32 GPU parity (head
|
| comparison vs per-row B=1 | bf16 | fp32 |
|---|---|---|
| B=1 run twice | identical | identical |
| B=1 static cache (PR mask path) vs dynamic cache | prefill bit-identical; row 3 diverges at token 33 | identical, logp Δ 0 |
| 4 copies of one prompt (no padding) | KL 0.067; copies 0–2 diverge at 33, copy 3 identical | identical, KL 1.6e-7 |
| mixed lengths, B=4 | KL 0.04–0.28, common prefix [0, 38, 28, 33] |
48/48 identical on every row, prefill logit |Δ| ≤ 5e-5, logp |Δ| ≤ 5e-5 |
There was no NaN anywhere. In fp32 the batched path matches B=1 exactly, so the mask, the token placement and the position_ids padding are all correct.
The bf16 gaps are numerical:
- They appear with no batching at all, from static vs dynamic cache at B=1.
- They appear between identical rows of the same batch.
- They vanish in fp32.
Upstream MoE inference uses per-token easy top-k with no capacity dropping, so padding cannot couple rows through expert capacity. That leaves shape-dependent bf16 GEMM/SDPA accumulation flipping near-tied routing decisions.
I also ran SKIP=no-commit-to-branch pre-commit run --from-ref <merge-base> --to-ref HEAD; all hooks passed.
Corrections to the PR description
- "B=1 is unchanged / bit-identical" holds only when
real_pos == L. A right-padded B=1 prompt used to get its own pad token as the next input, and this PR fixes that. - "Both in-tree producers (
embed_for_arand the vLLM adapter) populate it" is wrong. The vLLM adapter only setsfused.prompt_lengths, it does so for replay, and it never reachesautoregress. - The helper names it cites (
_build_kv_cache,_cache_attention_mask,_place_token) no longer exist. The logic lives ininit_state/step. - Finished-sequence masking from fix(hi3): support batched trainside AR with dynamic KV cache #452 is not done, so fix(hi3): support batched trainside AR with dynamic KV cache #452 is only partly resolved by this PR.
- In-tree t2t/i2t still cannot build B>1 AR inputs, because the upstream tokenizer encodes only the first sample (fix(hi3): batched chat-template tokenization encodes only the first sample #520). Until that is fixed, the batched path is only reachable per row plus
concat.
Follow-ups
- fix(hi3): batched chat-template tokenization encodes only the first sample #520: batched
apply_chat_templatedrops rows. - fix(hi3): AR replay and in-process autoregress need opposite ForCausalMM.training states #521:
replayin eval mode. - refactor(vllm-omni/hi3): reuse FusedMultimodalCondition.concat for DiT capture padding (drops position_ids zero-pad) #525: the vLLM DiT capture adapter duplicates this
concatpadding and still zero-padsposition_ids. It's harmless today because replay builds no KV cache.
Summary
Cached train-side HunyuanImage3 AR only worked at
B=1. The crash reported in #452 isupstream's
assert len(cache_position) == 1inHunyuanStaticCache.update, wherelen()on a2-D cache position is the batch dim. A dynamic cache trims the returned KV to one scalar end
position, which cannot represent per-row right-pad, so
_build_kv_cachenow requestsdynamic=Trueonly atB == 1and letsB > 1use the full static cache. That makes theassertion unreachable (it is guarded by
if self.dynamic) instead of monkey-patching upstreamremote code — the direction the issue suggested first.
Removing the crash was not sufficient for the behavior the issue asks for (batched logits/KV
matching independent
B=1). Three further defects were hiding behind it, each invisible atB=1becausereal_pos == Lthere:gen_textdecode ("Removeattention mask to use full attention of 1 x seqlen"), which is only safe for a cache trimmed
to the valid length. Over a static cache the decode query also attends the right-pad KV
written at prefill and the never-written tail slots.
_cache_attention_masksupplies a boolmask admitting row
i's keys[0, real_pos_i + step_idx), and right-pads the prefillmask's key axis, since a static cache returns its full
max_cache_lenkey axis while theprefill mask is only
Lwide.input_ids.gather(1, position_ids)whereposition_idsisreal_pos(the first<pad>index, advanced once per step). The sampled token must therefore be scattered to
real_pos_i + step_idx(_place_token); appending it at the padded tail fed every shorterrow its own pad token as the next input.
position_idspadding.FusedMultimodalCondition.concatpadded the raggedLaxis with0, butposition_idsare the KV-cache write indices (upstream passes them straightthrough as
cache_position, applied withindex_copy_). Every pad position of every shortrow therefore wrote its KV onto cache slot 0 — last write wins, silently replacing that row's
first real token. Measured on the repro shape: slot 0 off by ~4.0 while every other slot was
bit-identical.
_pad_positionspads with each row's own continuing indices instead.B=1keeps the dynamic cache and stays bit-identical. Batched AR withoutreal_posnow raisesin
init_staterather than failing later on an opaque SDPA shape mismatch. The four upstreamconstraints are written down in a new
unirl/models/hunyuan_image3/README.md## Gotchas, withthe one-line docstrings pointing at it.
Related Issue
Fixes #452
Test Plan
Hardware available for this change was one 4 GB RTX 3050, so the ~80B checkpoint could not be
loaded. Validation instead builds a tiny random-init
HunyuanImage3ForCausalMMfrom the realupstream remote code (
tencent/HunyuanImage-3.0-Instruct,skip_load_module={"vae","vit"},2 layers / 128 hidden / dense MLP, fp32 CPU, greedy
top_k=1) and asserts that a batchedright-padded run matches running each row independently at
B=1. Batch-vs-B=1parity is aplumbing property, so random weights exercise it faithfully. Per
CLAUDE.md§5 the harnessesare run and quoted here, not committed.
Harness 1 — uniformly-batched conditions, prefill + decode parity, plus ablations that each
disable one part of the fix:
Harness 2 — the issue's exact construction (per-row
embed_for_ar, thenHunyuanImage3FusedMultimodalCondition.concat), which is what exposed theposition_idspadding defect:
Residual
logp/KV deltas of ~1e-6 are fp32 batched-GEMM noise: the batched path reduces over awider masked key axis than the
B=1trim, so bitwise identity is not achievable. Tokensequences are exactly equal. Reverting the diff turns both harnesses back into the reported
AssertionErroratmodeling_hunyuan_image_3.py:965.Against the issue's requested checklist:
B=2+with unequal right-paddedreal_pos✅;batched-vs-independent
B=1logits and KV parity ✅; two-plus decode steps after prefill ✅(3 and 5); cache growth ✅. Not run: FSDP train-side forwards, the real checkpoint, and
finished-sequence masking beyond the per-row cache mask — all need the 8×H20 pod.
Hooks:
check-experimental-boundariesfails only because the local venv is Python 3.10 and the scriptimports
tomllib(3.11+); re-run underpython3.11it reportscheck-experimental-boundaries: ok. Every other hook passes.Compatibility / Risk
B=1is unchanged by construction: it keepsdynamic=True, gets no injected mask, and itstoken placement is arithmetically identical to the previous tail append. Verified bit-identical.
B>1previously raisedAssertionError, so there is no prior batched behavior to regress._pad_positionschangesconcat-producedposition_idspadding, which thegen_imagediffusion path also consumes. Unique per-row cache write indices are required there for the
same
index_copy_reason, but only the AR path was validated in this PR; a reviewer with podaccess may want a diffusion smoke test.
init_statenow raisesValueErrorwhenB>1andconditions.tokenizer_output.real_posisabsent. Both in-tree producers (
embed_for_arand the vLLM adapter) populate it; anout-of-tree caller that does not would newly fail fast.
modeling_hunyuan_image_3.py), so theREADME gotchas should be re-checked when the checkpoint revision moves.
Reviewer Notes
_cache_attention_maskand_place_tokeninar.py, then_pad_positionsin
conditions.py. TheREADME.mdgotchas explain what stock upstream does and why eachpiece is load-bearing; the harness ablations above fail deliberately to prove that.
assertion alone is only the crash. Fixing it in isolation would have produced silently wrong
batched logits — worse than the crash. The issue's suggested direction 1 was conditional on
"if the full static cache remains correctly masked"; it is not, which is why the decode mask
is here rather than a bare
dynamic=False.this path (fix(vllm-omni): retire per-request AR seed patch #474 is vLLM-Omni AR seeding, perf(qwen3-omni): project only final prefill logits #446 is qwen3-omni prefill logits). No overlap.
harnesses, their ablations, and the pre-fix reproduction were run locally with the results
quoted verbatim above.
HunyuanStaticCache.updatecould support batched dynamic trimming by slicing at themax active position, which was the issue's direction 2. It is not used here because it still
needs the same per-row decode mask, and it would mean patching remote code that the checkpoint
owns. Worth an upstream report regardless.
Checklist