mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 17:24:32 +02:00
feat(lora): load dora magnitude vectors saved as magnitude keys
ai-toolkit DoRA saves lora_A/B plus a 1-D per-output magnitude key in place of alpha; PEFT and diffusers name the same quantity lora_magnitude_vector. Neither key was in the suffix table, so such adapters loaded as plain LoRA with the magnitude renormalization silently missing. The semantics match LyCORIS wd_on_out=True row norms, so both keys convert onto the existing dora_scale path. - accept .magnitude and .lora_magnitude_vector in LORA_SUFFIXES and convert at group level in try_load_lora - reshape 1-D vectors to (out, 1): on square layers the apply-time orientation detection would otherwise renormalize the wrong axis - cover square-layer numeric equality, fused-qkv slicing and the PEFT key form in the offline suite
This commit is contained in:
@@ -102,7 +102,10 @@ LORA_SUFFIXES = (
|
||||
".lora_down.weight", ".lora_up.weight", ".lora_mid.weight",
|
||||
".lora_A.weight", ".lora_B.weight",
|
||||
# diff_b: bias delta some saves pair with the weight LoRA, applied as ex_bias.
|
||||
".alpha", ".dora_scale", ".bias", ".diff_b", ".scale",
|
||||
# magnitude / lora_magnitude_vector: DoRA row norms (ai-toolkit / PEFT key
|
||||
# names); converted onto the dora_scale path by try_load_lora.
|
||||
".alpha", ".dora_scale", ".magnitude", ".lora_magnitude_vector",
|
||||
".bias", ".diff_b", ".scale",
|
||||
)
|
||||
LOKR_SUFFIXES = (
|
||||
".lokr_w1", ".lokr_w2",
|
||||
@@ -563,6 +566,16 @@ def try_load_lora(name, network_on_disk, lora_scale, *,
|
||||
for (prefix, base), w in groups.items():
|
||||
if "lora_down.weight" not in w or "lora_up.weight" not in w:
|
||||
continue
|
||||
# DoRA magnitude vectors: ai-toolkit saves `magnitude`, PEFT/diffusers
|
||||
# `lora_magnitude_vector`. Both are 1-D per-output row norms with
|
||||
# dora_scale semantics; reshape to (out, 1) so the apply-time
|
||||
# orientation detection cannot misread square layers as per-input.
|
||||
for mag_key in ("magnitude", "lora_magnitude_vector"):
|
||||
mag = w.get(mag_key)
|
||||
if mag is not None and "dora_scale" not in w:
|
||||
w = dict(w)
|
||||
w.pop(mag_key)
|
||||
w["dora_scale"] = mag.reshape(-1, 1) if mag.ndim == 1 else mag
|
||||
arch_prefix = _resolve_prefix(network_prefix, prefix)
|
||||
for diffusers_path, chunk in resolve_group_targets(resolve_targets, prefix, base):
|
||||
network_key = arch_prefix + diffusers_path.replace(".", "_")
|
||||
|
||||
@@ -347,6 +347,23 @@ def sd_lora_diff_b_fused_qkv():
|
||||
return sd
|
||||
|
||||
|
||||
def sd_lora_aitk_magnitude_dora(w_base):
|
||||
"""ai-toolkit DoRA save: lora_A/B + 1-D magnitude row norms, no alpha key.
|
||||
|
||||
Targets img_attn.proj -> attn.to_out.0, a square Linear (HIDDEN == QKV_OUT
|
||||
in the mock): the case where a raw 1-D vector cannot be disambiguated
|
||||
between per-output and per-input at apply time.
|
||||
"""
|
||||
down = torch.randn(RANK_LORA, QKV_OUT)
|
||||
up = torch.randn(HIDDEN, RANK_LORA)
|
||||
magnitude = w_base.reshape(HIDDEN, -1).norm(dim=1) # 1-D (out,)
|
||||
return {
|
||||
'diffusion_model.double_blocks.1.img_attn.proj.lora_A.weight': down,
|
||||
'diffusion_model.double_blocks.1.img_attn.proj.lora_B.weight': up,
|
||||
'diffusion_model.double_blocks.1.img_attn.proj.magnitude': magnitude,
|
||||
}, down, up, magnitude
|
||||
|
||||
|
||||
def sd_lokr_dora_fused_qkv(per_input=False):
|
||||
"""Kohya LoKR on fused img_attn.qkv with a dora_scale companion.
|
||||
|
||||
@@ -1001,6 +1018,59 @@ def test_legacy_bias_fused_skipped():
|
||||
return True
|
||||
|
||||
|
||||
def test_lora_aitk_magnitude_dora():
|
||||
"""ai-toolkit `magnitude` converts onto the dora_scale path, (out, 1)-shaped.
|
||||
|
||||
Numeric check on a square layer: with the raw 1-D vector the apply-time
|
||||
orientation detection would fall to per-input and renormalize columns;
|
||||
the (out, 1) reshape pins per-output. Reference: (W+D) * m/||W+D||_row - W
|
||||
at scale 1 (ai-toolkit DoRA saves no alpha key).
|
||||
"""
|
||||
torch.manual_seed(0)
|
||||
w_base = torch.randn(HIDDEN, QKV_OUT)
|
||||
sd, down, up, magnitude = sd_lora_aitk_magnitude_dora(w_base)
|
||||
net = _load_via(F.try_load_lora, sd)
|
||||
assert net is not None and len(net.modules) == 1, f'got {net.modules if net else None}'
|
||||
mod = next(iter(net.modules.values()))
|
||||
assert mod.dora_scale is not None and tuple(mod.dora_scale.shape) == (HIDDEN, 1), \
|
||||
f'dora_scale shape {tuple(mod.dora_scale.shape) if mod.dora_scale is not None else None}'
|
||||
updown, _ex_bias = mod.calc_updown(w_base.clone())
|
||||
merged = w_base + up @ down
|
||||
norm = merged.reshape(HIDDEN, -1).norm(dim=1).reshape(HIDDEN, 1) + torch.finfo(w_base.dtype).eps
|
||||
ref = merged * (magnitude.reshape(HIDDEN, 1) / norm) - w_base
|
||||
rel = ((updown - ref).norm() / (ref.norm() + 1e-12)).item()
|
||||
assert torch.allclose(updown, ref, rtol=1e-4, atol=1e-5), f'rel err {rel:.4f}'
|
||||
return True
|
||||
|
||||
|
||||
def test_lora_magnitude_fused_qkv_sliced():
|
||||
"""A fused-qkv magnitude vector converts, reshapes and slices per chunk."""
|
||||
sd = dict(sd_lora_dora_fused_qkv())
|
||||
del sd['lora_unet_double_blocks_0_img_attn_qkv.dora_scale']
|
||||
sd['lora_unet_double_blocks_0_img_attn_qkv.magnitude'] = torch.rand(3 * QKV_OUT) + 0.5
|
||||
net = _load_via(F.try_load_lora, sd)
|
||||
assert net is not None and len(net.modules) == 3, f'got {net.modules if net else None}'
|
||||
for nk, mod in net.modules.items():
|
||||
assert mod.dora_scale is not None and tuple(mod.dora_scale.shape) == (QKV_OUT, 1), \
|
||||
f'{nk}: dora_scale shape {tuple(mod.dora_scale.shape) if mod.dora_scale is not None else None}'
|
||||
return True
|
||||
|
||||
|
||||
def test_lora_peft_magnitude_vector():
|
||||
"""PEFT/diffusers `lora_magnitude_vector` converts the same way."""
|
||||
sd = {
|
||||
'transformer.transformer_blocks.0.attn.to_q.lora_A.weight': torch.randn(RANK_LORA, HIDDEN),
|
||||
'transformer.transformer_blocks.0.attn.to_q.lora_B.weight': torch.randn(QKV_OUT, RANK_LORA),
|
||||
'transformer.transformer_blocks.0.attn.to_q.lora_magnitude_vector': torch.rand(QKV_OUT) + 0.5,
|
||||
}
|
||||
net = _load_via(F.try_load_lora, sd)
|
||||
assert net is not None and len(net.modules) == 1, f'got {net.modules if net else None}'
|
||||
mod = next(iter(net.modules.values()))
|
||||
assert mod.dora_scale is not None and tuple(mod.dora_scale.shape) == (QKV_OUT, 1), \
|
||||
f'dora_scale shape {tuple(mod.dora_scale.shape) if mod.dora_scale is not None else None}'
|
||||
return True
|
||||
|
||||
|
||||
def test_lokr_shape_mismatch_rejected():
|
||||
"""Kron dims that disagree with the module are rejected at load, not at apply."""
|
||||
sd = sd_lokr_bfl_proj()
|
||||
@@ -1487,6 +1557,9 @@ def run_tests():
|
||||
test_dora_per_input_fused_skipped,
|
||||
test_lora_diff_b_fused_qkv_sliced,
|
||||
test_legacy_bias_fused_skipped,
|
||||
test_lora_aitk_magnitude_dora,
|
||||
test_lora_magnitude_fused_qkv_sliced,
|
||||
test_lora_peft_magnitude_vector,
|
||||
test_lokr_shape_mismatch_rejected,
|
||||
test_lokr_fused_shape_mismatch_rejected,
|
||||
test_lokr_bfl_non_fused,
|
||||
|
||||
Reference in New Issue
Block a user