Qwen-Image-Layered-MRP-MLX/src/mflux/dreambooth/lora_layers/lora_layers.py

241 lines
9.8 KiB
Python

from pathlib import Path
from typing import TYPE_CHECKING
import mlx
import mlx.core as mx
from mlx import nn
from mlx.utils import tree_flatten
from mflux.dreambooth.lora_layers.linear_lora_layer import LoRALinear
from mflux.dreambooth.state.training_spec import SingleTransformerBlocks, TrainingSpec, TransformerBlocks
from mflux.dreambooth.state.zip_util import ZipUtil
from mflux.models.transformer.joint_transformer_block import JointTransformerBlock
from mflux.models.transformer.single_transformer_block import SingleTransformerBlock
from mflux.post_processing.generated_image import GeneratedImage
from mflux.weights.weight_handler import MetaData, WeightHandler
if TYPE_CHECKING:
from mflux import Flux1
class LoRALayers:
def __init__(self, weights: "WeightHandler"):
self.layers = weights
@staticmethod
def from_spec(flux: "Flux1", training_spec: TrainingSpec) -> "LoRALayers":
if training_spec.lora_layers.state_path is not None:
# Load from state if present in the spec
from mflux.weights.weight_handler_lora import WeightHandlerLoRA
weights = ZipUtil.unzip(
zip_path=training_spec.checkpoint_path,
filename=training_spec.lora_layers.state_path,
loader=lambda x: WeightHandlerLoRA.load_lora_weights(
transformer=flux.transformer, lora_files=[x], lora_scales=[1.0]
),
)
return LoRALayers(weights=weights[0])
else:
# Construct the LoRA weights from the spec
transformer_lora_layers = {}
single_transformer_lora_layers = {}
if training_spec.lora_layers.transformer_blocks:
transformer_lora_layers = LoRALayers._construct_layers(
blocks=flux.transformer.transformer_blocks,
block_spec=training_spec.lora_layers.transformer_blocks,
block_prefix="transformer.transformer_blocks",
)
if training_spec.lora_layers.single_transformer_blocks:
single_transformer_lora_layers = LoRALayers._construct_layers(
blocks=flux.transformer.single_transformer_blocks,
block_spec=training_spec.lora_layers.single_transformer_blocks,
block_prefix="transformer.single_transformer_blocks",
)
lora_layers = {**transformer_lora_layers, **single_transformer_lora_layers}
weights = WeightHandler(
meta_data=MetaData(is_mflux=True),
transformer=mlx.utils.tree_unflatten(list(lora_layers.items()))['transformer'],
) # fmt:off
return LoRALayers(weights=weights)
@staticmethod
def _construct_layers(
block_spec: TransformerBlocks | SingleTransformerBlocks,
blocks: list[JointTransformerBlock] | list[SingleTransformerBlock],
block_prefix: str,
) -> dict:
block_indices = block_spec.block_range.get_blocks()
lora_layers = {}
for idx in block_indices:
if idx >= len(blocks):
raise IndexError(f"Index {idx} over range")
block = blocks[idx]
for layer_type in block_spec.layer_types:
original_layer = LoRALayers._get_nested_attr(block, layer_type)
is_list = isinstance(original_layer, list)
lora_layer = LoRALinear.from_linear(
linear=original_layer[0] if is_list else original_layer,
r=block_spec.lora_rank,
)
layer_path = f"{block_prefix}.{idx}.{layer_type}"
lora_layers[layer_path] = [lora_layer] if is_list else lora_layer
return lora_layers
@staticmethod
def transformer_dict_from_template(weights: dict, transformer: nn.Module, scale: float) -> dict:
lora_layers = {}
for key in weights.keys():
if key.endswith(".lora_A"):
base_path = key[: -len(".lora_A")]
parts = base_path.split(".")
if parts[1] == "transformer_blocks":
LoRALayers._handle_transformer_blocks(
weights=weights,
scale=scale,
transformer=transformer,
lora_layers=lora_layers,
base_path=base_path,
)
if parts[1] == "single_transformer_blocks":
LoRALayers._handle_single_transformer_blocks(
weights=weights,
scale=scale,
transformer=transformer,
lora_layers=lora_layers,
base_path=base_path,
)
return lora_layers
@staticmethod
def _handle_transformer_blocks(
weights: dict, scale: float, transformer: nn.Module, lora_layers: dict, base_path: str
):
parts = base_path.split(".")
block_idx = int(parts[2])
module_name = parts[3]
attr_name = parts[4]
block = transformer.transformer_blocks[block_idx]
module = getattr(block, module_name)
original_layer = getattr(module, attr_name)
if len(parts) == 6:
original_layer = original_layer[0] # Special case here
# Create LoRA layer
lora_A = weights[f"{base_path}.lora_A"]
rank = lora_A.shape[1]
lora_layer = LoRALinear.from_linear(linear=original_layer, r=rank, scale=scale)
# Set the weights
lora_layer.lora_A = weights[f"{base_path}.lora_A"]
lora_layer.lora_B = weights[f"{base_path}.lora_B"]
# Store the layer
lora_layers[base_path] = lora_layer
@staticmethod
def _handle_single_transformer_blocks(
weights: dict, scale: float, transformer: nn.Module, lora_layers: dict, base_path: str
):
parts = base_path.split(".")
block_idx = int(parts[2])
module_name = parts[3]
if len(parts) == 4:
original_layer = getattr(transformer.single_transformer_blocks[block_idx], module_name)
elif len(parts) == 5:
attr_name = parts[4]
block = transformer.single_transformer_blocks[block_idx]
module = getattr(block, module_name)
original_layer = getattr(module, attr_name)
# Create LoRA layer
lora_A = weights[f"{base_path}.lora_A"]
rank = lora_A.shape[1]
lora_layer = LoRALinear.from_linear(linear=original_layer, r=rank, scale=scale)
# Set the weights
lora_layer.lora_A = weights[f"{base_path}.lora_A"]
lora_layer.lora_B = weights[f"{base_path}.lora_B"]
# Store the layer
lora_layers[base_path] = lora_layer
@staticmethod
def set_transformer_block(transformer_block, dictionary: dict):
for key, val in dictionary.items():
if key == "attn":
LoRALayers._set_attribute(transformer_block, key, val, "to_q")
LoRALayers._set_attribute(transformer_block, key, val, "to_k")
LoRALayers._set_attribute(transformer_block, key, val, "to_v")
LoRALayers._set_attribute(transformer_block, key, val, "to_out")
LoRALayers._set_attribute(transformer_block, key, val, "add_q_proj")
LoRALayers._set_attribute(transformer_block, key, val, "add_k_proj")
LoRALayers._set_attribute(transformer_block, key, val, "add_v_proj")
LoRALayers._set_attribute(transformer_block, key, val, "to_add_out")
elif key == "ff" or key == "ff_context":
LoRALayers._set_attribute(transformer_block, key, val, "linear1")
LoRALayers._set_attribute(transformer_block, key, val, "linear2")
elif key == "norm1" or key == "norm1_context":
LoRALayers._set_attribute(transformer_block, key, val, "linear")
else:
raise Exception("Could not set LoRA weights")
@staticmethod
def set_single_transformer_block(single_transformer_block, dictionary: dict):
for key, val in dictionary.items():
if key == "attn":
LoRALayers._set_attribute(single_transformer_block, key, val, "to_q")
LoRALayers._set_attribute(single_transformer_block, key, val, "to_k")
LoRALayers._set_attribute(single_transformer_block, key, val, "to_v")
elif key == "norm":
LoRALayers._set_attribute(single_transformer_block, key, val, "linear")
elif key == "proj_mlp" or key == "proj_out":
single_transformer_block[key] = val
else:
raise Exception("Could not set LoRA weights")
@staticmethod
def _set_attribute(block, key: str, val: dict, name: str):
if block[key].get(name, False) and val.get(name, False):
block[key][name] = val[name]
@staticmethod
def _get_nested_attr(obj, attr_path):
attrs = attr_path.split(".")
for attr in attrs:
obj = getattr(obj, attr)
return obj
def save(self, path: Path, training_spec: TrainingSpec) -> None:
weights = {}
for entry in tree_flatten(self.layers.transformer):
name = entry[0]
weight = entry[1]
if name.endswith(".lora_A") or name.endswith(".lora_B"):
weights[name] = weight
weights = {key: mx.transpose(val) for key, val in weights.items()}
weights = {"transformer": weights}
mx.save_safetensors(
str(path),
dict(tree_flatten(weights)),
metadata={
"mflux_version": GeneratedImage.get_version(),
"transformer_blocks": str(training_spec.lora_layers.transformer_blocks),
"single_transformer_blocks": str(training_spec.lora_layers.single_transformer_blocks),
},
)