diff --git a/finetrainers/ltx_video/full_finetune_condition.py b/finetrainers/ltx_video/full_finetune_condition.py index a82ddbd..a6fbd9d 100644 --- a/finetrainers/ltx_video/full_finetune_condition.py +++ b/finetrainers/ltx_video/full_finetune_condition.py @@ -3,21 +3,14 @@ from typing import Dict, List, Optional import torch import torch.nn as nn -from diffusers import FlowMatchEulerDiscreteScheduler, LTXVideoTransformer3DModel +from diffusers import FlowMatchEulerDiscreteScheduler, LTXVideoTransformer3DModel, AutoencoderKLLTXVideo from finetrainers.conditioning.LTXVideoConditionedTransformer3DModel import LTXVideoConditionedTransformer3DModel from finetrainers.conditioning.conditioned_pipeline import LTXConditionedPipeline from PIL import Image -from diffusers import LTXVideoConditionedTransformer3DModel, FlowMatchEulerDiscreteScheduler -from finetrainers.conditioning.LTXVideoConditionedTransformer3DModel import LTXVideoConditionedTransformer3DModel -from finetrainers.conditioning.conditioned_pipeline import LTXConditionedPipeline -from PIL import Image -from transformers import T5Tokenizer, T5EncoderModel,AutoencoderKLLTXVideo - +from transformers import T5Tokenizer, T5EncoderModel # Exisiting Pipeline -from .lora import ( - initialize_pipeline, - load_conditioned_diffusion_models, # use this for conditioning +from finetrainers.ltx_video.lora import ( load_latent_models, post_latent_preparation, prepare_conditions, diff --git a/finetrainers/ltx_video/lora.py b/finetrainers/ltx_video/lora.py index 4b600b8..dac9a38 100644 --- a/finetrainers/ltx_video/lora.py +++ b/finetrainers/ltx_video/lora.py @@ -44,11 +44,12 @@ def load_condition_models( cache_dir: Optional[str] = None, **kwargs, ) -> Dict[str, nn.Module]: - transformer = LTXVideoConditionedTransformer3DModel.from_pretrained( - model_id, subfolder="transformer", torch_dtype=transformer_dtype, revision=revision, cache_dir=cache_dir + tokenizer = T5Tokenizer.from_pretrained(model_id, subfolder="tokenizer", revision=revision, cache_dir=cache_dir) + text_encoder = T5EncoderModel.from_pretrained( + model_id, subfolder="text_encoder", torch_dtype=text_encoder_dtype, revision=revision, cache_dir=cache_dir ) - scheduler = FlowMatchEulerDiscreteScheduler() - return {"transformer": transformer, "scheduler": scheduler} + return {"tokenizer": tokenizer, "text_encoder": text_encoder} + def load_diffusion_models(