Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
33 changes: 33 additions & 0 deletions src/maxdiffusion/tpu_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
41 changes: 31 additions & 10 deletions src/maxdiffusion/utils/ltx2_block_benchmark.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand All @@ -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,
Expand All @@ -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,
Expand All @@ -90,6 +89,7 @@ def _forward(
height,
width,
):
model = nnx.merge(graphdef, state)
return model(
hidden_states=latents,
audio_hidden_states=audio_latents,
Expand Down Expand Up @@ -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
Expand All @@ -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)
Comment thread
jitendra-jalwaniya marked this conversation as resolved.
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()

Expand All @@ -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
Comment thread
jitendra-jalwaniya marked this conversation as resolved.

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,
Expand All @@ -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,
Expand Down Expand Up @@ -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())
Expand Down
Loading