Fixed some issues with imports and the original lora.py code.

This commit is contained in:
CrossProduct
2025-01-25 02:22:05 +00:00
parent 06e26ec092
commit e99d70ff46
2 changed files with 8 additions and 14 deletions
@@ -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,
+5 -4
View File
@@ -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(