From 21962751c27b9775673475dbede3db2d82adaa1c Mon Sep 17 00:00:00 2001 From: Jitendra Jalwaniya Date: Tue, 29 Sep 2026 06:53:39 +0000 Subject: [PATCH] ltx2_block_benchmark: scoped VMEM via tpu_utils, CFG batch multiplier, bool masks --- src/maxdiffusion/tpu_utils.py | 33 +++++++++++++++ .../utils/ltx2_block_benchmark.py | 41 ++++++++++++++----- 2 files changed, 64 insertions(+), 10 deletions(-) diff --git a/src/maxdiffusion/tpu_utils.py b/src/maxdiffusion/tpu_utils.py index ad39686c9..640b5fbc9 100644 --- a/src/maxdiffusion/tpu_utils.py +++ b/src/maxdiffusion/tpu_utils.py @@ -78,3 +78,36 @@ def get_tpu_type() -> TpuType: return TpuType.UNKNOWN except Exception: return TpuType.UNKNOWN + + +# XLA's per-op scratch budget used across maxdiffusion (README, CI, run scripts). +_DEFAULT_SCOPED_VMEM_LIMIT_KIB = 64 * 1024 + + +def get_vmem_capacity_bytes() -> int | None: + """Returns the physical VMEM per TensorCore, or None when not on a TPU. + + For example, 128 MiB on TPU v6e and 64 MiB on TPU7x. + """ + try: + if jax.devices()[0].platform != "tpu": + return None + from jax.experimental.pallas import tpu as pltpu # pylint: disable=import-outside-toplevel + + return int(pltpu.get_tpu_info().vmem_capacity_bytes) + except Exception: # pylint: disable=broad-exception-caught + return None + + +def get_scoped_vmem_limit_kib(requested_kib: int = _DEFAULT_SCOPED_VMEM_LIMIT_KIB) -> int | None: + """Returns a value for `xla_tpu_scoped_vmem_limit_kib`, or None when not on a TPU. + + Scoped VMEM is the scratch budget XLA grants a single fusion (and a Pallas + kernel that does not set its own `vmem_limit_bytes`); it is a slice of the + physical VMEM, not the whole of it. The request is clamped to the physical + capacity so the same value is valid on every TPU generation. + """ + capacity = get_vmem_capacity_bytes() + if capacity is None: + return None + return min(requested_kib, capacity // 1024) diff --git a/src/maxdiffusion/utils/ltx2_block_benchmark.py b/src/maxdiffusion/utils/ltx2_block_benchmark.py index 564c03594..406a7bb06 100644 --- a/src/maxdiffusion/utils/ltx2_block_benchmark.py +++ b/src/maxdiffusion/utils/ltx2_block_benchmark.py @@ -32,7 +32,6 @@ "--xla_tpu_enable_async_collective_fusion_multiple_steps=true", "--xla_tpu_overlap_compute_collective_tc=true", "--xla_enable_async_all_gather=true", - "--xla_tpu_scoped_vmem_limit_kib=65536", "--xla_tpu_enable_async_all_to_all=true", "--xla_tpu_enable_all_experimental_scheduler_features=true", "--xla_tpu_enable_latency_hiding_scheduler=true", @@ -48,7 +47,7 @@ from flax import nnx from flax.linen import partitioning as nn_partitioning -from maxdiffusion import max_logging, max_utils, pyconfig +from maxdiffusion import max_logging, max_utils, pyconfig, tpu_utils from maxdiffusion.models.ltx2.transformer_ltx2 import LTX2VideoTransformer3DModel from maxdiffusion.utils.tile_size_grid_search import ( BenchResult, @@ -75,9 +74,9 @@ def tiled_seq_len(full_seq: int, attention: str, context_shards: int, ulysses_sh return local_tiled_seq_len(full_seq, attention, context_shards, ulysses_shards) -@functools.partial(nnx.jit, static_argnames=("num_frames", "height", "width")) def _forward( - model, + graphdef, + state, latents, timestep, prompt_embeds, @@ -90,6 +89,7 @@ def _forward( height, width, ): + model = nnx.merge(graphdef, state) return model( hidden_states=latents, audio_hidden_states=audio_latents, @@ -134,6 +134,13 @@ def __init__( self._ulysses_shards = int(getattr(config, "ulysses_shards", 1) or 1) self._vmem = int(vmem_limit_bytes) self.label = f"ltx2/{self._attention}/u{self._ulysses_shards}" + # Passed as a compile option (not via LIBTPU_INIT_ARGS) because libtpu is + # already initialized by the time the benchmark runs, after which + # LIBTPU_INIT_ARGS is no longer read. + scoped_vmem_kib = tpu_utils.get_scoped_vmem_limit_kib() + self._compiler_options = {} + if scoped_vmem_kib is not None: + self._compiler_options["xla_tpu_scoped_vmem_limit_kib"] = str(scoped_vmem_kib) self._lf = (num_frames - 1) // 8 + 1 self._lh, self._lw = height // 32, width // 32 @@ -143,7 +150,13 @@ def __init__( self._width_orig = width data_shards = int(mesh.shape.get("data", 1)) * int(mesh.shape.get("fsdp", 1)) - self._batch = batch if batch is not None else max(1, data_shards) + guidance_scale = getattr(config, "guidance_scale", 1.0) + stg_scale = getattr(config, "stg_scale", 0.0) + do_cfg = guidance_scale > 1.0 + do_stg = stg_scale > 0.0 + cfg_mult = 4 if (do_cfg and do_stg) else (2 if do_cfg else 1) + default_batch = max(1, data_shards) * cfg_mult + self._batch = batch if batch is not None else default_batch self._hf_cfg = LTX2VideoTransformer3DModel.load_config(config.pretrained_model_name_or_path, subfolder="transformer") self._inputs = self._make_inputs() @@ -163,13 +176,21 @@ def tiled_seq_lens(self): return (s, s) def vmem_bytes(self): - return self._vmem + data_shards = int(self._mesh.shape.get("data", 1)) * int(self._mesh.shape.get("fsdp", 1)) + local_batch = max(1, self._batch // max(1, data_shards)) + return self._vmem // local_batch def run(self, bq, bkv, *, bkv_compute=None, iters=10, warmup=2): cmp = bkv_compute or bkv try: with self._mesh: model = self._build_model(bq, bkv, cmp) + graphdef, state = nnx.split(model) + forward = jax.jit( + functools.partial(_forward, graphdef), + static_argnames=("num_frames", "height", "width"), + compiler_options=self._compiler_options, + ) ( latents, timestep, @@ -181,8 +202,8 @@ def run(self, bq, bkv, *, bkv_compute=None, iters=10, warmup=2): ) = self._inputs with self._mesh, nn_partitioning.axis_rules(self._rules): mean, std, times, compile_ms = time_callable( - lambda: _forward( - model, + lambda: forward( + state, latents, timestep, prompt_embeds, @@ -268,10 +289,10 @@ def _make_inputs(self): # Prompts prompt_embeds = jax.random.normal(k2, (self._batch, 1024, 3840), dtype) - prompt_attention_mask = jnp.ones((self._batch, 1024), dtype=jnp.int32) + prompt_attention_mask = jnp.ones((self._batch, 1024), dtype=jnp.bool_) audio_prompt_embeds = jax.random.normal(k4, (self._batch, 1024, 3840), dtype) - audio_prompt_attention_mask = jnp.ones((self._batch, 1024), dtype=jnp.int32) + audio_prompt_attention_mask = jnp.ones((self._batch, 1024), dtype=jnp.bool_) timestep = jnp.zeros((self._batch,), jnp.float32) repl = jax.sharding.NamedSharding(self._mesh, jax.sharding.PartitionSpec())