Skip to content

[PR 3/5] ltx2_block_benchmark: per-generation scoped VMEM, CFG batch multiplier, bool masks - #496

Closed
jitendra-jalwaniya wants to merge 1 commit into
ltx2_tpu_fixesfrom
ltx2_block_benchmark_fixes
Closed

jitendra-jalwaniya wants to merge 1 commit into
ltx2_tpu_fixesfrom
ltx2_block_benchmark_fixes

Conversation

@jitendra-jalwaniya

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

Copy link
Copy Markdown
Collaborator

Fixes to the LTX-2 tile-search block benchmark (utils/ltx2_block_benchmark.py) so it better reflects the real inference step. Split out of #489.

- **Scoped VMEM per TPU generation**: drop `--xla_tpu_scoped_vmem_limit_kib=65536` from `LIBTPU_INIT_ARGS`. By the time the benchmark calls `jax.devices()`, libtpu is already initialized, so the flag was never read. The limit is now

passed as a jax.jit compiler option, looked up by device_kind: 32 MiB on v6e, 64 MiB on TPU7x, 64 MiB for other TPUs.
- jit entry point: _forward now takes (graphdef, state) and uses nnx.merge, instead of being wrapped in nnx.jit. This makes it possible to pass compiler_options.
- CFG/STG batch: the default batch is now multiplied by the guidance fan-out (×2 for CFG, ×4 for CFG+STG), matching the effective batch the pipeline runs through the transformer.
- vmem_bytes(): divides the VMEM budget by the per-device batch (batch / (data × fsdp)), so larger per-device batches get a proportionally smaller per-tile budget.
- Masks: prompt attention masks are now bool_ instead of int32, matching the dtype the pipeline passes.

@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 updates the LTX2 block benchmark to dynamically configure XLA compiler options (specifically xla_tpu_scoped_vmem_limit_kib) based on the detected TPU device kind, and refactors the model execution to use nnx.split and nnx.merge with JIT compilation. It also adjusts the batch size calculation to account for CFG and STG scales, and changes prompt attention masks to boolean types. The review feedback suggests expanding the TPU VMEM limit mappings to support TPU v4 and v5 lite, avoiding the division of VMEM by the local batch size since batch dimensions are processed sequentially, and correcting the batch multiplier logic when STG is enabled without CFG.

Comment thread src/maxdiffusion/utils/ltx2_block_benchmark.py Outdated
Comment thread src/maxdiffusion/utils/ltx2_block_benchmark.py
Comment thread src/maxdiffusion/utils/ltx2_block_benchmark.py
@jitendra-jalwaniya jitendra-jalwaniya changed the title ltx2_block_benchmark: per-generation scoped VMEM, CFG batch multiplier, bool masks [PR 3/5] ltx2_block_benchmark: per-generation scoped VMEM, CFG batch multiplier, bool masks Sep 29, 2026
@jitendra-jalwaniya
jitendra-jalwaniya requested review from Perseus14 and removed request for entrpn September 29, 2026 07:55
Comment on lines +146 to +153
device_kind = jax.devices()[0].device_kind
if "tpu" in device_kind.lower():
scoped_vmem = _SCOPED_VMEM_LIMIT_KIB.get(device_kind, 65536)
self._compiler_options = {
"xla_tpu_scoped_vmem_limit_kib": str(scoped_vmem),
}
else:
self._compiler_options = {}

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.

You can use src/maxdiffusion/tpu_utils.py to determine tpu type. Let's import that instead of building custom logic

Comment on lines +80 to +83
_SCOPED_VMEM_LIMIT_KIB = {
"TPU v6 lite": 32768,
"TPU7x": 65536,
}

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.

Let's move this to src/maxdiffusion/tpu_utils.py. Addiitonally TPU v6 has ~128 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.

Yeah agree, dropping these changes.
Had some confusion around XLA's scoped VMEM.

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