[PR 3/5] ltx2_block_benchmark: per-generation scoped VMEM, CFG batch multiplier, bool masks - #496
jitendra-jalwaniya wants to merge 1 commit into
Conversation
There was a problem hiding this comment.
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.
54af45c to
5d55cc9
Compare
ecc5b5e to
55e6847
Compare
| 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 = {} |
There was a problem hiding this comment.
You can use src/maxdiffusion/tpu_utils.py to determine tpu type. Let's import that instead of building custom logic
| _SCOPED_VMEM_LIMIT_KIB = { | ||
| "TPU v6 lite": 32768, | ||
| "TPU7x": 65536, | ||
| } |
There was a problem hiding this comment.
Let's move this to src/maxdiffusion/tpu_utils.py. Addiitonally TPU v6 has ~128 MiB.
There was a problem hiding this comment.
Yeah agree, dropping these changes.
Had some confusion around XLA's scoped VMEM.
5d55cc9 to
278ed61
Compare
55e6847 to
2196275
Compare
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.passed as a
jax.jitcompiler option, looked up bydevice_kind: 32 MiB on v6e, 64 MiB on TPU7x, 64 MiB for other TPUs.- jit entry point:
_forwardnow takes(graphdef, state)and usesnnx.merge, instead of being wrapped innnx.jit. This makes it possible to passcompiler_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 ofint32, matching the dtype the pipeline passes.