def __init__(self, *args, **kwargs) -> None:
super().__init__(*args, **kwargs)
self.device = get_local_torch_device()
# Adapter tensors and wrapped model layers belong to this pipeline's module
# instances. Sharing either cache across two generators can apply one model's
# adapter to another model's layers.
self.trainable_transformer_modules = {}
self.lora_adapters = defaultdict(dict)
self.lora_adapter_paths = {}
self.lora_layers = {}
self.exclude_lora_layers = {}
self.cur_adapter_name = ""
self.cur_adapter_path = ""
self.cur_adapter_strength = 1.0
# build list of trainable transformers
for transformer_name in self.trainable_transformer_names:
if (transformer_name in self.modules and self.modules[transformer_name] is not None):
self.trainable_transformer_modules[transformer_name] = (self.modules[transformer_name])
# check for transformer_2 in case of Wan2.2 MoE or fake_score_transformer_2
if transformer_name.endswith("_2"):
raise ValueError(
f"trainable_transformer_name override in pipelines should not include _2 suffix: {transformer_name}"
)
secondary_transformer_name = transformer_name + "_2"
if (secondary_transformer_name in self.modules and self.modules[secondary_transformer_name] is not None):
self.trainable_transformer_modules[secondary_transformer_name] = self.modules[
secondary_transformer_name]
logger.info(
"trainable_transformer_modules: %s",
self.trainable_transformer_modules.keys(),
)
# Only override the pipeline class's own default when the caller actually set
# one. Assigning unconditionally erases per-model defaults, and a model that
# declares one usually does so because wrapping every linear breaks its forward.
if self.fastvideo_args.lora_target_modules is not None:
self.lora_target_modules = self.fastvideo_args.lora_target_modules
self.lora_path = self.fastvideo_args.lora_path
self.lora_nickname = self.fastvideo_args.lora_nickname
self.lora_strength = self.fastvideo_args.lora_strength
constructor_patch = DenseLoRAPatch.from_adapter(self.lora_path) if self.lora_path else None
self._constructor_dense_lora_path = self.lora_path if constructor_patch is not None else None
self.training_mode = self.fastvideo_args.training_mode
if self.training_mode and getattr(self.fastvideo_args, "lora_training", False):
assert isinstance(self.fastvideo_args, TrainingArgs)
if self.fastvideo_args.lora_alpha is None:
self.fastvideo_args.lora_alpha = self.fastvideo_args.lora_rank
self.lora_rank = self.fastvideo_args.lora_rank # type: ignore
self.lora_alpha = self.fastvideo_args.lora_alpha # type: ignore
logger.info(
"Using LoRA training with rank %d and alpha %d",
self.lora_rank,
self.lora_alpha,
)
if self.lora_target_modules is None:
self.lora_target_modules = [
"q_proj",
"k_proj",
"v_proj",
"o_proj",
"to_q",
"to_k",
"to_v",
"to_out",
"to_qkv",
"to_gate_compress",
]
logger.info(
"Using default lora_target_modules for all transformers: %s",
self.lora_target_modules,
)
else:
logger.warning(
"Using custom lora_target_modules for all transformers, which may not be intended: %s",
self.lora_target_modules,
)
self.convert_to_lora_layers()
# Inference
elif not self.training_mode and self.lora_path is not None:
self.convert_to_lora_layers()
if not any(is_lazy_module(module) for module in self.trainable_transformer_modules.values()):
self._setting_constructor_adapter = True
try:
self.set_lora_adapter(
self.lora_nickname, # type: ignore
self.lora_path,
strength=self.lora_strength,
) # type: ignore
finally:
self._setting_constructor_adapter = False