From e5a15a4777082c2d836aae52b460d6ce373139d4 Mon Sep 17 00:00:00 2001 From: sync-bot Date: Mon, 12 Jan 2026 14:14:33 +0000 Subject: [PATCH] Automated PR - 2026-01-12 --- README.md | 2 +- .../src/ltx_core/model/video_vae/resnet.py | 2 +- .../src/ltx_core/model/video_vae/video_vae.py | 21 ++++++++------- .../gemma/encoders/base_encoder.py | 26 ++++++++++--------- .../src/ltx_pipelines/distilled.py | 4 ++- .../src/ltx_pipelines/ic_lora.py | 4 ++- .../ltx_pipelines/keyframe_interpolation.py | 4 ++- .../src/ltx_pipelines/ti2vid_one_stage.py | 2 +- .../src/ltx_pipelines/ti2vid_two_stages.py | 4 ++- .../ltx-trainer/configs/ltx2_av_lora.yaml | 4 +++ .../ltx-trainer/configs/ltx2_v2v_ic_lora.yaml | 4 +++ .../docs/configuration-reference.md | 6 +++-- .../ltx-trainer/src/ltx_trainer/config.py | 5 ++++ .../ltx-trainer/src/ltx_trainer/trainer.py | 9 +++++++ 14 files changed, 66 insertions(+), 31 deletions(-) diff --git a/README.md b/README.md index 777beec..eb99a45 100644 --- a/README.md +++ b/README.md @@ -3,7 +3,7 @@ [![Website](https://img.shields.io/badge/Website-LTX-181717?logo=google-chrome)](https://ltx.io) [![Model](https://img.shields.io/badge/HuggingFace-Model-orange?logo=huggingface)](https://huggingface.co/Lightricks/LTX-2) [![Demo](https://img.shields.io/badge/Demo-Try%20Now-brightgreen?logo=vercel)](https://app.ltx.studio/ltx-2-playground/i2v) -[![Paper](https://img.shields.io/badge/Paper-PDF-EC1C24?logo=adobeacrobatreader&logoColor=white)](https://videos.ltx.io/LTX-2/grants/LTX_2_Technical_Report_compressed.pdf) +[![Paper](https://img.shields.io/badge/Paper-PDF-EC1C24?logo=adobeacrobatreader&logoColor=white)](https://arxiv.org/abs/2601.03233) [![Discord](https://img.shields.io/badge/Join-Discord-5865F2?logo=discord)](https://discord.gg/ltxplatform) **LTX-2** is the first DiT-based audio-video foundation model that contains all core capabilities of modern video generation in one model: synchronized audio and video, high fidelity, multiple performance modes, production-ready outputs, API access, and open access. diff --git a/packages/ltx-core/src/ltx_core/model/video_vae/resnet.py b/packages/ltx-core/src/ltx_core/model/video_vae/resnet.py index 300e66b..1423f2f 100644 --- a/packages/ltx-core/src/ltx_core/model/video_vae/resnet.py +++ b/packages/ltx-core/src/ltx_core/model/video_vae/resnet.py @@ -99,7 +99,7 @@ class ResnetBlock3D(nn.Module): self.timestep_conditioning = timestep_conditioning if timestep_conditioning: - self.scale_shift_table = nn.Parameter(torch.randn(4, in_channels) / in_channels**0.5) + self.scale_shift_table = nn.Parameter(torch.zeros(4, in_channels)) def _feed_spatial_noise( self, diff --git a/packages/ltx-core/src/ltx_core/model/video_vae/video_vae.py b/packages/ltx-core/src/ltx_core/model/video_vae/video_vae.py index c65a3e0..4638407 100644 --- a/packages/ltx-core/src/ltx_core/model/video_vae/video_vae.py +++ b/packages/ltx-core/src/ltx_core/model/video_vae/video_vae.py @@ -1,5 +1,5 @@ from dataclasses import replace -from typing import Any, Callable, Iterator, List, Optional, Tuple +from typing import Any, Callable, Iterator, List, Tuple import torch from einops import rearrange @@ -521,12 +521,11 @@ class VideoDecoder(nn.Module): ) self.last_scale_shift_table = nn.Parameter(torch.empty(2, feature_channels)) - # def forward(self, sample: torch.Tensor, target_shape) -> torch.Tensor: def forward( self, sample: torch.Tensor, - timestep: Optional[torch.Tensor] = None, - generator: Optional[torch.Generator] = None, + timestep: torch.Tensor | None = None, + generator: torch.Generator | None = None, ) -> torch.Tensor: r""" Decode latent representation into video frames. @@ -651,8 +650,8 @@ class VideoDecoder(nn.Module): self, latent: torch.Tensor, tiling_config: TilingConfig | None = None, - timestep: Optional[torch.Tensor] = None, - generator: Optional[torch.Generator] = None, + timestep: torch.Tensor | None = None, + generator: torch.Generator | None = None, ) -> Iterator[torch.Tensor]: """ Decode a latent tensor into video frames using tiled processing. @@ -769,8 +768,8 @@ class VideoDecoder(nn.Module): group_tiles: List[Tile], buffer: torch.Tensor, latent: torch.Tensor, - timestep: Optional[torch.Tensor], - generator: Optional[torch.Generator], + timestep: torch.Tensor | None, + generator: torch.Generator | None, ) -> torch.Tensor: """ Decode and accumulate all tiles of a temporal group into a local buffer. @@ -815,6 +814,7 @@ def decode_video( latent: torch.Tensor, video_decoder: VideoDecoder, tiling_config: TilingConfig | None = None, + generator: torch.Generator | None = None, ) -> Iterator[torch.Tensor]: """ Decode a video latent tensor with the given decoder. @@ -822,6 +822,7 @@ def decode_video( latent: Tensor [c, f, h, w] video_decoder: Decoder module. tiling_config: Optional tiling settings. + generator: Optional random generator for deterministic decoding. Yields: Decoded chunk [f, h, w, c], uint8 in [0, 255]. """ @@ -832,10 +833,10 @@ def decode_video( return frames if tiling_config is not None: - for frames in video_decoder.tiled_decode(latent, tiling_config): + for frames in video_decoder.tiled_decode(latent, tiling_config, generator=generator): yield convert_to_uint8(frames) else: - decoded_video = video_decoder(latent) + decoded_video = video_decoder(latent, generator=generator) yield convert_to_uint8(decoded_video) diff --git a/packages/ltx-core/src/ltx_core/text_encoders/gemma/encoders/base_encoder.py b/packages/ltx-core/src/ltx_core/text_encoders/gemma/encoders/base_encoder.py index e689c1a..9fca260 100644 --- a/packages/ltx-core/src/ltx_core/text_encoders/gemma/encoders/base_encoder.py +++ b/packages/ltx-core/src/ltx_core/text_encoders/gemma/encoders/base_encoder.py @@ -32,7 +32,6 @@ class GemmaTextEncoderModelBase(torch.nn.Module): dtype: torch.dtype = torch.bfloat16, ) -> None: super().__init__() - self._gemma_root = None self.tokenizer = tokenizer self.model = model self.processor = img_processor @@ -73,12 +72,6 @@ class GemmaTextEncoderModelBase(torch.nn.Module): ) return projected, attention_mask - def _init_image_processor(self) -> None: - img_processor = AutoImageProcessor.from_pretrained(self._gemma_root, local_files_only=True) - if not self.tokenizer: - raise ValueError("Tokenizer is not loaded, cannot load image processor") - self.processor = Gemma3Processor(image_processor=img_processor, tokenizer=self.tokenizer.tokenizer) - def _enhance( self, messages: list[dict[str, str]], @@ -86,8 +79,6 @@ class GemmaTextEncoderModelBase(torch.nn.Module): max_new_tokens: int = 512, seed: int = 42, ) -> str: - if self.processor is None: - self._init_image_processor() text = self.processor.tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True) model_inputs = self.processor( @@ -242,17 +233,23 @@ def _find_matching_dir(root_path: str, pattern: str) -> str: def module_ops_from_gemma_root(gemma_root: str) -> tuple[ModuleOps, ...]: gemma_path = _find_matching_dir(gemma_root, "model*.safetensors") tokenizer_path = _find_matching_dir(gemma_root, "tokenizer.model") + processor_path = _find_matching_dir(gemma_root, "preprocessor_config.json") def load_gemma(module: GemmaTextEncoderModelBase) -> GemmaTextEncoderModelBase: module.model = Gemma3ForConditionalGeneration.from_pretrained( gemma_path, local_files_only=True, torch_dtype=torch.bfloat16 ) - module._gemma_root = module._gemma_root or gemma_root return module def load_tokenizer(module: GemmaTextEncoderModelBase) -> GemmaTextEncoderModelBase: module.tokenizer = LTXVGemmaTokenizer(tokenizer_path, 1024) - module._gemma_root = module._gemma_root or gemma_root + return module + + def load_processor(module: GemmaTextEncoderModelBase) -> GemmaTextEncoderModelBase: + image_processor = AutoImageProcessor.from_pretrained(processor_path, local_files_only=True) + if not module.tokenizer: + raise ValueError("Tokenizer model operation must be performed before processor model operation") + module.processor = Gemma3Processor(image_processor=image_processor, tokenizer=module.tokenizer.tokenizer) return module gemma_load_ops = ModuleOps( @@ -265,7 +262,12 @@ def module_ops_from_gemma_root(gemma_root: str) -> tuple[ModuleOps, ...]: matcher=lambda module: isinstance(module, GemmaTextEncoderModelBase) and module.tokenizer is None, mutator=load_tokenizer, ) - return (gemma_load_ops, tokenizer_load_ops) + processor_load_ops = ModuleOps( + "ProcessorLoad", + matcher=lambda module: isinstance(module, GemmaTextEncoderModelBase) and module.processor is None, + mutator=load_processor, + ) + return (gemma_load_ops, tokenizer_load_ops, processor_load_ops) def encode_text(text_encoder: GemmaTextEncoderModelBase, prompts: list[str]) -> list[tuple[torch.Tensor, torch.Tensor]]: diff --git a/packages/ltx-pipelines/src/ltx_pipelines/distilled.py b/packages/ltx-pipelines/src/ltx_pipelines/distilled.py index e3b51ba..cd7a6f2 100644 --- a/packages/ltx-pipelines/src/ltx_pipelines/distilled.py +++ b/packages/ltx-pipelines/src/ltx_pipelines/distilled.py @@ -185,7 +185,9 @@ class DistilledPipeline: del video_encoder cleanup_memory() - decoded_video = vae_decode_video(video_state.latent, self.model_ledger.video_decoder(), tiling_config) + decoded_video = vae_decode_video( + video_state.latent, self.model_ledger.video_decoder(), tiling_config, generator + ) decoded_audio = vae_decode_audio( audio_state.latent, self.model_ledger.audio_decoder(), self.model_ledger.vocoder() ) diff --git a/packages/ltx-pipelines/src/ltx_pipelines/ic_lora.py b/packages/ltx-pipelines/src/ltx_pipelines/ic_lora.py index f9826a0..42b66ae 100644 --- a/packages/ltx-pipelines/src/ltx_pipelines/ic_lora.py +++ b/packages/ltx-pipelines/src/ltx_pipelines/ic_lora.py @@ -223,7 +223,9 @@ class ICLoraPipeline: del video_encoder cleanup_memory() - decoded_video = vae_decode_video(video_state.latent, self.stage_2_model_ledger.video_decoder(), tiling_config) + decoded_video = vae_decode_video( + video_state.latent, self.stage_2_model_ledger.video_decoder(), tiling_config, generator + ) decoded_audio = vae_decode_audio( audio_state.latent, self.stage_2_model_ledger.audio_decoder(), self.stage_2_model_ledger.vocoder() ) diff --git a/packages/ltx-pipelines/src/ltx_pipelines/keyframe_interpolation.py b/packages/ltx-pipelines/src/ltx_pipelines/keyframe_interpolation.py index 553a850..30648be 100644 --- a/packages/ltx-pipelines/src/ltx_pipelines/keyframe_interpolation.py +++ b/packages/ltx-pipelines/src/ltx_pipelines/keyframe_interpolation.py @@ -223,7 +223,9 @@ class KeyframeInterpolationPipeline: del video_encoder cleanup_memory() - decoded_video = vae_decode_video(video_state.latent, self.stage_2_model_ledger.video_decoder(), tiling_config) + decoded_video = vae_decode_video( + video_state.latent, self.stage_2_model_ledger.video_decoder(), tiling_config, generator + ) decoded_audio = vae_decode_audio( audio_state.latent, self.stage_2_model_ledger.audio_decoder(), self.stage_2_model_ledger.vocoder() ) diff --git a/packages/ltx-pipelines/src/ltx_pipelines/ti2vid_one_stage.py b/packages/ltx-pipelines/src/ltx_pipelines/ti2vid_one_stage.py index bfb801c..40c5d60 100644 --- a/packages/ltx-pipelines/src/ltx_pipelines/ti2vid_one_stage.py +++ b/packages/ltx-pipelines/src/ltx_pipelines/ti2vid_one_stage.py @@ -147,7 +147,7 @@ class TI2VidOneStagePipeline: del transformer cleanup_memory() - decoded_video = vae_decode_video(video_state.latent, self.model_ledger.video_decoder()) + decoded_video = vae_decode_video(video_state.latent, self.model_ledger.video_decoder(), generator=generator) decoded_audio = vae_decode_audio( audio_state.latent, self.model_ledger.audio_decoder(), self.model_ledger.vocoder() ) diff --git a/packages/ltx-pipelines/src/ltx_pipelines/ti2vid_two_stages.py b/packages/ltx-pipelines/src/ltx_pipelines/ti2vid_two_stages.py index b835bfe..528096d 100644 --- a/packages/ltx-pipelines/src/ltx_pipelines/ti2vid_two_stages.py +++ b/packages/ltx-pipelines/src/ltx_pipelines/ti2vid_two_stages.py @@ -225,7 +225,9 @@ class TI2VidTwoStagesPipeline: del video_encoder cleanup_memory() - decoded_video = vae_decode_video(video_state.latent, self.stage_2_model_ledger.video_decoder(), tiling_config) + decoded_video = vae_decode_video( + video_state.latent, self.stage_2_model_ledger.video_decoder(), tiling_config, generator + ) decoded_audio = vae_decode_audio( audio_state.latent, self.stage_2_model_ledger.audio_decoder(), self.stage_2_model_ledger.vocoder() ) diff --git a/packages/ltx-trainer/configs/ltx2_av_lora.yaml b/packages/ltx-trainer/configs/ltx2_av_lora.yaml index 8cf917d..4223618 100644 --- a/packages/ltx-trainer/configs/ltx2_av_lora.yaml +++ b/packages/ltx-trainer/configs/ltx2_av_lora.yaml @@ -253,6 +253,10 @@ checkpoints: # Set to -1 to keep all checkpoints keep_last_n: -1 + # Precision to use when saving checkpoint weights + # Options: "bfloat16" (default, smaller files) or "float32" (full precision) + precision: "bfloat16" + # ----------------------------------------------------------------------------- # Flow Matching Configuration # ----------------------------------------------------------------------------- diff --git a/packages/ltx-trainer/configs/ltx2_v2v_ic_lora.yaml b/packages/ltx-trainer/configs/ltx2_v2v_ic_lora.yaml index 166b1ec..7e647ac 100644 --- a/packages/ltx-trainer/configs/ltx2_v2v_ic_lora.yaml +++ b/packages/ltx-trainer/configs/ltx2_v2v_ic_lora.yaml @@ -263,6 +263,10 @@ checkpoints: # Set to -1 to keep all checkpoints keep_last_n: 3 + # Precision to use when saving checkpoint weights + # Options: "bfloat16" (default, smaller files) or "float32" (full precision) + precision: "bfloat16" + # ----------------------------------------------------------------------------- # Flow Matching Configuration # ----------------------------------------------------------------------------- diff --git a/packages/ltx-trainer/docs/configuration-reference.md b/packages/ltx-trainer/docs/configuration-reference.md index 363d0ec..385f753 100644 --- a/packages/ltx-trainer/docs/configuration-reference.md +++ b/packages/ltx-trainer/docs/configuration-reference.md @@ -292,8 +292,9 @@ Model checkpointing configuration. ```yaml checkpoints: - interval: 250 # Steps between checkpoint saves (null = disabled) - keep_last_n: 3 # Number of recent checkpoints to retain + interval: 250 # Steps between checkpoint saves (null = disabled) + keep_last_n: 3 # Number of recent checkpoints to retain + precision: bfloat16 # Precision for saved weights (bfloat16 or float32) ``` **Key parameters:** @@ -302,6 +303,7 @@ checkpoints: |---------------|------------------------------------------------------------------------| | `interval` | Steps between intermediate checkpoint saves (set to `null` to disable) | | `keep_last_n` | Number of most recent checkpoints to keep (-1 = keep all) | +| `precision` | Precision for saved checkpoint weights: `"bfloat16"` (default) or `"float32"` | ### HubConfig diff --git a/packages/ltx-trainer/src/ltx_trainer/config.py b/packages/ltx-trainer/src/ltx_trainer/config.py index 690f59a..b999665 100644 --- a/packages/ltx-trainer/src/ltx_trainer/config.py +++ b/packages/ltx-trainer/src/ltx_trainer/config.py @@ -350,6 +350,11 @@ class CheckpointsConfig(ConfigBaseModel): ge=-1, ) + precision: Literal["bfloat16", "float32"] = Field( + default="bfloat16", + description="Precision to use when saving checkpoint weights. Options: 'bfloat16' or 'float32'.", + ) + class HubConfig(ConfigBaseModel): """Configuration for Hugging Face Hub integration""" diff --git a/packages/ltx-trainer/src/ltx_trainer/trainer.py b/packages/ltx-trainer/src/ltx_trainer/trainer.py index 72e6364..c8eda57 100644 --- a/packages/ltx-trainer/src/ltx_trainer/trainer.py +++ b/packages/ltx-trainer/src/ltx_trainer/trainer.py @@ -873,6 +873,9 @@ class LtxvTrainer: save_dir.mkdir(exist_ok=True, parents=True) + # Determine save precision + save_dtype = torch.bfloat16 if self._config.checkpoints.precision == "bfloat16" else torch.float32 + # For LoRA: extract only adapter weights; for full: use as-is if is_lora: unwrapped = self._accelerator.unwrap_model(self._transformer, keep_torch_compile=False) @@ -885,9 +888,15 @@ class LtxvTrainer: # Convert to ComfyUI-compatible format (add "diffusion_model." prefix) state_dict = {f"diffusion_model.{k}": v for k, v in state_dict.items()} + # Cast to configured precision + state_dict = {k: v.to(save_dtype) if isinstance(v, Tensor) else v for k, v in state_dict.items()} + # Save to disk save_file(state_dict, saved_weights_path) else: + # Cast to configured precision + full_state_dict = {k: v.to(save_dtype) if isinstance(v, Tensor) else v for k, v in full_state_dict.items()} + # Save to disk self._accelerator.save(full_state_dict, saved_weights_path)