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
2 changes: 2 additions & 0 deletions src/maxdiffusion/configs/ltx2_3_video.yml
Original file line number Diff line number Diff line change
Expand Up @@ -125,6 +125,8 @@ profiler_steps: 5
enable_jax_named_scopes: False

replicate_vae: False
enable_vae_tiling: False
enable_vae_slicing: False

run_text_encoder_on_tpu: False
# Dynamically disables VAE slicing and distributes the batch dimension to avoid HBM OOM for larger batch sizes.
Expand Down
2 changes: 2 additions & 0 deletions src/maxdiffusion/configs/ltx2_video.yml
Original file line number Diff line number Diff line change
Expand Up @@ -131,6 +131,8 @@ enable_jax_named_scopes: False

replicate_vae: False
use_bwe: False
enable_vae_tiling: False
enable_vae_slicing: False

run_text_encoder_on_tpu: False
# Dynamically disables VAE slicing and distributes the batch dimension to avoid HBM OOM for larger batch sizes.
Expand Down
11 changes: 7 additions & 4 deletions src/maxdiffusion/models/attention_flax.py
Original file line number Diff line number Diff line change
Expand Up @@ -381,7 +381,7 @@ def _extract_custom_block_sizes(flash_block_sizes):
bq = 4864
bkv = 1024
bkv_compute = 1024
bkv_compute_in = 1024
bkv_compute_in = None
heads_per_tile = 1
vmem_limit_bytes = None
if flash_block_sizes is not None:
Expand All @@ -390,22 +390,25 @@ def _extract_custom_block_sizes(flash_block_sizes):
bq = get("block_q", None) or bq
bkv = get("block_kv", None) or bkv
bkv_compute = get("block_kv_compute", None) or bkv_compute
bkv_compute_in = get("block_kv_compute_in", None) or bkv_compute_in
bkv_compute_in = get("block_kv_compute_in", None)
heads_per_tile = get("heads_per_tile", None) or heads_per_tile
vmem_limit_bytes = get("vmem_limit_bytes", None) or vmem_limit_bytes
else:
bq = getattr(flash_block_sizes, "block_q", None) or bq
bkv = getattr(flash_block_sizes, "block_kv", None) or bkv
bkv_compute = getattr(flash_block_sizes, "block_kv_compute", None) or bkv_compute
bkv_compute_in = getattr(flash_block_sizes, "block_kv_compute_in", None) or bkv_compute_in
bkv_compute_in = getattr(flash_block_sizes, "block_kv_compute_in", None)
heads_per_tile = getattr(flash_block_sizes, "heads_per_tile", None) or heads_per_tile
vmem_limit_bytes = getattr(flash_block_sizes, "vmem_limit_bytes", None) or vmem_limit_bytes
# A BlockSizes object carries heads_per_tile=None when the config dict omitted
# it; getattr then returns that None instead of the default, so coerce it back
# to 1 (the custom-kernel default) to keep the `heads_per_tile > 1` guards safe.
if heads_per_tile is None:
heads_per_tile = 1
bkv_compute_in = min(bkv_compute, bkv_compute_in)
if bkv_compute_in is None:
bkv_compute_in = bkv_compute
else:
bkv_compute_in = min(bkv_compute, bkv_compute_in)
return bq, bkv, bkv_compute, bkv_compute_in, heads_per_tile, vmem_limit_bytes


Expand Down
2 changes: 2 additions & 0 deletions src/maxdiffusion/pipelines/ltx2/ltx2_pipeline.py
Original file line number Diff line number Diff line change
Expand Up @@ -558,9 +558,11 @@ def enable_vae_tiling(self):
if hasattr(self.vae, "enable_tiling"):
self.vae.enable_tiling()
self.vae.use_tiling = True
self.vae.use_framewise_decoding = True

def disable_vae_tiling(self):
self.vae.use_tiling = False
self.vae.use_framewise_decoding = False

@classmethod
def load_tokenizer(cls, config: HyperParameters):
Expand Down
Loading