mirror of
https://github.com/vladmandic/automatic
synced 2026-09-18 16:54:33 +02:00
fix(flux2_klein): route through native_transformer with explicit converter
Klein shares Flux2Transformer2DModel with full Flux 2 but uses a smaller config. from_single_file instantiates with the class default (= Flux 2 full size); loading a Klein-shaped file crashed at load_model_dict_into_meta with: Cannot load because double_stream_modulation_img.linear.weight expected shape (36864, 6144), but got (24576, 4096) FLUX2_KLEIN_SPEC fetches Klein's transformer/config.json from the base repo and uses convert_flux2_transformer_checkpoint_to_diffusers explicitly. Probed against snofsSexNudesAndOtherFunStuff_v14Base: 233/233 keys overlap.
This commit is contained in:
@@ -1,8 +1,25 @@
|
||||
import transformers
|
||||
import diffusers
|
||||
from diffusers.loaders.single_file_utils import convert_flux2_transformer_checkpoint_to_diffusers
|
||||
from modules import shared, devices, sd_models, model_quant, sd_hijack_te, sd_hijack_vae
|
||||
from modules.logger import log
|
||||
from pipelines import generic
|
||||
from pipelines.native_transformer import TransformerSpec
|
||||
|
||||
|
||||
# Klein shares Flux2Transformer2DModel with full Flux 2, but uses a smaller
|
||||
# config (hidden_size and friends). diffusers' from_single_file picks the
|
||||
# class default (= Flux 2 full), so loading a Klein-shaped community file
|
||||
# crashes at load_model_dict_into_meta with a shape mismatch like
|
||||
# "expected (36864, 6144), got (24576, 4096)". Routing through
|
||||
# native_transformer pulls the Klein transformer/config.json from the base
|
||||
# repo first and instantiates Flux2Transformer2DModel at the right size,
|
||||
# then runs the diffusers Flux 2 converter to split fused QKV blocks and
|
||||
# rename BFL keys into the diffusers-expected names.
|
||||
FLUX2_KLEIN_SPEC = TransformerSpec(
|
||||
cls=diffusers.Flux2Transformer2DModel,
|
||||
converter=convert_flux2_transformer_checkpoint_to_diffusers,
|
||||
)
|
||||
|
||||
|
||||
def load_flux2_klein(checkpoint_info, diffusers_load_config=None):
|
||||
@@ -15,7 +32,7 @@ def load_flux2_klein(checkpoint_info, diffusers_load_config=None):
|
||||
log.debug(f'Load model: type=Flux2Klein repo="{repo_id}" config={diffusers_load_config} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}')
|
||||
|
||||
# Load transformer - Klein uses Flux2Transformer2DModel (same class as Flux2, different size)
|
||||
transformer = generic.load_transformer(repo_id, cls_name=diffusers.Flux2Transformer2DModel, load_config=diffusers_load_config)
|
||||
transformer = generic.load_transformer(repo_id, cls_name=diffusers.Flux2Transformer2DModel, load_config=diffusers_load_config, native_spec=FLUX2_KLEIN_SPEC)
|
||||
|
||||
# Load text encoder - Klein uses Qwen3 (4B for Klein-4B, 8B for Klein-9B)
|
||||
text_encoder = generic.load_text_encoder(repo_id, cls_name=transformers.Qwen3ForCausalLM, load_config=diffusers_load_config)
|
||||
|
||||
Reference in New Issue
Block a user