-
Notifications
You must be signed in to change notification settings - Fork 442
Expand file tree
/
Copy pathlora_pipeline.py
More file actions
457 lines (414 loc) · 20.8 KB
/
Copy pathlora_pipeline.py
File metadata and controls
457 lines (414 loc) · 20.8 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
# SPDX-License-Identifier: Apache-2.0
from collections import defaultdict
from collections.abc import Hashable
from contextlib import nullcontext
import re
from typing import Any
from collections.abc import Generator
import torch
import torch.distributed as dist
import torch.nn as nn
from safetensors.torch import load_file
from torch.distributed.device_mesh import DeviceMesh, init_device_mesh
from torch.distributed.tensor import DTensor
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.hooks.hooks import ModuleHookManager
from fastvideo.hooks.layerwise_offload import LayerwiseOffloadHook
from fastvideo.layers.lora.linear import (
BaseLayerWithLoRA,
get_lora_layer,
replace_submodule,
)
from fastvideo.logger import init_logger
from fastvideo.models.loader.utils import get_param_names_mapping
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.utils import maybe_download_lora
logger = init_logger(__name__)
def _normalize_lora_key(name: str) -> str:
"""Normalize known DiffSynth/PEFT wrappers without rewriting module names."""
name = name.removeprefix("diffusion_model.")
name = name.removeprefix("pipe.dit.")
name = re.sub(r"(\.lora_[AB])\.default(?=\.weight$)", r"\1", name)
return name.removesuffix(".weight")
def _get_hook_ctx(module: nn.Module | None):
if module is None:
return nullcontext()
hook_mgr = ModuleHookManager.get_from(module)
if hook_mgr is not None:
offload_hook = hook_mgr.forward_hooks.get(LayerwiseOffloadHook.name())
if offload_hook is not None:
return offload_hook.mutate_params_scope() # type: ignore
return nullcontext()
def _named_module_by_prefix(module: nn.Module,
prefixes: list[str]) -> list[tuple[str | None, list[tuple[str, nn.Module]]]]:
none_list: list[tuple[str, nn.Module]] = []
prefix_list: list[tuple[str, list[tuple[str, nn.Module]]]] = [(prefix, []) for prefix in prefixes]
for name, submodule in module.named_modules():
for cur_prefix, cur_list in prefix_list:
# we should exclude e.g. block.1 and block.12.attn
if name.startswith(cur_prefix + "."):
cur_list.append((name, submodule))
break
else:
none_list.append((name, submodule))
return prefix_list + [(None, none_list)] # type: ignore
class LoRAModelLayers:
def __init__(self, block_list: list[tuple[str, nn.Module]]) -> None:
# block_name -> {layer_name -> layer}
self.block_to_lora_layers: dict[str, dict[str, BaseLayerWithLoRA]] = {}
# layer_name -> block_name
self.lora_layers_to_block: dict[str, str | None] = {}
self.other_lora_layers: dict[str, BaseLayerWithLoRA] = {}
self.block_mapping = dict(block_list)
def add_lora_layer(self, block_name: str | None, layer_name: str, layer: BaseLayerWithLoRA):
if block_name is None:
self.other_lora_layers[layer_name] = layer
self.lora_layers_to_block[layer_name] = None
else:
if block_name not in self.block_to_lora_layers:
self.block_to_lora_layers[block_name] = {}
self.block_to_lora_layers[block_name][layer_name] = layer
self.lora_layers_to_block[layer_name] = block_name
def all_lora_layers(self, ) -> Generator[tuple[str, BaseLayerWithLoRA], Any, None]:
for block_layers in self.block_to_lora_layers.values():
for name, layer in block_layers.items():
yield name, layer
for name, layer in self.other_lora_layers.items():
yield name, layer
def lora_layers_by_block(self, ) -> Generator[
tuple[nn.Module | None, dict[str, BaseLayerWithLoRA]],
Any,
None,
]:
for block_name, layers in self.block_to_lora_layers.items():
yield self.block_mapping[block_name], layers
yield None, self.other_lora_layers
class LoRAPipeline(ComposedPipelineBase):
"""
Pipeline that supports injecting LoRA adapters into the diffusion transformer.
TODO: support training.
"""
lora_adapters: dict[str, dict[str, torch.Tensor]] = defaultdict(
dict) # state dicts of loaded lora adapters (includes lora_A, lora_B, and lora_alpha)
cur_adapter_name: str = ""
cur_adapter_path: str = ""
cur_adapter_strength: float = 1.0
lora_adapter_paths: dict[str, str] = {}
# model_name -> layers
lora_layers: dict[str, LoRAModelLayers] = {}
fastvideo_args: FastVideoArgs | TrainingArgs
exclude_lora_layers: dict[str, list[str]] = {}
device: torch.device = get_local_torch_device()
lora_target_modules: list[str] | None = None
lora_path: str | None = None
lora_nickname: str = "default"
lora_rank: int | None = None
lora_alpha: int | None = None
lora_initialized: bool = False
def __init__(self, *args, **kwargs) -> None:
super().__init__(*args, **kwargs)
self.device = get_local_torch_device()
self.lora_adapter_paths = {}
# 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(),
)
for (
transformer_name,
transformer_module,
) in self.trainable_transformer_modules.items():
self.exclude_lora_layers[transformer_name] = (transformer_module.config.arch_config.exclude_lora_layers)
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.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()
self.set_lora_adapter(
self.lora_nickname, # type: ignore
self.lora_path,
) # type: ignore
def is_target_layer(self, module_name: str) -> bool:
if self.lora_target_modules is None:
return True
return any(target_name in module_name for target_name in self.lora_target_modules)
def set_trainable(self) -> None:
def set_lora_grads(lora_layers: LoRAModelLayers, device_mesh: DeviceMesh):
for name, layer in lora_layers.all_lora_layers():
layer.lora_A.requires_grad_(True)
layer.lora_B.requires_grad_(True)
layer.base_layer.requires_grad_(False)
layer.lora_A = nn.Parameter(DTensor.from_local(layer.lora_A, device_mesh=device_mesh))
layer.lora_B = nn.Parameter(DTensor.from_local(layer.lora_B, device_mesh=device_mesh))
is_lora_training = self.training_mode and getattr(self.fastvideo_args, "lora_training", False)
if not is_lora_training:
super().set_trainable()
return
device_mesh = init_device_mesh(
"cuda",
(dist.get_world_size(), 1),
mesh_dim_names=["fake", "replicate"],
)
for (
transformer_name,
transformer_module,
) in self.trainable_transformer_modules.items():
transformer_module.train()
transformer_module.requires_grad_(False)
if transformer_name in self.lora_layers:
set_lora_grads(self.lora_layers[transformer_name], device_mesh)
else:
raise ValueError(f"Transformer {transformer_name} should be trainable but not found in lora_layers")
def convert_to_lora_layers(self) -> None:
"""
Unified method to convert the transformer to a LoRA transformer.
"""
if self.lora_initialized:
return
self.lora_initialized = True
for (
transformer_name,
transformer_module,
) in self.trainable_transformer_modules.items():
converted_count = 0
# init bookkeeping structures
if transformer_name not in self.lora_layers:
# get block list
block_list = []
for name, submodule in transformer_module.named_children():
if isinstance(submodule, nn.ModuleList):
block_list = [(f"{name}.{i}", m) for i, m in enumerate(submodule)]
break
self.lora_layers[transformer_name] = LoRAModelLayers(block_list)
logger.info("Converting %s to LoRA Transformer", transformer_name)
# scan every module and convert to LoRA layer if applicable
for block_name, block_modules in _named_module_by_prefix(
transformer_module,
list(self.lora_layers[transformer_name].block_mapping),
):
if block_name is not None and (not self.fastvideo_args.training_mode
and self.fastvideo_args.dit_layerwise_offload):
scope_ctx = _get_hook_ctx(self.lora_layers[transformer_name].block_mapping[block_name])
else:
scope_ctx = nullcontext()
with scope_ctx:
for name, layer in block_modules:
if not self.is_target_layer(name):
continue
excluded = False
for exclude_layer in self.exclude_lora_layers[transformer_name]:
if exclude_layer in name:
excluded = True
break
if excluded:
continue
layer = get_lora_layer(
layer,
lora_rank=self.lora_rank,
lora_alpha=self.lora_alpha,
training_mode=self.training_mode,
)
if layer is not None:
block_name_split = name.split(".", 2)
if len(block_name_split) > 2:
block_name = (block_name_split[0] + "." + block_name_split[1])
else:
block_name = None
if (block_name not in self.lora_layers[transformer_name].block_mapping):
block_name = None
self.lora_layers[transformer_name].add_lora_layer(block_name, name, layer)
replace_submodule(transformer_module, name, layer)
converted_count += 1
logger.info("Converted %d layers to LoRA layers", converted_count)
def set_lora_adapter(self,
lora_nickname: str,
lora_path: str | None = None,
strength: float = 1.0,
accumulate: bool = False): # type: ignore
"""
Load a LoRA adapter into the pipeline and merge it into the transformer.
Args:
lora_nickname: The "nick name" of the adapter when referenced in the pipeline.
lora_path: The path to the adapter, either a local path or a Hugging Face repo id.
"""
if lora_nickname not in self.lora_adapters and lora_path is None:
raise ValueError(f"Adapter {lora_nickname} not found in the pipeline. Please provide lora_path to load it.")
if not self.lora_initialized:
self.convert_to_lora_layers()
adapter_updated = False
rank = dist.get_rank()
if lora_path is not None and self.lora_adapter_paths.get(lora_nickname) != lora_path:
self.lora_adapters[lora_nickname] = {}
lora_local_path = maybe_download_lora(lora_path)
lora_state_dict = load_file(lora_local_path)
# Map the hf layer names to our custom layer names
param_names_mapping_fn = get_param_names_mapping(self.modules["transformer"].param_names_mapping)
lora_param_names_mapping_fn = get_param_names_mapping(self.modules["transformer"].lora_param_names_mapping)
# Extract alpha values and weights in a single pass
to_merge_params: defaultdict[Hashable, dict[Any, Any]] = (defaultdict(dict))
for name, weight in lora_state_dict.items():
# Extract weights (lora_A, lora_B, and lora_alpha)
name = _normalize_lora_key(name)
if "lora_alpha" in name:
# Store alpha with minimal mapping - same processing as lora_A/lora_B
# but store in lora_adapters with ".lora_alpha" suffix
layer_name = name.replace(".lora_alpha", "")
layer_name, _, _ = lora_param_names_mapping_fn(layer_name)
target_name, _, _ = param_names_mapping_fn(layer_name)
# Store alpha alongside weights with same target_name base
alpha_key = target_name + ".lora_alpha"
self.lora_adapters[lora_nickname][alpha_key] = (weight.item()
if weight.numel() == 1 else float(weight.mean()))
continue
name, _, _ = lora_param_names_mapping_fn(name)
target_name, merge_index, num_params_to_merge = (param_names_mapping_fn(name))
# for (in_dim, r) @ (r, out_dim), we only merge (r, out_dim * n) where n is the number of linear layers to fuse
# see param mapping in HunyuanVideoArchConfig
if merge_index is not None and "lora_B" in name:
to_merge_params[target_name][merge_index] = weight
if len(to_merge_params[target_name]) == num_params_to_merge:
# cat at output dim according to the merge_index order
sorted_tensors = [to_merge_params[target_name][i] for i in range(num_params_to_merge)]
weight = torch.cat(sorted_tensors, dim=1)
del to_merge_params[target_name]
else:
continue
if target_name in self.lora_adapters[lora_nickname]:
raise ValueError(f"Target name {target_name} already exists in lora_adapters[{lora_nickname}]")
self.lora_adapters[lora_nickname][target_name] = weight.to(self.device)
adapter_updated = True
self.cur_adapter_path = lora_path
self.lora_adapter_paths[lora_nickname] = lora_path
logger.info("Rank %d: loaded LoRA adapter %s", rank, lora_path)
if (not adapter_updated and self.cur_adapter_name == lora_nickname and self.cur_adapter_strength == strength
and not accumulate):
return
self.cur_adapter_name = lora_nickname
self.cur_adapter_strength = strength
# Merge the new adapter
adapted_count = 0
for (
transformer_name,
transformer_lora_layers,
) in self.lora_layers.items():
for (
module,
layers,
) in transformer_lora_layers.lora_layers_by_block():
with _get_hook_ctx(module):
for name, layer in layers.items():
lora_A_name = name + ".lora_A"
lora_B_name = name + ".lora_B"
lora_alpha_name = name + ".lora_alpha"
if (lora_A_name in self.lora_adapters[lora_nickname]
and lora_B_name in self.lora_adapters[lora_nickname]):
# Get alpha value for this layer (defaults to None if not present)
lora_A = self.lora_adapters[lora_nickname][lora_A_name]
lora_B = self.lora_adapters[lora_nickname][lora_B_name]
# Simple lookup - alpha stored with same naming scheme as lora_A/lora_B
alpha = self.lora_adapters[lora_nickname].get(lora_alpha_name)
try:
layer.set_lora_weights(
lora_A,
lora_B,
lora_alpha=alpha,
training_mode=self.fastvideo_args.training_mode,
lora_path=lora_path,
strength=strength,
accumulate=accumulate,
)
except Exception as e:
logger.error(
"Error setting LoRA weights for layer %s: %s",
name,
str(e),
)
raise e
adapted_count += 1
else:
if rank == 0:
logger.warning(
"LoRA adapter %s does not contain the weights for layer %s. LoRA will not be applied to it.",
lora_path,
name,
)
layer.disable_lora = True
logger.info(
"Rank %d: LoRA adapter %s applied to %d layers",
rank,
lora_path,
adapted_count,
)
def merge_lora_weights(self) -> None:
for (
transformer_name,
transformer_lora_layers,
) in self.lora_layers.items():
for (
module,
layers,
) in transformer_lora_layers.lora_layers_by_block():
with _get_hook_ctx(module):
for name, layer in layers.items():
layer.merge_lora_weights()
def unmerge_lora_weights(self) -> None:
for (
transformer_name,
transformer_lora_layers,
) in self.lora_layers.items():
for (
module,
layers,
) in transformer_lora_layers.lora_layers_by_block():
with _get_hook_ctx(module):
for name, layer in layers.items():
layer.unmerge_lora_weights()