[diffusion] fix: fix flux2 lora (#20200)

Signed-off-by: Lancer <maruixiang6688@gmail.com>
This commit is contained in:
Lancer
2026-03-10 16:57:01 +08:00
committed by GitHub
parent 8517da5d08
commit 2c2003158f
3 changed files with 203 additions and 1 deletions

View File

@@ -23,6 +23,15 @@ class FluxArchConfig(DiTArchConfig):
stacked_params_mapping: list[tuple[str, str, str]] = field(default_factory=list)
exclude_lora_layers: list[str] = field(
default_factory=lambda: [
"time_guidance_embed.timestep_embedder.linear_1",
"time_guidance_embed.timestep_embedder.linear_2",
"time_guidance_embed.guidance_embedder.linear_1",
"time_guidance_embed.guidance_embedder.linear_2",
]
)
# nunchaku checkpoint uses different weight names; map to sglang flux layout
param_names_mapping: dict = field(
default_factory=lambda: {

View File

@@ -19,6 +19,7 @@ class LoRAFormat(str, Enum):
XLABS_FLUX = "xlabs-ai"
KOHYA_FLUX = "kohya-flux"
WAN = "wan"
AI_TOOLKIT_FLUX = "ai-toolkit-flux"
def _sample_keys(keys: Iterable[str], k: int = 20) -> list[str]:
@@ -104,6 +105,32 @@ def _looks_like_qwen_image(state_dict: Mapping[str, torch.Tensor]) -> bool:
)
def _looks_like_ai_toolkit_flux_lora(state_dict: Mapping[str, torch.Tensor]) -> bool:
"""Detect ai-toolkit/ComfyUI trained Flux LoRA with double_blocks/single_blocks naming.
Key patterns: double_blocks.{N}.img_attn.proj.lora_A.weight
"""
keys = list(state_dict.keys())
if not keys:
return False
has_double_blocks = any(
k.startswith("double_blocks.")
or k.startswith("base_model.model.double_blocks.")
for k in keys
)
has_single_blocks = any(
k.startswith("single_blocks.")
or k.startswith("base_model.model.single_blocks.")
for k in keys
)
has_lora_ab = _has_substring_key(keys, ".lora_A") or _has_substring_key(
keys, ".lora_B"
)
return (has_double_blocks or has_single_blocks) and has_lora_ab
def detect_lora_format_from_state_dict(
state_dict: Mapping[str, torch.Tensor],
) -> LoRAFormat:
@@ -112,6 +139,9 @@ def detect_lora_format_from_state_dict(
if not keys:
return LoRAFormat.STANDARD
if _looks_like_ai_toolkit_flux_lora(state_dict):
return LoRAFormat.AI_TOOLKIT_FLUX
if _has_substring_key(keys, ".lora_A") or _has_substring_key(keys, ".lora_B"):
return LoRAFormat.STANDARD
@@ -155,7 +185,6 @@ def _convert_qwen_image_standard(
out[new_name] = tensor
sample = _sample_keys(out.keys(), 20)
return out
@@ -309,6 +338,155 @@ def _convert_kohya_flux_via_diffusers(
)
def _convert_ai_toolkit_flux_lora(
state_dict: Mapping[str, torch.Tensor],
log: logging.Logger,
) -> Dict[str, torch.Tensor]:
"""Convert ai-toolkit/ComfyUI trained Flux LoRA to SGLang format.
Handles the naming convention conversion:
- double_blocks.{N}.img_attn.qkv -> transformer_blocks.{N}.attn.to_q/k/v
- double_blocks.{N}.txt_attn.qkv -> transformer_blocks.{N}.attn.add_q/k/v_proj
- double_blocks.{N}.img_attn.proj -> transformer_blocks.{N}.attn.to_out.0
- double_blocks.{N}.txt_attn.proj -> transformer_blocks.{N}.attn.to_add_out
- double_blocks -> transformer_blocks
- single_blocks -> single_transformer_blocks
"""
out: Dict[str, torch.Tensor] = {}
original_state_dict: Dict[str, torch.Tensor] = {}
for name, tensor in state_dict.items():
new_name = name
if new_name.startswith("diffusion_model."):
new_name = new_name[len("diffusion_model.") :]
if new_name.startswith("base_model.model."):
new_name = new_name[len("base_model.model.") :]
original_state_dict[new_name] = tensor
num_double_layers = 0
num_single_layers = 0
for key in original_state_dict.keys():
if key.startswith("single_blocks."):
parts = key.split(".")
if len(parts) > 1 and parts[1].isdigit():
num_single_layers = max(num_single_layers, int(parts[1]) + 1)
elif key.startswith("double_blocks."):
parts = key.split(".")
if len(parts) > 1 and parts[1].isdigit():
num_double_layers = max(num_double_layers, int(parts[1]) + 1)
lora_keys = ("lora_A", "lora_B")
attn_types = ("img_attn", "txt_attn")
for sl in range(num_single_layers):
single_block_prefix = f"single_blocks.{sl}"
attn_prefix = f"single_transformer_blocks.{sl}.attn"
for lora_key in lora_keys:
linear1_key = f"{single_block_prefix}.linear1.{lora_key}.weight"
if linear1_key in original_state_dict:
out[f"{attn_prefix}.to_qkv_mlp_proj.{lora_key}.weight"] = (
original_state_dict.pop(linear1_key)
)
linear2_key = f"{single_block_prefix}.linear2.{lora_key}.weight"
if linear2_key in original_state_dict:
out[f"{attn_prefix}.to_out.{lora_key}.weight"] = (
original_state_dict.pop(linear2_key)
)
for dl in range(num_double_layers):
transformer_block_prefix = f"transformer_blocks.{dl}"
for lora_key in lora_keys:
for attn_type in attn_types:
attn_prefix = f"{transformer_block_prefix}.attn"
qkv_key = f"double_blocks.{dl}.{attn_type}.qkv.{lora_key}.weight"
if qkv_key not in original_state_dict:
continue
fused_qkv_weight = original_state_dict.pop(qkv_key)
if lora_key == "lora_A":
diff_attn_proj_keys = (
["to_q", "to_k", "to_v"]
if attn_type == "img_attn"
else ["add_q_proj", "add_k_proj", "add_v_proj"]
)
for proj_key in diff_attn_proj_keys:
out[f"{attn_prefix}.{proj_key}.{lora_key}.weight"] = (
fused_qkv_weight
)
else:
if fused_qkv_weight.shape[0] % 3 != 0:
log.warning(
"[LoRAFormatAdapter] QKV weight shape %s not divisible by 3, "
"may cause shape mismatch for %s",
fused_qkv_weight.shape,
qkv_key,
)
sample_q, sample_k, sample_v = torch.chunk(
fused_qkv_weight, 3, dim=0
)
if attn_type == "img_attn":
out[f"{attn_prefix}.to_q.{lora_key}.weight"] = sample_q
out[f"{attn_prefix}.to_k.{lora_key}.weight"] = sample_k
out[f"{attn_prefix}.to_v.{lora_key}.weight"] = sample_v
else:
out[f"{attn_prefix}.add_q_proj.{lora_key}.weight"] = sample_q
out[f"{attn_prefix}.add_k_proj.{lora_key}.weight"] = sample_k
out[f"{attn_prefix}.add_v_proj.{lora_key}.weight"] = sample_v
proj_mappings = [
("img_attn.proj", "attn.to_out.0"),
("txt_attn.proj", "attn.to_add_out"),
]
for org_proj, diff_proj in proj_mappings:
for lora_key in lora_keys:
original_key = f"double_blocks.{dl}.{org_proj}.{lora_key}.weight"
if original_key in original_state_dict:
diffusers_key = (
f"{transformer_block_prefix}.{diff_proj}.{lora_key}.weight"
)
out[diffusers_key] = original_state_dict.pop(original_key)
for key, tensor in original_state_dict.items():
new_key = key.replace("double_blocks.", "transformer_blocks.")
new_key = new_key.replace("single_blocks.", "single_transformer_blocks.")
out[new_key] = tensor
extra_mappings = {
"img_in": "x_embedder",
"txt_in": "context_embedder",
"time_in.in_layer": "time_guidance_embed.timestep_embedder.linear_1",
"time_in.out_layer": "time_guidance_embed.timestep_embedder.linear_2",
"final_layer.linear": "proj_out",
"final_layer.adaLN_modulation.1": "norm_out.linear",
"single_stream_modulation.lin": "single_stream_modulation.linear",
"double_stream_modulation_img.lin": "double_stream_modulation_img.linear",
"double_stream_modulation_txt.lin": "double_stream_modulation_txt.linear",
}
final_out: Dict[str, torch.Tensor] = {}
for key, tensor in out.items():
new_key = key
for org_key, diff_key in extra_mappings.items():
if key.startswith(org_key):
new_key = key.replace(org_key, diff_key, 1)
break
final_out[new_key] = tensor
sample = _sample_keys(final_out.keys(), 20)
log.info(
"[LoRAFormatAdapter] after AI_TOOLKIT_FLUX conversion, "
"sample keys (<=20): %s",
", ".join(sample),
)
return final_out
def convert_lora_state_dict_by_format(
state_dict: Mapping[str, torch.Tensor],
fmt: LoRAFormat,
@@ -318,6 +496,9 @@ def convert_lora_state_dict_by_format(
if fmt == LoRAFormat.QWEN_IMAGE_STANDARD:
return _convert_qwen_image_standard(state_dict, log)
if fmt == LoRAFormat.AI_TOOLKIT_FLUX:
return _convert_ai_toolkit_flux_lora(state_dict, log)
if fmt == LoRAFormat.XLABS_FLUX:
converted = _convert_xlabs_ai_via_diffusers(state_dict, log)
return _convert_non_diffusers_sd_simple(converted, log)

View File

@@ -208,6 +208,18 @@ def _run_all_tests() -> List[Dict]:
)
)
# AI-Toolkit Flux LoRA (non-diffusers → diffusers).
results.append(
run_single_test(
name="AI-Toolkit Flux LoRA",
repo_id="fal/flux-2-klein-4b-spritesheet-lora",
filename="flux-spritesheet-lora.safetensors",
local_name="flux_spritesheet_lora.safetensors",
expected_before=LoRAFormat.AI_TOOLKIT_FLUX,
expected_after=LoRAFormat.STANDARD,
)
)
# Classic Kohya/A1111 SD LoRA (non-diffusers SD → diffusers).
results.append(
run_single_test(