Skip to content

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

Description

@CjhHa1

System info

Rollout / inference engine

Train-side only

Domain

unified_model / HunyuanImage3 AR

Recipe / config-name

Based on unified_model/hi3_vllmomni, with the AR prefill isolated under its production FSDP/LoRA settings.

How you ran it

  • Modified recipe or my own script

Reproduction

Build four unequal-length text prompts through HunyuanImage3TextEmbedStage.embed_for_ar, combine them with HunyuanImage3FusedMultimodalCondition.concat, then initialize and run the first cached AR step:

embedded_rows = [
    embedder.embed_for_ar(Texts(texts=[prompt]), bot_task="think", max_length=512)
    for prompt in prompts
]
fused = HunyuanImage3FusedMultimodalCondition.concat(
    [row["fused"] for row in embedded_rows]
)
real_pos = torch.cat(
    [row["tokenizer_output"].real_pos for row in embedded_rows], dim=0
)
conditions = HunyuanImage3ARConditions(
    fused=fused,
    tokenizer_output=SimpleNamespace(real_pos=real_pos),
)
state = HunyuanImage3ARStep().init_state(bundle, conditions, max_new_tokens=1)
HunyuanImage3ARStep().step(bundle, conditions, state)

Observed input:

input_ids.shape = [4, 490]
real_pos = [13, 220, 490, 450]

Traceback:

HunyuanImage3ARStep.step
  -> HunyuanImage3ForCausalMM.forward
  -> HunyuanImage3Attention.forward
  -> HunyuanStaticCache.update

File modeling_hunyuan_image_3.py, line 965, in update
    assert len(cache_position) == 1
AssertionError

UniRL constructs this cache with dynamic=True in unirl/models/hunyuan_image3/ar.py::_build_kv_cache. For a 2D cache position, the upstream cache updates each batch row, then explicitly rejects every dynamic batch larger than one.

B=1, L=490 runs successfully with identical logits and KV across repeated first-prefill forwards.

Expected behavior

Cached trainside HunyuanImage3 AR should support B > 1, including right-padded prompts with different real_pos values. Batched logits and KV should match running each sample independently at B=1 for both prefill and subsequent decode steps.

Suggested validation / possible fixes

Possible implementation directions:

  1. Use dynamic=False when batch_size > 1, if the full static cache remains correctly masked; the diffusion path already constructs HunyuanStaticCache(dynamic=False).
  2. Extend the dynamic 2D path to return a slice ending at the maximum active cache position instead of asserting len(cache_position) == 1.

Please validate before choosing either approach:

  • B=2+ with unequal right-padded prompt lengths / real_pos;
  • batched vs independent B=1 logits and KV parity;
  • at least two decode steps after prefill;
  • finished-sequence masking and cache growth;
  • both plain and FSDP trainside forwards.

Impact

The surrounding UniRL AR implementation is batch-shaped, but this cache assertion currently limits the actual cached trainside path to one sample. External vLLM-Omni rollout is unaffected.

Activity

  1. PushpakAg commented on Sep 16, 2026

    @PushpakAg
    Contributor

    Looks like there's an assertion error in HunyuanStaticCache.update when trying to handle batches larger than one with a dynamic cache. I'd check how the cache positions are being managed in modeling_hunyuan_image_3.py.

  2. CjhHa1 commented on Sep 17, 2026

    @CjhHa1
    CollaboratorAuthor

    Looks like there's an assertion error in HunyuanStaticCache.update when trying to handle batches larger than one with a dynamic cache. I'd check how the cache positions are being managed in modeling_hunyuan_image_3.py.

    good, I'll assign you

  3. added a commit that references this issue on Sep 27, 2026
    503ba82
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Labels

bugSomething isn't working

Type

No type

Projects

No projects

    Milestone

    No milestone

    Relationships

    None yet

    Development

    No branches or pull requests

    Issue actions