Automated PR - 2026-03-05
This commit is contained in:
@@ -27,14 +27,15 @@ from torch.optim.lr_scheduler import (
|
||||
from torch.utils.data import DataLoader
|
||||
from torchvision.transforms import functional as F # noqa: N812
|
||||
|
||||
from ltx_core.text_encoders.gemma import convert_to_additive_mask
|
||||
from ltx_trainer import logger
|
||||
from ltx_trainer.config import LtxTrainerConfig
|
||||
from ltx_trainer.config_display import print_config
|
||||
from ltx_trainer.datasets import PrecomputedDataset
|
||||
from ltx_trainer.gpu_utils import free_gpu_memory, free_gpu_memory_context, get_gpu_memory_gb
|
||||
from ltx_trainer.hf_hub_utils import push_to_hub
|
||||
from ltx_trainer.model_loader import load_embeddings_processor, load_text_encoder
|
||||
from ltx_trainer.model_loader import load_model as load_ltx_model
|
||||
from ltx_trainer.model_loader import load_text_encoder
|
||||
from ltx_trainer.progress import TrainingProgress
|
||||
from ltx_trainer.quantization import quantize_model
|
||||
from ltx_trainer.timestep_samplers import SAMPLERS
|
||||
@@ -320,8 +321,8 @@ class LtxvTrainer:
|
||||
audio_features = conditions["prompt_embeds"]
|
||||
|
||||
mask = conditions["prompt_attention_mask"]
|
||||
additive_mask = self._text_encoder._convert_to_additive_mask(mask, video_features.dtype)
|
||||
video_embeds, audio_embeds, attention_mask = self._text_encoder.embeddings_processor.create_embeddings(
|
||||
additive_mask = convert_to_additive_mask(mask, video_features.dtype)
|
||||
video_embeds, audio_embeds, attention_mask = self._embeddings_processor.create_embeddings(
|
||||
video_features, audio_features, additive_mask
|
||||
)
|
||||
|
||||
@@ -346,26 +347,31 @@ class LtxvTrainer:
|
||||
|
||||
@free_gpu_memory_context(after=True)
|
||||
def _load_text_encoder_and_cache_embeddings(self) -> list[CachedPromptEmbeddings] | None:
|
||||
"""Load text encoder, computes and returns validation embeddings."""
|
||||
"""Load text encoder + embeddings processor, compute and cache validation embeddings."""
|
||||
|
||||
# This method:
|
||||
# 1. Loads the text encoder on GPU
|
||||
# 2. If validation prompts are configured, computes and caches their embeddings
|
||||
# 3. Unloads the heavy Gemma model while keeping the lightweight embedding connectors
|
||||
# The text encoder is kept (as self._text_encoder) but with model/tokenizer/feature_extractor
|
||||
# set to None. Only the embedding connectors remain for use during training.
|
||||
# 1. Loads the pure Gemma text encoder on GPU
|
||||
# 2. Loads the embeddings processor (feature extractor + connectors)
|
||||
# 3. If validation prompts are configured, computes and caches their embeddings
|
||||
# 4. Unloads the Gemma model entirely, keeps the embeddings processor for training
|
||||
|
||||
# Load text encoder on GPU
|
||||
# Load text encoder (pure Gemma LLM) on GPU
|
||||
logger.debug("Loading text encoder...")
|
||||
|
||||
self._text_encoder = load_text_encoder(
|
||||
checkpoint_path=self._config.model.model_path,
|
||||
text_encoder = load_text_encoder(
|
||||
gemma_model_path=self._config.model.text_encoder_path,
|
||||
device="cuda",
|
||||
dtype=torch.bfloat16,
|
||||
load_in_8bit=self._config.acceleration.load_text_encoder_in_8bit,
|
||||
)
|
||||
|
||||
# Load embeddings processor (feature extractor + connectors)
|
||||
logger.debug("Loading embeddings processor...")
|
||||
self._embeddings_processor = load_embeddings_processor(
|
||||
checkpoint_path=self._config.model.model_path,
|
||||
device="cuda",
|
||||
dtype=torch.bfloat16,
|
||||
)
|
||||
|
||||
# Cache validation embeddings if prompts are configured
|
||||
cached_embeddings = None
|
||||
if self._config.validation.prompts:
|
||||
@@ -373,22 +379,26 @@ class LtxvTrainer:
|
||||
cached_embeddings = []
|
||||
with torch.inference_mode():
|
||||
for prompt in self._config.validation.prompts:
|
||||
v_ctx_pos, a_ctx_pos, _ = self._text_encoder(prompt)
|
||||
v_ctx_neg, a_ctx_neg, _ = self._text_encoder(self._config.validation.negative_prompt)
|
||||
pos_hs, pos_mask = text_encoder.encode(prompt)
|
||||
pos_out = self._embeddings_processor.process_hidden_states(pos_hs, pos_mask)
|
||||
|
||||
neg_hs, neg_mask = text_encoder.encode(self._config.validation.negative_prompt)
|
||||
neg_out = self._embeddings_processor.process_hidden_states(neg_hs, neg_mask)
|
||||
|
||||
cached_embeddings.append(
|
||||
CachedPromptEmbeddings(
|
||||
video_context_positive=v_ctx_pos.cpu(),
|
||||
audio_context_positive=a_ctx_pos.cpu(),
|
||||
video_context_negative=v_ctx_neg.cpu() if v_ctx_neg is not None else None,
|
||||
audio_context_negative=a_ctx_neg.cpu() if a_ctx_neg is not None else None,
|
||||
video_context_positive=pos_out.video_encoding.cpu(),
|
||||
audio_context_positive=pos_out.audio_encoding.cpu(),
|
||||
video_context_negative=neg_out.video_encoding.cpu(),
|
||||
audio_context_negative=(
|
||||
neg_out.audio_encoding.cpu() if neg_out.audio_encoding is not None else None
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
# Unload heavy components to free VRAM, keeping only the embedding connectors
|
||||
self._text_encoder.model = None
|
||||
self._text_encoder.tokenizer = None
|
||||
self._text_encoder.feature_extractor = None
|
||||
# Unload Gemma model and feature extractor, keep only connectors for training
|
||||
del text_encoder
|
||||
self._embeddings_processor.feature_extractor = None
|
||||
|
||||
logger.debug("Validation prompt embeddings cached. Gemma model unloaded")
|
||||
return cached_embeddings
|
||||
@@ -426,7 +436,7 @@ class LtxvTrainer:
|
||||
self._scheduler = components.scheduler
|
||||
self._audio_vae = components.audio_vae_decoder
|
||||
self._vocoder = components.vocoder
|
||||
# Note: self._text_encoder was set in _load_text_encoder_and_cache_embeddings
|
||||
# Note: self._embeddings_processor was set in _load_text_encoder_and_cache_embeddings
|
||||
|
||||
# Determine initial dtype based on training mode.
|
||||
# Note: For FSDP + LoRA, we'll cast to FP32 later in _prepare_models_for_training()
|
||||
|
||||
Reference in New Issue
Block a user