From 278ed61daa40e3c961674711174bcb7d622c48d5 Mon Sep 17 00:00:00 2001 From: Jitendra Jalwaniya Date: Tue, 29 Sep 2026 06:53:39 +0000 Subject: [PATCH] ltx2: fix custom block_kv_compute_in fallback, framewise VAE tiling --- src/maxdiffusion/configs/ltx2_3_video.yml | 2 ++ src/maxdiffusion/configs/ltx2_video.yml | 2 ++ src/maxdiffusion/models/attention_flax.py | 11 +++++++---- src/maxdiffusion/pipelines/ltx2/ltx2_pipeline.py | 2 ++ 4 files changed, 13 insertions(+), 4 deletions(-) diff --git a/src/maxdiffusion/configs/ltx2_3_video.yml b/src/maxdiffusion/configs/ltx2_3_video.yml index 9c9c5432a..63796ed83 100644 --- a/src/maxdiffusion/configs/ltx2_3_video.yml +++ b/src/maxdiffusion/configs/ltx2_3_video.yml @@ -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. diff --git a/src/maxdiffusion/configs/ltx2_video.yml b/src/maxdiffusion/configs/ltx2_video.yml index 23a4b104a..e66d69fc7 100644 --- a/src/maxdiffusion/configs/ltx2_video.yml +++ b/src/maxdiffusion/configs/ltx2_video.yml @@ -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. diff --git a/src/maxdiffusion/models/attention_flax.py b/src/maxdiffusion/models/attention_flax.py index 3db542586..de6712235 100644 --- a/src/maxdiffusion/models/attention_flax.py +++ b/src/maxdiffusion/models/attention_flax.py @@ -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: @@ -390,14 +390,14 @@ 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 @@ -405,7 +405,10 @@ def _extract_custom_block_sizes(flash_block_sizes): # 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 diff --git a/src/maxdiffusion/pipelines/ltx2/ltx2_pipeline.py b/src/maxdiffusion/pipelines/ltx2/ltx2_pipeline.py index eeaad473d..27e554052 100644 --- a/src/maxdiffusion/pipelines/ltx2/ltx2_pipeline.py +++ b/src/maxdiffusion/pipelines/ltx2/ltx2_pipeline.py @@ -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):