Repository navigation
[LoRA] add LoKr adapter support (Z-Image, Flux2/Klein) - #14163
christopher5106 wants to merge 3 commits into
Conversation
|
Will review next week but triggering an AI review as well. |
There was a problem hiding this comment.
🤗 Serge says:
Solid, well-tested addition of LoKr adapter loading for Z-Image and Flux2/Klein. The design is coherent: shape-inferred LoKrConfig, alpha baked into the left Kronecker factor per LyCORIS convention, and a fuse-then-map strategy for the fused-QKV case that avoids the lossy per-chunk re-factorization. I verified the key integration points and found no blocking issues.
Correctness (verified, no defects)
_create_lokr_config—collectionsis imported;factorization/LoKrConfigare imported lazily from peft; thedecompose_factorsearch andrank/rank_patternderivation are internally consistent and reproduced exactly by the tests._bake_lokr_alpha— rank is read from the inner dimension of the decomposed factor (w2_b.shape[0]/w1_b.shape[0]), and full-matrix modules correctly drop alpha via therank is Noneguard. Scaling matches the test expectations ((alpha/rank) * kron(w1, w2_a@w2_b)).- Flux2 fused-QKV path —
Flux2Transformer2DModelinheritsfuse_qkv_projectionsfromAttentionMixin; the"Added" in processor nameguard does not spuriously trip (Flux2 processors areFlux2AttnProcessor/Flux2KVAttnProcessor), andfuse_projectionscreatesto_qkv/to_added_qkvused by_get_fused_projectionsat inference. The.attn.to_qkv.substring check does not falsely matchto_qkv_mlp_proj.is_quantizedis a real attribute on the model. reis imported inlora_conversion_utils.py;loggerexists inlora_pipeline.py.
Minor (non-blocking)
- In
_bake_lokr_alpha, a malformed module that has alokr_w*_bbut neitherlokr_w1norlokr_w1_awould raiseKeyErrorat the final assignment. Not reachable for the four in-the-wild formats, so purely defensive — fine to leave.
Tests
Flux2LoKrTestsandZImageLoKrTestscover fused QKV (img+txt), rank-decomposed factors with meaningful alpha, full-rank factors with placeholder alpha, and the LyCORIS underscore format, assertingget_delta_weightequalskron(w1, w2). All referenced imports resolve.
LGTM.
serge v0.1.0 · model: claude-opus-4-8 · 19 LLM turns · 23 tool calls · 110.2s · 880940 in / 7039 out tokens
BenjaminBossan
left a comment
There was a problem hiding this comment.
From the PEFT perspective, this looks correct to me. I have two small comments, but I'll leave it to the Diffusers maintainers to decide if they're worth addressing.
| if not is_correct_format: | ||
| raise ValueError("Invalid LoRA checkpoint. Make sure all LoRA param names contain `'lora'` substring.") | ||
| raise ValueError( | ||
| "Invalid LoRA checkpoint. Make sure all LoRA param names contain the `'lora'` or `'lokr'` substring." |
There was a problem hiding this comment.
The message could be confusing because it still mentions "LoRA checkpoint".
How about "Invalid adapter checkpoint. We currently support LoRA and LoKr."
I wouldn't really mention substring matching, as it's an implementation detail. Same applies to the error messages below.
There was a problem hiding this comment.
Done in 96bb95c: "Invalid adapter checkpoint. We currently support LoRA and LoKr." at all 25 sites this PR touches, and no mention of the substring check anymore. The three pre-existing messages this PR did not touch keep their old text.
| # be split exactly into separate Q/K/V factors. Fuse the model's projections so the adapter maps 1:1. | ||
| needs_fused_qkv = any(".attn.to_qkv." in k or ".attn.to_added_qkv." in k for k in state_dict) | ||
| if needs_fused_qkv: | ||
| if getattr(transformer, "is_quantized", False): |
There was a problem hiding this comment.
Another potential error case would be if there is already an adapter loaded onto a non-fused q, k, or v layer, right? Maybe not terribly likely, but could be worth checking.
There was a problem hiding this comment.
Good catch, done in 96bb95c: before calling fuse_qkv_projections() the loader now walks transformer.named_modules() and raises if any to_q/to_k/to_v or add_{q,k,v}_proj is already a PEFT BaseTunerLayer, naming the first one and pointing at unload_lora_weights(). Fusing would have replaced those modules and silently orphaned the adapter. Covered by test_lokr_fused_qkv_checkpoint_refuses_when_unfused_projections_are_adapted, which injects a plain LoRA on to_q first and checks the model is left unfused.
Adds loading of LoKr (LyCORIS Kronecker product) adapters: - `load_lora_adapter` detects `lokr_` keys and injects a peft `LoKrConfig`, inferred from the tensor shapes via `_create_lokr_config` (decompose factor, per-module rank/alpha patterns). - State dict conversions for the formats in the wild: ai-toolkit Z-Image (dotted diffusers paths under `diffusion_model.`), ai-toolkit BFL Flux2 (fused qkv), LyCORIS underscore format, and bare dotted diffusers paths. - BFL fused-QKV LoKr cannot be split exactly into separate Q/K/V Kronecker factors, so `Flux2LoraLoaderMixin.load_lora_weights` fuses the model's QKV projections and maps the adapter 1:1 (exact). - Alpha follows the LyCORIS convention: scaling applies only to rank-decomposed factors and is baked into the weights at conversion. Fixes huggingface#13221
e6a1235 to
557d685
Compare
|
Both review comments are addressed in 96bb95c (message wording, and refusing a fused-QKV LoKr load when an adapter already sits on the unfused projections, with a test). The branch was also rebased onto current main in September; |
BenjaminBossan
left a comment
There was a problem hiding this comment.
Thanks for addressing my previous comments. There is no issue there, so from that point of view, it's a 👍 from me.
When going over the PR again, I saw that the unit tests specifically check the fused QKV part of the conversion step, but never the whole pipeline. This is probably not easy to achieve, as it would require the reference output. Just flagging this here for @sayakpaul, not sure what the requirements are in Diffusers. At a glance, I couldn't find LoRA conversion unit tests to compare to.
There was a problem hiding this comment.
Nit about naming: The file name should mention "conversion", whereas "lora" isn't really tested here.
There was a problem hiding this comment.
Renamed to tests/lora/test_lokr_conversion.py in c6d5763, thanks!
sayakpaul
left a comment
There was a problem hiding this comment.
Thanks a lot for working on this! I apologize for the delay. The PR looks in a great shape.
| return ait_sd | ||
|
|
||
|
|
||
| def _bake_lokr_alpha(state_dict): |
There was a problem hiding this comment.
| def _bake_lokr_alpha(state_dict): | |
| def _bake_lokr_alpha_(state_dict): |
(nit): since this modifies the state_dict inline.
|
|
||
| converted_state_dict = {} | ||
|
|
||
| # Some Flux2 LoKr checkpoints already store expanded diffusers block names; accept those as-is. |
There was a problem hiding this comment.
If we have, let's also include a checkpoint as an example.
| if len(original_state_dict) > 0: | ||
| raise ValueError(f"`original_state_dict` should be empty at this point but has {original_state_dict.keys()=}.") | ||
|
|
||
| return {f"transformer.{k}": v for k, v in converted_state_dict.items()} |
| state_dict = dict(state_dict) | ||
| _bake_lokr_alpha(state_dict) | ||
|
|
||
| lycoris_key_pattern = re.compile(r"^lycoris_((?:single_)?transformer_blocks)_(\d+)_(.+)\.(.+)$") |
There was a problem hiding this comment.
We expect the stat_dict to always contain lycoris prefix? If so, do we want to check against it and raise if needed?
| state_dict, network_alphas, metadata = self.lora_state_dict(pretrained_model_name_or_path_or_dict, **kwargs) | ||
|
|
||
| is_correct_format = all("lora" in key for key in state_dict.keys()) | ||
| is_correct_format = all("lora" in key or "lokr" in key for key in state_dict.keys()) |
There was a problem hiding this comment.
I would expect to see the LoKR conversion utilities to be called here if this is the case. But that's not the case in:
https://github.com/scenario-labs/diffusers/blob/c6d57637b9b05f0aa841f9619cfa4b947bf95d2f/src/diffusers/loaders/lora_pipeline.py#L243
This pattern seems to be repeating in a couple of places.
| raise ValueError( | ||
| "This LoKr checkpoint targets fused QKV projections, which requires fusing the transformer's " | ||
| "QKV projections, and that is not supported on quantized models. Please load the transformer " | ||
| "without quantization." | ||
| ) |
There was a problem hiding this comment.
Do we know why that is the case?
| # BFL-format LoKr checkpoints apply LoKr to the fused QKV projections, whose Kronecker product delta cannot | ||
| # be split exactly into separate Q/K/V factors. Fuse the model's projections so the adapter maps 1:1. |
There was a problem hiding this comment.
Should this go into peft_utils.py?
| from peft import LoKrConfig | ||
| from peft.tuners.lokr.layer import factorization |
There was a problem hiding this comment.
Do we know the minimum version of PEFT required for this? Does diffusers meet that expectation already?
| # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. | ||
| # See the License for the specific language governing permissions and | ||
| # limitations under the License. | ||
| import unittest |
There was a problem hiding this comment.
We intend on keeping tests/lora very minimal. As such we don't test against community LoRA checkpoints because we cannot control them.
Instead, we could create a lokr.py under tests/models/testing_utils similar to https://github.com/huggingface/diffusers/blob/main/tests/models/testing_utils/lora.py and create a LoKRTesterMixin and maintain a very small set of tests there. WDYT?
… fused-QKV load over an unfused adapter
- "Invalid adapter checkpoint. We currently support LoRA and LoKr." replaces the
message that still said "LoRA checkpoint" and described the substring check.
- Loading a fused-QKV LoKr checkpoint now refuses when an adapter is already
injected on to_q/to_k/to_v or add_{q,k,v}_proj: fuse_qkv_projections() would
replace those modules and orphan it. Covered by a test.
Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
The file exercises checkpoint conversion and loading for LoKr adapters, not LoRA. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
c6d5763 to
7330b4e
Compare
|
@christopher5106 LMK when this is ready for another review. |
What does this PR do?
Adds support for loading LoKr (LyCORIS Kronecker product) adapters, for Z-Image and Flux2/Klein checkpoints in the wild, plus a generic path that works for any model whose LoKr state dict already uses diffusers module names.
Fixes #13221. Also addresses #13261 and #13137 (Flux 2 Klein LoKr).
Follows up on the discussion in #13326 — see "Why not split the fused QKV" below for measurements that motivated a different approach for the fused-QKV part. cc @CalamitousFelicitousness whose prototype this builds on.
How it works
PeftAdapterMixin.load_lora_adapterdetectslokr_keys and injects a peftLoKrConfiginstead of aLoraConfig. The new_create_lokr_configinfers the config from tensor shapes: it reconstructs each module's Kronecker factorization and finds thedecompose_factorunder which peft's ownfactorization()reproduces every shape; per-modulerank_pattern/alpha_patternmake peft recreate full vs. rank-decomposed factors exactly as stored (for full-matrix factors the rank is set tomax(lokr_w2.shape)so peft also creates full matrices).alpha / rankscaling applies only when a factor is rank-decomposed and is baked intolokr_w1at conversion; full-rank modules ignore alpha (ai-toolkit stores a ~1e10 placeholder there). The config setsalpha = rso the runtime scaling is 1.0.diffusion_model.layers.0.attention.to_q.lokr_w1)F16/z-image-turbo-flow-dpo(the checkpoint from #13221)_convert_non_diffusers_lokr_to_diffusers(paths already match the model)diffusion_model.double_blocks.0.img_attn.qkv.lokr_w1)puttmorbidly233/loraklein_snofs_v1_2.safetensors(the checkpoint from #13261)_convert_non_diffusers_flux2_lokr_to_diffuserslycoris_transformer_blocks_0_attn_to_q.lokr_w1)gattaplayer/besch-flux2-klein-9b-lokr-lion-3e-6-bs2-ga2-v02_convert_lycoris_flux2_lokr_to_diffusers.alphakeysbghira/flux2-klein-9b-distillation-lokr_convert_non_diffusers_lokr_to_diffusersWhy not split the fused QKV
BFL-format checkpoints apply LoKr to the fused QKV projection of the double blocks, while the diffusers Flux2 model has separate Q/K/V. For LoRA this split is exact (shared
lora_A, chunkedlora_B), but a Kronecker product delta over the fused projection is mathematically not a Kronecker product per chunk. #13326 approximated each chunk with a rank-1 Van Loan re-factorization; measuring that on the realklein_snofs_v1_2checkpoint gives a mean relative delta error of 69% on the 48 QKV projections (max 79%, min 32%), i.e. the adapter loads but its QKV contribution is mostly destroyed — the same class of problem that got #13997 rejected.Instead, the converter maps these modules to the model's fused
attn.to_qkv/attn.to_added_qkvprojections, andFlux2LoraLoaderMixin.load_lora_weightscallstransformer.fuse_qkv_projections()(idempotent) before injection, with an info log. Loading is then exact. Quantized transformers raise a clear error for such checkpoints since fusing would concatenate packed quantized weights.Tests / verification
ZImageLoKrTestsandFlux2LoKrTests: injectedget_delta_weightequalskron(w1, w2)exactly, covering fused QKV (img + txt streams), rank-decomposed factors with meaningful alpha, full-rank factors with placeholder alpha, and the LyCORIS format.F16/z-image-turbo-flow-dpoloaded into a real-dimension (dim=3840) Z-Image transformer: all deltas exact, scaling 1.0.tests/lora/test_lora_layers_z_image.py(36 passed) andtests/lora/test_lora_layers_flux2.py(39 passed).Who can review?
@sayakpaul @asomoza