Automated PR - 2026-01-12
This commit is contained in:
@@ -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]]:
|
||||
|
||||
Reference in New Issue
Block a user