Skip to content

fix(hi3): support batched trainside AR with dynamic KV cache - #479

Merged
CjhHa1 merged 11 commits into
Tencent-Hunyuan:mainfrom
PushpakAg:fix/issue-452
Sep 27, 2026
Merged

CjhHa1 merged 11 commits into
Tencent-Hunyuan:mainfrom
PushpakAg:fix/issue-452

Conversation

@PushpakAg

Copy link
Copy Markdown
Contributor

Summary

Cached train-side HunyuanImage3 AR only worked at B=1. The crash reported in #452 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 requests
dynamic=True only at B == 1 and lets B > 1 use the full static cache. That makes the
assertion unreachable (it is guarded by if self.dynamic) instead of monkey-patching upstream
remote 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 at
B=1 because real_pos == L there:

  • Decode mask. Upstream deletes the attention mask during gen_text decode ("Remove
    attention 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_mask supplies a bool
    mask admitting row i's keys [0, real_pos_i + step_idx), and right-pads the prefill
    mask's key axis, since a static cache returns its full max_cache_len key axis while the
    prefill mask is only L wide.
  • Token placement. Upstream re-reads the next input token with
    input_ids.gather(1, position_ids) where position_ids is real_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 shorter
    row its own pad token as the next input.
  • position_ids padding. FusedMultimodalCondition.concat padded the ragged L axis with
    0, but position_ids are the KV-cache write indices (upstream passes them straight
    through as cache_position, applied with index_copy_). Every pad position of every short
    row 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_positions pads with each row's own continuing indices instead.

B=1 keeps the dynamic cache and stays bit-identical. Batched AR without real_pos now raises
in init_state rather than failing later on an opaque SDPA shape mismatch. The four upstream
constraints are written down in a new unirl/models/hunyuan_image3/README.md ## Gotchas, with
the 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 HunyuanImage3ForCausalMM from the real
upstream 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 batched
right-padded run matches running each row independently at B=1. Batch-vs-B=1 parity is a
plumbing property, so random weights exercise it faithfully. Per CLAUDE.md §5 the harnesses
are run and quoted here, not committed.

Harness 1 — uniformly-batched conditions, prefill + decode parity, plus ablations that each
disable one part of the fix:

$ python harness_452.py
=== parity: batched right-padded AR vs independent B=1 ===
[PASS] unequal real_pos, 3 decode steps
        tokens equal=True  logp max|d|=4.768e-07  prefill-KV max|d|=1.192e-06
[PASS] extreme length skew, 5 decode steps
        tokens equal=True  logp max|d|=4.768e-07  prefill-KV max|d|=1.192e-06
[PASS] all rows padded (max real_pos < L)
        tokens equal=True  logp max|d|=0.000e+00  prefill-KV max|d|=0.000e+00
[PASS] B=1 unchanged (dynamic cache path)
        tokens equal=True  logp max|d|=0.000e+00  prefill-KV max|d|=0.000e+00

=== ablations: each part of the fix must be load-bearing ===
[FAIL] decode mask removed -> expect FAIL
        tokens equal=False  logp max|d|=9.038e-02
[FAIL] tail-append restored -> expect FAIL
        tokens equal=False  logp max|d|=1.318e-01
ablation: decode mask load-bearing=True  token placement load-bearing=True

RESULT: PASS

Harness 2 — the issue's exact construction (per-row embed_for_ar, then
HunyuanImage3FusedMultimodalCondition.concat), which is what exposed the position_ids
padding defect:

$ python harness_452_concat.py
repro shape: input_ids=(4, 64) real_pos=[13, 37, 64, 51]

=== parity: concat-built batch vs independent B=1 ===
[PASS] unequal real_pos, 3 decode steps
        tokens equal=True  logp max|d|=4.768e-07  prefill-KV max|d|=1.192e-06  (cache slot 0 max|d|=0.000e+00)
[PASS] extreme length skew, 5 decode steps
        tokens equal=True  logp max|d|=4.768e-07  prefill-KV max|d|=1.192e-06  (cache slot 0 max|d|=0.000e+00)

=== ablation: constant position_ids padding must break parity ===
[FAIL] position_ids pad reverted to 0 -> expect FAIL
        tokens equal=False  logp max|d|=1.292e-01  prefill-KV max|d|=4.069e+00  (cache slot 0 max|d|=4.069e+00)
ablation: position_ids padding load-bearing=True

RESULT: PASS

Residual logp/KV deltas of ~1e-6 are fp32 batched-GEMM noise: the batched path reduces over a
wider masked key axis than the B=1 trim, so bitwise identity is not achievable. Token
sequences are exactly equal. Reverting the diff turns both harnesses back into the reported
AssertionError at modeling_hunyuan_image_3.py:965.

Against the issue's requested checklist: B=2+ with unequal right-padded real_pos ✅;
batched-vs-independent B=1 logits 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:

$ pre-commit run --files unirl/models/hunyuan_image3/{ar.py,conditions.py,README.md}
ruff check...............................................................Passed
ruff format..............................................................Passed
recipe _target_ paths resolve............................................Passed
docstrings are one line..................................................Passed
core dependency directions hold..........................................Passed
experimental-tier boundaries hold............................(see note)...Failed

check-experimental-boundaries fails only because the local venv is Python 3.10 and the script
imports tomllib (3.11+); re-run under python3.11 it reports check-experimental-boundaries: ok. Every other hook passes.

Compatibility / Risk

  • No config, recipe, checkpoint, data-format, or resource changes.
  • B=1 is unchanged by construction: it keeps dynamic=True, gets no injected mask, and its
    token placement is arithmetically identical to the previous tail append. Verified bit-identical.
  • B>1 previously raised AssertionError, so there is no prior batched behavior to regress.
  • _pad_positions changes concat-produced position_ids padding, which the gen_image
    diffusion 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 pod
    access may want a diffusion smoke test.
  • init_state now raises ValueError when B>1 and conditions.tokenizer_output.real_pos is
    absent. Both in-tree producers (embed_for_ar and the vLLM adapter) populate it; an
    out-of-tree caller that does not would newly fail fast.
  • The fix depends on upstream remote-code behavior (modeling_hunyuan_image_3.py), so the
    README gotchas should be re-checked when the checkpoint revision moves.

Reviewer Notes

  • Start here: _cache_attention_mask and _place_token in ar.py, then _pad_positions
    in conditions.py. The README.md gotchas explain what stock upstream does and why each
    piece is load-bearing; the harness ablations above fail deliberately to prove that.
  • Scope: the issue reports one assertion; this PR fixes four coupled defects, because the
    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.
  • Duplicate-work check: scanned the 40 open PRs and the hi3/AR/KV-cache keywords; none touch
    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.
  • AI assistance: written with Claude Code. The submitter has reviewed the full diff; the
    harnesses, their ablations, and the pre-fix reproduction were run locally with the results
    quoted verbatim above.
  • Upstream HunyuanStaticCache.update could support batched dynamic trimming by slicing at the
    max 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

  • I reviewed the changed code and removed unrelated/generated artifacts.
  • I updated tests, docs, and configs where needed, or explained why not.

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.
@github-actions github-actions Bot added the need review Ready and waiting for review label Sep 17, 2026
@CjhHa1

CjhHa1 commented Sep 24, 2026

Copy link
Copy Markdown
Collaborator

Thanks for the thorough write-up. I checked the four fixes against upstream modeling_hunyuan_image_3.py and the diagnosis holds: the len(cache_position) assert in HunyuanStaticCache.update hits the batch dim, gen_text decode drops the mask, prepare_inputs_for_generation gathers the next input at position_ids (= real_pos, then +1 per step), and position_ids go straight into index_copy_ as cache_position. The approach looks right. A few things before merging:

1. Rebase onto #508 (merged). #508 removed _build_kv_cache and its silent return None fallbacks. A None cache made prefill correct and every decode step context-free, because upstream never creates a cache itself. Cache construction now lives inline in init_state, so:

  • pass dynamic=batch_size == 1 to the HunyuanStaticCache(...) call in init_state instead of adding a helper parameter;
  • past_key_values can no longer be None in the AR loop, so getattr(past_key_values, "dynamic", True) in _cache_attention_mask can be plain past_key_values.dynamic. The default currently only absorbs that removed None case.

2. pad_id in _place_token is dead. state.input_ids is only consumed by upstream's input_ids.gather(1, position_ids), and every position it reads (real_pos + s) was scattered at the previous step. The appended column's fill value is never read. I simulated this with a sentinel fill over mixed real_pos and 5 steps, and the gather never returned it. So int(getattr(transformer.config, "pad_id", 0) or 0) and the pad_id parameter can go; torch.zeros_like(input_ids[:, :1]) is enough.

3. Diffusion path needs a check. _pad_positions changes HunyuanImage3FusedMultimodalCondition.concat, which gen_image also consumes. Diffusion also uses position_ids as cache_position, so continuing indices should be a fix there too. But it's unvalidated, and rope_cache is still zero-padded on the same axis. Could you run a mixed-length diffusion prefill, or at least confirm the diffusion capture path never hits concat with ragged L? Related but out of scope: adapters/hi3.py (_pad_to(c.get("position_ids"), ..., value=0)) has the same zero-pad on the vLLM-Omni capture path. Worth a follow-up issue.

