Automated PR - 2026-01-12

This commit is contained in:
sync-bot
2026-01-12 14:14:33 +00:00
parent 628956009c
commit e5a15a4777
14 changed files with 66 additions and 31 deletions
@@ -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,
@@ -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)
@@ -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]]: