Skip to content

[PR 2/5] LTX2: fix block_kv_compute_in fallback, framewise VAE tiling, v6e tile-search VMEM cap - #495

Closed
jitendra-jalwaniya wants to merge 1 commit into
pyink_fixes_480from
ltx2_tpu_fixes
Closed

jitendra-jalwaniya wants to merge 1 commit into
pyink_fixes_480from
ltx2_tpu_fixes

Conversation

@jitendra-jalwaniya

@jitendra-jalwaniya jitendra-jalwaniya commented Sep 29, 2026 •

Copy link
Copy Markdown
Collaborator

Small, independent LTX-2 fixes for TPU, split out of #489.

- **`attention_flax.py`**: `_extract_custom_block_sizes` used to default `block_kv_compute_in` to 1024 and then take `min(block_kv_compute, 1024)`. So any tuned `block_kv_compute` above 1024 was silently cut down unless

block_kv_compute_in was set explicitly. It now falls back to block_kv_compute when unset, and is only clamped when it is given.
- ltx2_pipeline.py: enable_vae_tiling() / disable_vae_tiling() now also toggle vae.use_framewise_decoding, so tiled decoding also decodes frame by frame.
- ltx2_video.yml / ltx2_3_video.yml: add enable_vae_tiling: False and enable_vae_slicing: False so the options are declared and can be overridden from the command line.
- generate_ltx2.py: on v6e, cap the tile-search VMEM budget at 32 MiB, down from the 64 MiB default. This keeps the grid search from picking tiles that don't fit v6e's scoped VMEM.

@github-actions

Copy link
Copy Markdown

@gemini-code-assist gemini-code-assist Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

Code Review

This pull request introduces several configuration and logic updates for LTX2 video generation. Specifically, it adds VAE tiling and slicing options to the configuration files, caps the VMEM limit to 32MB on v6 JAX devices, refactors block size extraction in attention_flax.py to handle bkv_compute_in dynamically, and updates the VAE pipeline to toggle use_framewise_decoding alongside VAE tiling. There are no review comments, and I have no feedback to provide.

@jitendra-jalwaniya jitendra-jalwaniya changed the title ltx2: fix custom block size fallback, framewise VAE tiling, v6 tile-s… [PR 2/5] LTX2: fix block_kv_compute_in fallback, framewise VAE tiling, v6e tile-search VMEM cap Sep 29, 2026
@jitendra-jalwaniya
jitendra-jalwaniya requested review from Perseus14 and removed request for entrpn September 29, 2026 07:54
Comment thread src/maxdiffusion/generate_ltx2.py Outdated
keys.get("tile_search_vmem_limit_bytes") or config.flash_block_sizes.get("vmem_limit_bytes") or 64 * 1024 * 1024
)
if "v6" in jax.devices()[0].device_kind.lower():
vmem_limit_bytes = min(vmem_limit_bytes, 32 * 1024 * 1024)

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.

v6e has ~128 MiB, why are we scoping it to 32 MiB

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

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

Thanks for highlighting this issue.
I had some confusion around XLA's scoped VMEM vs full VMEM.
Remove these changes.

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.

2 participants