4. Description accuracy. "B=1 is unchanged / bit-identical" holds only when real_pos == L. For B=1 with right-pad, the old tail append fed the pad token back as the next input, and _place_token now fixes that. That's a behavior change (a fix), so please say so rather than calling B=1 unchanged.

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 TextSegment stops at the stop token. It's fine to leave, but please list it as not done so #452 isn't closed as fully resolved.

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

CjhHa1 commented Sep 27, 2026

Copy link
Copy Markdown
Collaborator

GPU validation on the real checkpoint (head 01d30dbe)

Setup: 1 node, 8×H20. tencent/HunyuanImage-3.0-Instruct loaded in bf16 with device_map="balanced". torch 2.10.0-rc6, transformers 5.12.1, Python 3.13. Greedy decoding (top_k=1), max_new_tokens=96.

Batched tokenization currently returns only the first row (#520), so each batch was built per row: embed_for_ar on one prompt at a time, then HunyuanImage3FusedMultimodalCondition.concat, with tokenizer_output.real_pos stacked. This exercises exactly the ragged-concat path that _pad_positions fixes. Replay ran with transformer.training = True (see #521).

  • t2t: 4 prompts, lengths [1238, 1265, 1238, 1410].
  • i2t: 2 rows with different image sizes, lengths [5340, 5798].

Results

Baseline: main's ar.py on the same batched t2t inputs raises the #452 AssertionError. This PR runs t2t and i2t to completion, with no NaN in any output. That includes the fully-masked pad query rows _pad_attn produces, so the NaN concern on this path did not show up for AR.

Rollout vs replay (the GRPO ratio check: per-sample unpadded replay over the generated tokens):

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;
  • main raises on batched input;
  • a missing tokenizer_output raises at B=1 and at B=4.

In the parity checks tokens are identical and logp differs by at most 5e-7.

Not covered

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

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

LGTM

@github-actions github-actions Bot removed the need review Ready and waiting for review label Sep 27, 2026
@CjhHa1

CjhHa1 commented Sep 27, 2026

Copy link
Copy Markdown
Collaborator

Follow-up: maintainer commits and fp32 GPU parity (head 3575710a)

Pushed on top of the author's 010227fa. I left the PR description as is. The corrections that apply to it are listed at the end.

Reviewer notes: what changed since the original commit

ar.py

  • The cache is built inline per fix(hi3): fail fast when the AR KV cache cannot be built #508, with dynamic=batch_size == 1.
  • autoregress now requires fused.prompt_lengths and conditions.tokenizer_output at every batch size. Upstream _update_model_kwargs_for_generation reads tokenizer_output.real_pos as the decode position_ids. Without it, B=1 also silently re-feeds the shifted prompt.
  • The real_pos is None fallbacks are gone, because real_pos is now always present. That removes the tail append, logits[:, -1] at prefill, and the early mask return.
  • step reads fused.prompt_lengths, which embed_for_ar already normalizes to [B], instead of re-deriving it from tokenizer_output on every step.
  • _real_pos, _cache_attention_mask and _place_token each had one caller and are inlined into step, which now branches on step_idx. The unused pad_id fill is dropped.

conditions.py (HunyuanImage3FusedMultimodalCondition.concat)

  • The pad helpers are module-level: _pad_seq(t, length, dim) and _pad_positions(t, length).
  • _pad_attn is now two _pad_seq calls, with the same zero fill as before.
  • The local _materialize is replaced by the repo's unirl.distributed.tensor.ref.hydrate.
  • The always-zero value parameter is dropped.
  • dataclasses.replace replaces the hand-listed pass-through fields, and super().concat replaces Batch.concat.__func__.

Against base, the only behavior change in concat is the position_ids padding. I checked field by field on ragged, equal-length, single-item, empty, TensorRef, rope-longer-than-L and float-mask inputs.

README.md: the ## Gotchas now state only the upstream constraints.

Test plan: fp32 settles the bf16 divergence

This run used the same setup as the previous comment: 8×H20, real HunyuanImage-3.0-Instruct, and the same 4 t2t prompts (lengths [1238, 1265, 1238, 1410]), built per row and then merged with concat (#520). Decoding was greedy for 48 tokens. The model was loaded once in bf16 and once in fp32; TF32 was off, and the KV cache dtype was logged as float32. The code was the same patch as 3575710a.

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

Follow-ups

@github-actions github-actions Bot added the approved Approved by reviewer label Sep 27, 2026
@CjhHa1
CjhHa1 merged commit 503ba82 into Tencent-Hunyuan:main Sep 27, 2026
5 checks passed
@github-actions github-actions Bot removed the approved Approved by reviewer label Sep 27, 2026
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.

fix(hi3): support batched trainside AR with dynamic KV cache

2 participants