diff --git a/studio/backend/core/inference/diffusion.py b/studio/backend/core/inference/diffusion.py index b40add9dd2..7ce31d5ecf 100644 --- a/studio/backend/core/inference/diffusion.py +++ b/studio/backend/core/inference/diffusion.py @@ -42,6 +42,7 @@ from .diffusion_device import ( diffusion_device_target_from_torch_device, resolve_diffusion_device_target, ) +from .diffusion_krea2 import KREA2_FAMILY_NAME, load_krea2_pipeline from .diffusion_memory import ( OFFLOAD_NONE, apply_memory_plan, @@ -205,6 +206,12 @@ _TRUSTED_NON_GGUF_REPOS = frozenset( "black-forest-labs/flux.1-dev", "tongyi-mai/z-image-turbo", "qwen/qwen-image", + # Krea 2: official vendor repos, safetensors-only, no remote code. Loaded + # per-component via core/inference/diffusion_krea2.py (no GGUF variant yet). + # Turbo is the inference model; Raw is the undistilled base Krea recommends + # training LoRAs on (train on Raw, run adapters on Turbo). + "krea/krea-2-turbo", + "krea/krea-2-raw", } ) @@ -1042,15 +1049,21 @@ class DiffusionBackend: # (transformer + VAE + text encoders + scheduler) from the repo # and re-applies any embedded quantization_config (e.g. bnb-4bit), # so a pre-quantized pipeline reloads quantized with no extra config. - pipe_kwargs: dict[str, Any] = {"torch_dtype": dtype} - if hf_token: - pipe_kwargs["token"] = hf_token - # The prefetched snapshot dir keeps from_pretrained off the hub: - # its own snapshot sweep re-downloads files the scoped prefetch - # skipped (root packaged singles, e.g. 24 GB per FLUX.1 repo). - pipe = pipeline_cls.from_pretrained( - _base_local_dir or repo_id, **pipe_kwargs - ) + if fam.name == KREA2_FAMILY_NAME: + # The krea repo ships transformers-5.x style configs the 4.x + # line cannot parse; assemble the pipeline per-component + # (see diffusion_krea2.py for the exact compat story). + pipe = load_krea2_pipeline(repo_id, dtype, hf_token = hf_token) + else: + pipe_kwargs: dict[str, Any] = {"torch_dtype": dtype} + if hf_token: + pipe_kwargs["token"] = hf_token + # The prefetched snapshot dir keeps from_pretrained off the + # hub: its own snapshot sweep re-downloads files the scoped + # prefetch skipped (packaged root singles, 24 GB per FLUX.1). + pipe = pipeline_cls.from_pretrained( + _base_local_dir or repo_id, **pipe_kwargs + ) elif kind == "single_file" and fam.single_file_is_pipeline: # A single-file SDXL-style checkpoint is the WHOLE pipeline # (U-Net + VAE + both text encoders), not a transformer-only file, @@ -1083,10 +1096,17 @@ class DiffusionBackend: single_file_path, **sf_kwargs ) - pipe_kwargs = {"torch_dtype": dtype, "transformer": transformer} - if hf_token: - pipe_kwargs["token"] = hf_token - pipe = pipeline_cls.from_pretrained(_base_local_dir or base, **pipe_kwargs) + if fam.name == KREA2_FAMILY_NAME: + pipe = load_krea2_pipeline( + base, dtype, hf_token = hf_token, transformer = transformer + ) + else: + pipe_kwargs = {"torch_dtype": dtype, "transformer": transformer} + if hf_token: + pipe_kwargs["token"] = hf_token + pipe = pipeline_cls.from_pretrained( + _base_local_dir or base, **pipe_kwargs + ) # Resolve the effective speed mode: GGUF models default to the # near-lossless `default` profile (compile is ~2.2x and sits below diff --git a/studio/backend/core/inference/diffusion_arch_patches.py b/studio/backend/core/inference/diffusion_arch_patches.py index e079e09603..f73085197e 100644 --- a/studio/backend/core/inference/diffusion_arch_patches.py +++ b/studio/backend/core/inference/diffusion_arch_patches.py @@ -27,13 +27,14 @@ source-body check (``_body_has``) confirms the exact lines we rewrite are still future diffusers that changed the block body simply leaves the block UNPATCHED (correctness first) rather than running a stale copy. Kill-switch: ``UNSLOTH_DIFFUSION_ARCH_PATCHES=0``. -Implemented for all four families (extend by adding entries to ``_SPECS``): +Implemented for all five families (extend by adding entries to ``_SPECS``): * qwen-image ``QwenImageTransformerBlock._modulate`` -- modulation addcmul (all 4 sites). * z-image ``ZImageTransformerBlock.forward`` -- the 2 gated-residual addcmuls. * flux.1 ``FluxTransformerBlock.forward`` (inline norm2 modulation + 4 gated residuals) + ``FluxSingleTransformerBlock.forward`` (residual + gate*proj_out). * flux.2-klein ``Flux2TransformerBlock.forward`` (4 inline modulations + 4 gated residuals) + ``Flux2SingleTransformerBlock.forward`` (inline modulation + gated residual). + * krea-2 ``Krea2TransformerBlock.forward`` (2 inline modulations + 2 gated residuals). """ from __future__ import annotations @@ -468,6 +469,53 @@ def _spec_flux2_single(): return (cls, "forward", _flux2_single_forward) +# ===================================================================================== +# krea-2: Krea2TransformerBlock.forward (2 inline modulations + 2 gated residuals) +# ===================================================================================== +def _krea2_block_forward( + self, + hidden_states, + temb, + image_rotary_emb, + attention_mask = None, +): + """diffusers 0.39 ``Krea2TransformerBlock.forward`` with the two inline modulations + ``(1 + scale) * norm(x) + shift`` and the two gated residuals ``x + gate * out`` each + fused to one ``torch.addcmul``.""" + # temb: (B, 1, 6 * hidden_size), shared across all blocks; each block only learns an + # additive table. + modulation = temb.unflatten(-1, (6, -1)) + self.scale_shift_table + prescale, preshift, pregate, postscale, postshift, postgate = modulation.unbind(-2) + + norm1 = self.norm1(hidden_states) + attn_out = self.attn( + torch.addcmul(preshift, norm1, 1 + prescale), + attention_mask = attention_mask, + image_rotary_emb = image_rotary_emb, + ) + hidden_states = torch.addcmul(hidden_states, pregate, attn_out) + norm2 = self.norm2(hidden_states) + ff_out = self.ff(torch.addcmul(postshift, norm2, 1 + postscale)) + return torch.addcmul(hidden_states, postgate, ff_out) + + +def _spec_krea2_forward(): + try: + from diffusers.models.transformers.transformer_krea2 import Krea2TransformerBlock as cls + except Exception: # noqa: BLE001 + return None + orig = getattr(cls, "forward", None) + if orig is None or not _body_has( + orig, + "(1.0 + prescale) * self.norm1(hidden_states) + preshift", + "hidden_states = hidden_states + pregate * attn_out", + "(1.0 + postscale) * self.norm2(hidden_states) + postshift", + "hidden_states = hidden_states + postgate * ff_out", + ): + return None + return (cls, "forward", _krea2_block_forward) + + # ===================================================================================== # registry + lifecycle # ===================================================================================== @@ -480,6 +528,7 @@ _SPECS: tuple[Callable[[], Optional[tuple]], ...] = ( _spec_flux_single, _spec_flux2_double, _spec_flux2_single, + _spec_krea2_forward, ) # (cls, attr) pairs we successfully patched, for an exact reverse. diff --git a/studio/backend/core/inference/diffusion_families.py b/studio/backend/core/inference/diffusion_families.py index 88ce4007dc..a6cf3d7b03 100644 --- a/studio/backend/core/inference/diffusion_families.py +++ b/studio/backend/core/inference/diffusion_families.py @@ -126,6 +126,13 @@ class DiffusionFamily: # Recommended base repos to train FROM, most-preferred first (e.g. a QLoRA-friendly # prequant repo, then a bf16 repo). Surfaced by the Train UI as the base-model choices. train_base_repos: tuple[str, ...] = field(default_factory = tuple) + # When set, deploying a LoRA trained on this family loads THIS repo instead of the + # checkpoint it was trained on -- for families whose release guidance is to train on one + # checkpoint but run adapters on another (Krea: train on Raw, preview on Turbo). Both + # sides must be the same precision so the swap never enlarges the load (unlike the + # nf4 -> bf16 gap that would risk an OOM on deploy). Unset elsewhere, so every other + # family deploys on the base it was trained on. + deploy_base_repo: Optional[str] = None # Keyed by architecture, not per model variant: a checkpoint's specific base repo @@ -282,6 +289,31 @@ _FAMILIES: tuple[DiffusionFamily, ...] = ( ("Comfy-Org/z_image_turbo", "split_files/text_encoders/qwen_3_4b.safetensors", "llm"), ), ), + # Krea 2 (diffusers >= 0.39): a ~12B single-stream flow-matching DiT with a + # Qwen3-VL-4B text encoder (12 tapped hidden layers fused in-transformer) and the + # Qwen-Image VAE. Loaded per-component through core/inference/diffusion_krea2.py + # because the krea repo ships transformers-5.x style configs (see that module). + # No GGUF/sd.cpp mapping yet, so the no-GPU route falls back to diffusers. + DiffusionFamily( + name = "krea-2", + pipeline_class = "Krea2Pipeline", + transformer_class = "Krea2Transformer2DModel", + base_repo = "krea/Krea-2-Turbo", + aliases = ("krea2",), + # LoRA training via the DiT trainer (no prequant repo yet, so nf4 quantizes the + # 12B transformer on the fly under the default precision). Krea's release guidance + # is explicit: train LoRAs on the undistilled Raw checkpoint and apply them on + # Turbo for inference, so Raw is the default training base and Turbo stays the + # inference/base repo. + trainable = True, + train_base_repos = ("krea/Krea-2-Raw", "krea/Krea-2-Turbo"), + # Per Krea's guidance, adapters trained on Raw are meant to run on Turbo; deploy + # previews them on Turbo (same bf16 precision, so the swap never enlarges the load). + deploy_base_repo = "krea/Krea-2-Turbo", + # The checkpoint is exported bf16-only (the model card pins bfloat16); fp16 is + # unvalidated upstream, so keep the fp16 fallback off like z-image. + fp16_incompatible = True, + ), # SDXL is the one U-Net family here: the denoiser is ``pipe.unet`` # (UNet2DConditionModel), not a DiT ``pipe.transformer``, and a single-file # ``.safetensors`` is the WHOLE pipeline rather than a transformer-only file. @@ -429,6 +461,14 @@ def resolve_base_repo(fam: DiffusionFamily, base_repo: Optional[str]) -> str: # (studio/frontend/src/features/images/images-page.tsx); keep the two in sync. _GENERATION_DEFAULTS: tuple[tuple[str, int, float], ...] = ( ("z-image-turbo", 9, 0.0), + # Krea 2 Raw is the undistilled base (both Turbo and Raw are in _TRUSTED_NON_GGUF_REPOS, so + # either is inference-loadable): its model card runs 52 steps at guidance 3.5. It must precede + # the generic "krea" key, or the distilled recipe below would degrade a Raw load to garbage. + ("krea-2-raw", 52, 3.5), + # Krea 2 Turbo is distilled (TDM): 8 steps, no CFG -- matching the Create UI seed, so the + # OpenAI /v1/images/generations route uses the documented recipe instead of falling through + # to the generic (9, 0.0). "krea" then covers Turbo and any other krea id but Raw (above). + ("krea", 8, 0.0), ("flux.1-schnell", 4, 0.0), # Kontext (editing) before the generic flux.1: ~28 steps, lower guidance (~2.5). ("kontext", 28, 2.5), diff --git a/studio/backend/core/inference/diffusion_krea2.py b/studio/backend/core/inference/diffusion_krea2.py new file mode 100644 index 0000000000..86a1d0413c --- /dev/null +++ b/studio/backend/core/inference/diffusion_krea2.py @@ -0,0 +1,156 @@ +"""Krea 2 pipeline loader: assembles ``Krea2Pipeline`` from per-component loads. + +Why not ``Krea2Pipeline.from_pretrained``: the ``krea/Krea-2-Turbo`` repo was exported +with transformers 5.2, and two of its configs use 5.x-only conventions that the 4.x +line cannot parse: + +- ``tokenizer/tokenizer_config.json`` declares ``Qwen2Tokenizer`` (slow -- 5.x unified + slow/fast under the plain name) but ships only ``tokenizer.json``. 4.x's slow class + needs vocab.json/merges.txt (absent), and its fast class trips over + ``extra_special_tokens`` stored as a LIST (4.x expects a dict). Loading the fast + class with an explicit ``extra_special_tokens = {}`` override is id-identical: every + listed token is already registered as an added special token inside tokenizer.json, + and the pipeline templates prompts manually (it never uses a chat template). +- ``text_encoder/config.json`` keeps the rope settings under ``rope_parameters`` (the + 5.x name). 4.x reads ``rope_scaling`` + a top-level ``rope_theta`` and crashes on the + missing key (``NoneType.get``). The values are copied across verbatim -- and they + equal 4.x's Qwen3-VL defaults (theta 5e6, mrope_section [24, 20, 20], interleaved + mrope applied unconditionally), so the rotary embedding is numerically identical. + The state dict itself round-trips 1:1 (checkpoint keys == 4.x module keys). + +``from_pretrained`` additionally type-checks a passed ``tokenizer`` against the +declared SLOW class (a fast tokenizer does not subclass it), so the pipeline is built +through its constructor instead, forwarding the ``is_distilled`` / +``text_encoder_select_layers`` / ``patch_size`` init config from model_index.json -- +Turbo's fixed mu=1.15 timestep shift rides on ``is_distilled = True``, so dropping it +would silently degrade the schedule. + +Both workarounds are self-disabling on a transformers 5.x runtime: the plain tokenizer +load succeeds (no fallback taken) and ``rope_scaling`` parses non-None (no patch). +""" + +from __future__ import annotations + +import json +from pathlib import Path +from typing import Any, Optional + +from loggers import get_logger + +logger = get_logger(__name__) + +KREA2_FAMILY_NAME = "krea-2" + + +def load_krea2_tokenizer(repo_id: str, hf_token: Optional[str] = None): + """The Krea 2 tokenizer, tolerating the repo's transformers-5.x tokenizer config.""" + from transformers import AutoTokenizer + + kwargs: dict[str, Any] = {"subfolder": "tokenizer"} + if hf_token: + kwargs["token"] = hf_token + try: + return AutoTokenizer.from_pretrained(repo_id, **kwargs) + except Exception as exc: # noqa: BLE001 -- 4.x config-parse failure, retry with override + logger.info("diffusion.krea2 tokenizer compat fallback: %s", exc) + return AutoTokenizer.from_pretrained(repo_id, extra_special_tokens = {}, **kwargs) + + +def remap_rope_parameters(text_config) -> None: + """Copy 5.x ``rope_parameters`` onto the 4.x ``rope_scaling`` / ``rope_theta`` slots + in place. A no-op when ``rope_scaling`` already parsed non-None (a 5.x runtime) or + the config carries no ``rope_parameters`` dict.""" + rope_parameters = getattr(text_config, "rope_parameters", None) + if getattr(text_config, "rope_scaling", None) is None and isinstance(rope_parameters, dict): + text_config.rope_scaling = {k: v for k, v in rope_parameters.items() if k != "rope_theta"} + if "rope_theta" in rope_parameters: + text_config.rope_theta = rope_parameters["rope_theta"] + + +def load_krea2_text_encoder( + repo_id: str, + dtype, + hf_token: Optional[str] = None, +): + """The Qwen3-VL text encoder, remapping 5.x ``rope_parameters`` for a 4.x runtime.""" + from transformers import AutoConfig, Qwen3VLModel + + kwargs: dict[str, Any] = {"subfolder": "text_encoder"} + if hf_token: + kwargs["token"] = hf_token + config = AutoConfig.from_pretrained(repo_id, **kwargs) + remap_rope_parameters(getattr(config, "text_config", config)) + return Qwen3VLModel.from_pretrained(repo_id, config = config, dtype = dtype, **kwargs) + + +def _load_model_index(repo_id: str, hf_token: Optional[str] = None) -> dict[str, Any]: + """model_index.json as a dict, from a local path or the Hub cache.""" + is_local_dir = False + try: + root = Path(repo_id).expanduser() + is_local_dir = root.is_dir() + local = root / "model_index.json" + if local.is_file(): + return json.loads(local.read_text()) + except OSError: + pass + if is_local_dir: + # A local checkpoint dir without the file must fail clearly here: falling through + # to hf_hub_download with a filesystem path as the repo id would die with an + # opaque HFValidationError instead. + raise FileNotFoundError(f"model_index.json not found in local model dir {repo_id}") + from huggingface_hub import hf_hub_download + + path = hf_hub_download(repo_id, "model_index.json", token = hf_token or None) + return json.loads(Path(path).read_text()) + + +def load_krea2_pipeline( + repo_id: str, + dtype, + hf_token: Optional[str] = None, + transformer = None, + with_transformer: bool = True, +): + """A ready ``Krea2Pipeline`` for ``repo_id`` (still on CPU; caller places it). + + ``transformer`` lets the single-file/quant paths hand in a prebuilt denoiser; + ``with_transformer = False`` skips the (26 GB) denoiser entirely for a + conditioning-only pipeline (the trainer's phased load). The remaining components + (VAE, text encoder, tokenizer, scheduler) come from the repo. + """ + import diffusers + + # diffusers gained Krea2Pipeline in 0.39; on an older install the getattr chain below + # would die with a bare AttributeError mid-load, so fail first with the actionable fix. + if not hasattr(diffusers, "Krea2Pipeline"): + raise RuntimeError( + f"Krea 2 needs diffusers >= 0.39.0 (Krea2Pipeline); this environment has " + f"diffusers {getattr(diffusers, '__version__', 'unknown')}. " + f"Upgrade with: pip install -U diffusers" + ) + + token = hf_token or None + tokenizer = load_krea2_tokenizer(repo_id, hf_token = token) + text_encoder = load_krea2_text_encoder(repo_id, dtype, hf_token = token) + scheduler = diffusers.FlowMatchEulerDiscreteScheduler.from_pretrained( + repo_id, subfolder = "scheduler", token = token + ) + vae = diffusers.AutoencoderKLQwenImage.from_pretrained( + repo_id, subfolder = "vae", torch_dtype = dtype, token = token + ) + if transformer is None and with_transformer: + transformer = diffusers.Krea2Transformer2DModel.from_pretrained( + repo_id, subfolder = "transformer", torch_dtype = dtype, token = token + ) + model_index = _load_model_index(repo_id, hf_token = token) + return diffusers.Krea2Pipeline( + scheduler = scheduler, + vae = vae, + text_encoder = text_encoder, + tokenizer = tokenizer, + transformer = transformer, + text_encoder_select_layers = model_index.get("text_encoder_select_layers"), + is_distilled = bool(model_index.get("is_distilled", False)), + patch_size = int(model_index.get("patch_size", 2)), + ) diff --git a/studio/backend/core/inference/diffusion_transformer_quant.py b/studio/backend/core/inference/diffusion_transformer_quant.py index 4d31ae6f25..b990efc563 100644 --- a/studio/backend/core/inference/diffusion_transformer_quant.py +++ b/studio/backend/core/inference/diffusion_transformer_quant.py @@ -66,6 +66,10 @@ _INT8_EXCLUDE_NAME_TOKENS = ( "guidance_embed", "time_text_embed", # Flux/Qwen time_text_embed.* (pooled-text + timestep); NOT context_embedder "pooled", + # Krea 2's Krea2TimestepEmbedding ("time_embed.linear_2", 6144->6144 at M = batch); + # its other M=1 projection ("time_mod_proj") is already caught by "_mod", and + # img_in / final_layer.linear / text_fusion.projector fall under min_features. + "time_embed", ) diff --git a/studio/backend/core/training/diffusion_dit_trainer.py b/studio/backend/core/training/diffusion_dit_trainer.py index ff40c43f00..08b17bede3 100644 --- a/studio/backend/core/training/diffusion_dit_trainer.py +++ b/studio/backend/core/training/diffusion_dit_trainer.py @@ -72,6 +72,26 @@ _FLUX_TARGETS = ( ) _QWEN_TARGETS = _FLUX_TARGETS _ZIMAGE_TARGETS = ("to_q", "to_k", "to_v", "to_out.0") +# The Krea 2 authors' recommended default target set (their DreamBooth reference script): +# attention + SwiGLU + the text-fusion projector + the conditioning embedders. For long +# runs they suggest narrowing to the attention layers so prompt adherence doesn't drop. +_KREA2_TARGETS = ( + "img_in", + "final_layer.linear", + "to_q", + "to_k", + "to_v", + "to_out.0", + "to_gate", + "ff.up", + "ff.down", + "text_fusion.projector", + "txt_in.linear_1", + "txt_in.linear_2", + "time_embed.linear_1", + "time_embed.linear_2", + "time_mod_proj", +) def _select_lora_targets( @@ -698,6 +718,105 @@ def _zimage_save(pipe_cls, out_dir, transformer_lora_layers): ) +# ── Krea 2 ──────────────────────────────────────────────────────────────────── +def _krea2_load_conditioners(cfg, device, weight_dtype): + # The krea repo ships transformers-5.x style configs the pinned 4.x line cannot + # parse, so the conditioning pipeline is assembled per-component (with the tokenizer + # and rope compat) instead of from_pretrained(transformer = None); see + # core/inference/diffusion_krea2.py for the exact compat story. + import torch + from core.inference.diffusion_krea2 import load_krea2_pipeline + + pipe = load_krea2_pipeline( + cfg.base_model, torch.bfloat16, hf_token = cfg.hf_token, with_transformer = False + ) + pipe.vae.to(device, dtype = torch.float32) + return pipe, pipe.vae + + +def _krea2_load_transformer(cfg, device, weight_dtype, base_precision): + # The transformer subfolder is diffusers-format (no transformers compat needed). + # There is no prequant repo yet, so nf4 quantizes the 12B transformer on the fly. + from diffusers import Krea2Transformer2DModel + return _load_dit_transformer(Krea2Transformer2DModel, cfg, device, base_precision) + + +def _krea2_encode_prompts(pipe, captions, device): + import torch + + _encoders_to_device(pipe, device) + out = [] + with torch.no_grad(): + for cap in captions: + # encode_prompt pads/truncates to the fixed max_sequence_length, so every + # embed is [1, 512, num_text_layers, 2560] with a [1, 512] validity mask -- + # static shapes (the padding sits mid-template, BEFORE the assistant suffix, + # matching how the model was sampled at training time). + pe, mask = pipe.encode_prompt( + prompt = cap, + device = device, + num_images_per_prompt = 1, + max_sequence_length = 512, + ) + out.append((pe.cpu(), mask.cpu())) + return out + + +# Krea 2 conditions on the Qwen-Image VAE (AutoencoderKLQwenImage) with the same +# per-channel latents_mean / latents_std normalisation, so latent encoding is shared. +_krea2_encode_latents = _qwen_encode_latents +_krea2_encode_latent_stats = _qwen_encode_latent_stats + + +def _krea2_collate( + entries, + device, + weight_dtype, + pad_to = None, +): + import torch + + # Fixed-length embeds (see _krea2_encode_prompts), so collation is a plain concat + # with the mask riding along; ``pad_to`` is moot because the shapes are static. + pe_b = torch.cat([e[0] for e in entries]).to(device = device, dtype = weight_dtype) + mask_b = torch.cat([e[1] for e in entries]).to(device) + return (pe_b, mask_b) + + +def _krea2_forward(transformer, noisy, timesteps, sigmas, embeds_batch, cfg, device, weight_dtype): + from diffusers import Krea2Pipeline + + pe, mask = embeds_batch + # [B,16,1,H,W] -> [B, (H/2)*(W/2), 64] 2x2 patches. Krea2Pipeline._pack_latents / + # _unpack_latents are instance methods (they read self.patch_size), so the packing is + # inlined here exactly like the reference DreamBooth script (patch_size = 2). + bsz, c, _f, h, w = noisy.shape + packed = noisy.reshape(bsz, c, h // 2, 2, w // 2, 2) + packed = packed.permute(0, 2, 4, 1, 3, 5).reshape(bsz, (h // 2) * (w // 2), c * 4) + # Text tokens sit at the rotary origin, so one shared position grid serves the batch. + position_ids = Krea2Pipeline.prepare_position_ids(pe.shape[1], h // 2, w // 2, device) + pred = transformer( + hidden_states = packed, + encoder_hidden_states = pe, + timestep = timesteps / 1000, + position_ids = position_ids, + encoder_attention_mask = mask, + return_dict = False, + )[0] + pred = pred.view(bsz, h // 2, w // 2, c, 2, 2) + pred = pred.permute(0, 3, 1, 4, 2, 5) + return pred.reshape(bsz, c, 1, h, w) + + +def _krea2_save(pipe_cls, out_dir, transformer_lora_layers): + from diffusers import Krea2Pipeline + Krea2Pipeline.save_lora_weights( + save_directory = out_dir, + transformer_lora_layers = transformer_lora_layers, + weight_name = DEFAULT_LORA_FILENAME, + ) + + _SPECS: dict[str, _FamilySpec] = { "flux.1": _FamilySpec( family = "flux.1", @@ -741,6 +860,20 @@ _SPECS: dict[str, _FamilySpec] = { forward = _zimage_forward, save = _zimage_save, ), + "krea-2": _FamilySpec( + family = "krea-2", + lora_targets = _KREA2_TARGETS, + force_bf16 = True, + dense_bf16_gb = 26.3, + load_conditioners = _krea2_load_conditioners, + load_transformer = _krea2_load_transformer, + encode_prompts = _krea2_encode_prompts, + encode_latents = _krea2_encode_latents, + encode_latent_stats = _krea2_encode_latent_stats, + collate = _krea2_collate, + forward = _krea2_forward, + save = _krea2_save, + ), } diff --git a/studio/backend/core/training/diffusion_train_common.py b/studio/backend/core/training/diffusion_train_common.py index d6c58ae0bd..385d4d8be1 100644 --- a/studio/backend/core/training/diffusion_train_common.py +++ b/studio/backend/core/training/diffusion_train_common.py @@ -208,7 +208,7 @@ def get_trainer(family: str) -> Callable[..., str]: if key == "sdxl": from core.training.diffusion_lora_trainer import run_diffusion_lora_training return run_diffusion_lora_training - if key in ("flux.1", "qwen-image", "z-image"): + if key in ("flux.1", "qwen-image", "z-image", "krea-2"): from core.training.diffusion_dit_trainer import run_dit_lora_training return run_dit_lora_training raise ValueError(f"No trainer is registered for family {family!r}.") @@ -222,6 +222,9 @@ FAMILY_TRAIN_DEFAULTS: dict[str, dict[str, Any]] = { "flux.1": {"lora_rank": 16, "learning_rate": 1e-4, "resolution": 512}, "qwen-image": {"lora_rank": 16, "learning_rate": 5e-5, "resolution": 512}, "z-image": {"lora_rank": 16, "learning_rate": 1e-4, "resolution": 768}, + # The Krea 2 authors' recommended starting point (their DreamBooth script defaults): + # rank/alpha 32, lr 3e-4, 512px. + "krea-2": {"lora_rank": 32, "learning_rate": 3e-4, "resolution": 512}, } @@ -237,6 +240,7 @@ _FAMILY_LABELS = { "flux.1": "FLUX.1-dev", "qwen-image": "Qwen-Image", "z-image": "Z-Image", + "krea-2": "Krea 2", } _FAMILY_VRAM_NOTES = { "sdxl": "Trains on ~12 GB+ (bf16 LoRA). The lightest, fastest option.", @@ -246,13 +250,17 @@ _FAMILY_VRAM_NOTES = { ), "qwen-image": "20B model, QLoRA (nf4) by default (~24 GB+). The heaviest option.", "z-image": "6B model, QLoRA (nf4) by default (~12 GB+). bf16 only.", + "krea-2": ( + "12B model, QLoRA (nf4) by default (~18 GB+). bf16 only. Trains on the " + "undistilled Krea-2-Raw (Krea's guidance: train on Raw, run adapters on Turbo)." + ), } # The flow-matching DiT families (run by diffusion_dit_trainer). They expose the # base_precision / compile levers and require bf16 compute on CUDA; SDXL is absent because # it uses its own mixed_precision path. Kept as a set so the UI gate, the bf16 preflight, # and any future dispatch stay in sync. -_DIT_TRAIN_FAMILIES = frozenset({"flux.1", "qwen-image", "z-image"}) +_DIT_TRAIN_FAMILIES = frozenset({"flux.1", "qwen-image", "z-image", "krea-2"}) def bf16_unsupported_reason(resolved_family: str) -> Optional[str]: @@ -349,6 +357,8 @@ def family_train_infos() -> list[dict[str, Any]]: "precision_modes": fam_modes, "recommended_precision": "nf4" if (not is_dit or dit_block) else dit_recommended, "supports_compile": bool(is_dit and not dit_block), + # Krea trains on Raw but previews adapters on Turbo; None elsewhere. + "deploy_base": fam.deploy_base_repo, } ) return infos diff --git a/studio/backend/models/training.py b/studio/backend/models/training.py index 95a44c4b19..d6d5241072 100644 --- a/studio/backend/models/training.py +++ b/studio/backend/models/training.py @@ -870,6 +870,10 @@ class DiffusionTrainableFamily(BaseModel): precision_modes: List[str] = Field(default_factory = list) recommended_precision: str = "nf4" supports_compile: bool = False + # When set, deploying a LoRA trained on this family previews it on this repo instead of + # the training base (Krea trains on Raw but runs adapters on Turbo). Null for families + # that deploy on the base they trained on. + deploy_base: Optional[str] = None class DiffusionTrainingInfoResponse(BaseModel): diff --git a/studio/backend/tests/test_diffusion_arch_patches.py b/studio/backend/tests/test_diffusion_arch_patches.py index 77276a0fb1..ceeb047669 100644 --- a/studio/backend/tests/test_diffusion_arch_patches.py +++ b/studio/backend/tests/test_diffusion_arch_patches.py @@ -236,6 +236,37 @@ def test_flux2_single_forward_matches_stock(): _close_any(got, ref) +def test_krea2_forward_matches_stock(): + from diffusers.models.transformers.transformer_krea2 import ( + Krea2RotaryPosEmbed, + Krea2TransformerBlock, + ) + from diffusers.pipelines.krea2.pipeline_krea2 import Krea2Pipeline + + torch.manual_seed(4) + blk = Krea2TransformerBlock( + hidden_size = D, + intermediate_size = 2 * D, + num_heads = H, + num_kv_heads = H // 2, + norm_eps = 1e-6, + ).eval() + # Give the zero-init modulation table real values so all six scale/shift/gate + # branches contribute to the output. + with torch.no_grad(): + blk.scale_shift_table.normal_() + # A [text + 2x2 image grid] sequence with the real rotary embed (axes sum to head_dim). + position_ids = Krea2Pipeline.prepare_position_ids(4, 2, 2, torch.device("cpu")) + rope = Krea2RotaryPosEmbed(theta = 10000, axes_dim = [D // H // 2, D // H // 4, D // H // 4]) + image_rotary_emb = rope(position_ids) + hs = torch.randn(B, position_ids.shape[0], D) + tm = torch.randn(B, 1, 6 * D) + with torch.inference_mode(): + ref = Krea2TransformerBlock.forward(blk, hs, tm, image_rotary_emb).clone() + got = ap._krea2_block_forward(blk, hs, tm, image_rotary_emb) + torch.testing.assert_close(got, ref, atol = 1e-5, rtol = 1e-4) + + # ── lifecycle ─────────────────────────────────────────────────────────────────── @@ -246,7 +277,7 @@ def test_install_idempotent_and_reversible(): q_orig, z_orig = Q._modulate, Z.forward n1 = ap.install_arch_patches() n2 = ap.install_arch_patches() # idempotent - assert n1 == 6 and n2 == n1 # qwen + z-image + flux.1 x2 + flux.2 x2 + assert n1 == 7 and n2 == n1 # qwen + z-image + flux.1 x2 + flux.2 x2 + krea-2 assert Q._modulate is not q_orig and Z.forward is not z_orig assert ap.is_installed() diff --git a/studio/backend/tests/test_diffusion_backend.py b/studio/backend/tests/test_diffusion_backend.py index 65516fb5e9..8c8332ded1 100644 --- a/studio/backend/tests/test_diffusion_backend.py +++ b/studio/backend/tests/test_diffusion_backend.py @@ -82,6 +82,14 @@ def test_detect_family_from_repo_id(): assert detect_family("unsloth/FLUX.1-dev-GGUF").name == "flux.1" # A plain Qwen-Image checkpoint must still resolve to the base family, not edit. assert detect_family("unsloth/Qwen-Image-2512-GGUF").name == "qwen-image" + # Krea 2 (diffusers >= 0.39): bf16-only single-stream DiT, no GGUF/sd.cpp mapping. + krea2 = detect_family("krea/Krea-2-Turbo") + assert krea2.name == "krea-2" + assert krea2.pipeline_class == "Krea2Pipeline" + assert krea2.transformer_class == "Krea2Transformer2DModel" + assert krea2.cfg_kwarg == "guidance_scale" + assert krea2.fp16_incompatible is True + assert krea2.sd_cpp_text_encoders == () assert detect_family("meta-llama/Llama-3-8B") is None @@ -122,7 +130,7 @@ def test_detect_family_override(): def test_supported_family_names(): names = supported_family_names() # The unknown-model error lists these, so the key families must be present. - for expected in ("flux.1", "flux.2-klein", "flux.2-dev", "qwen-image", "z-image"): + for expected in ("flux.1", "flux.2-klein", "flux.2-dev", "qwen-image", "z-image", "krea-2"): assert expected in names # Every listed name is a valid family_override (round-trips through detect_family). for name in names: diff --git a/studio/backend/tests/test_diffusion_dit_trainer.py b/studio/backend/tests/test_diffusion_dit_trainer.py index 63cc01922c..c40a0575a3 100644 --- a/studio/backend/tests/test_diffusion_dit_trainer.py +++ b/studio/backend/tests/test_diffusion_dit_trainer.py @@ -29,15 +29,18 @@ from core.training.diffusion_train_common import ( ) -def test_specs_cover_the_three_dit_families(): - assert set(_SPECS) == {"flux.1", "qwen-image", "z-image"} - # FLUX / Qwen share the added-kv attention target set; Z-Image is single-stream. +def test_specs_cover_the_dit_families(): + assert set(_SPECS) == {"flux.1", "qwen-image", "z-image", "krea-2"} + # FLUX / Qwen share the added-kv attention target set; Z-Image and Krea 2 are + # single-stream. assert "add_q_proj" in _SPECS["flux.1"].lora_targets assert "add_q_proj" in _SPECS["qwen-image"].lora_targets assert "add_q_proj" not in _SPECS["z-image"].lora_targets - # Z-Image and Qwen are bf16-only. + assert "add_q_proj" not in _SPECS["krea-2"].lora_targets + # Z-Image, Qwen and Krea 2 are bf16-only. assert _SPECS["z-image"].force_bf16 is True assert _SPECS["qwen-image"].force_bf16 is True + assert _SPECS["krea-2"].force_bf16 is True def test_select_lora_targets_uses_family_default_for_generic_config(): diff --git a/studio/backend/tests/test_diffusion_krea2.py b/studio/backend/tests/test_diffusion_krea2.py new file mode 100644 index 0000000000..86ff8fdcb0 --- /dev/null +++ b/studio/backend/tests/test_diffusion_krea2.py @@ -0,0 +1,244 @@ +"""Unit tests for the Krea 2 per-component pipeline loader (CPU-only, no network).""" + +from __future__ import annotations + +import json +import sys +from types import SimpleNamespace + +import pytest + +from core.inference.diffusion_krea2 import ( + KREA2_FAMILY_NAME, + _load_model_index, + load_krea2_pipeline, + remap_rope_parameters, +) + + +# ── rope_parameters (transformers 5.x) -> rope_scaling (4.x) remap ────────── + + +def test_remap_rope_parameters_copies_5x_values(): + cfg = SimpleNamespace( + rope_scaling = None, + rope_theta = 1000000.0, + rope_parameters = { + "mrope_interleaved": True, + "mrope_section": [24, 20, 20], + "rope_theta": 5000000, + "rope_type": "default", + }, + ) + remap_rope_parameters(cfg) + # rope_theta is hoisted to the top-level slot, the rest lands in rope_scaling. + assert cfg.rope_theta == 5000000 + assert cfg.rope_scaling == { + "mrope_interleaved": True, + "mrope_section": [24, 20, 20], + "rope_type": "default", + } + + +def test_remap_rope_parameters_noop_on_5x_runtime_or_plain_4x_config(): + # rope_scaling already parsed (5.x runtime exposing the alias): untouched. + parsed = {"rope_type": "default", "mrope_section": [1, 2, 3]} + cfg = SimpleNamespace(rope_scaling = parsed, rope_theta = 7.0, rope_parameters = {"x": 1}) + remap_rope_parameters(cfg) + assert cfg.rope_scaling is parsed + assert cfg.rope_theta == 7.0 + # No rope_parameters at all (a plain 4.x-exported config): untouched. + cfg = SimpleNamespace(rope_scaling = None, rope_theta = 7.0) + remap_rope_parameters(cfg) + assert cfg.rope_scaling is None + + +# ── model_index.json resolution ────────────────────────────────────────────── + + +def test_load_model_index_from_local_path(tmp_path): + (tmp_path / "model_index.json").write_text(json.dumps({"is_distilled": True, "patch_size": 2})) + assert _load_model_index(str(tmp_path)) == {"is_distilled": True, "patch_size": 2} + + +# ── pipeline assembly threads the model_index init config ──────────────────── + + +def test_load_krea2_pipeline_threads_init_config(monkeypatch, tmp_path): + (tmp_path / "model_index.json").write_text( + json.dumps( + { + "is_distilled": True, + "patch_size": 2, + "text_encoder_select_layers": [2, 5, 8], + } + ) + ) + + captured: dict = {} + + class _FromPretrained: + def __init__(self, tag): + self.tag = tag + + def from_pretrained(self, repo_id, **kwargs): + captured.setdefault("components", {})[self.tag] = (repo_id, kwargs) + return SimpleNamespace(tag = self.tag) + + def _pipeline_ctor(**kwargs): + captured["pipeline"] = kwargs + return SimpleNamespace(**kwargs) + + fake_diffusers = SimpleNamespace( + FlowMatchEulerDiscreteScheduler = _FromPretrained("scheduler"), + AutoencoderKLQwenImage = _FromPretrained("vae"), + Krea2Transformer2DModel = _FromPretrained("transformer"), + Krea2Pipeline = _pipeline_ctor, + ) + monkeypatch.setitem(sys.modules, "diffusers", fake_diffusers) + monkeypatch.setattr( + "core.inference.diffusion_krea2.load_krea2_tokenizer", + lambda repo_id, hf_token = None: SimpleNamespace(tag = "tokenizer"), + ) + monkeypatch.setattr( + "core.inference.diffusion_krea2.load_krea2_text_encoder", + lambda repo_id, dtype, hf_token = None: SimpleNamespace(tag = "text_encoder"), + ) + + pipe = load_krea2_pipeline(str(tmp_path), "bf16") + + # Turbo's fixed-mu schedule rides on is_distilled; dropping any of these would + # silently degrade generations, so the ctor kwargs are asserted exactly. + assert captured["pipeline"]["is_distilled"] is True + assert captured["pipeline"]["patch_size"] == 2 + assert captured["pipeline"]["text_encoder_select_layers"] == [2, 5, 8] + assert pipe.transformer.tag == "transformer" + # A prebuilt transformer (single-file/quant path) must be used as-is. + prebuilt = SimpleNamespace(tag = "prebuilt") + pipe = load_krea2_pipeline(str(tmp_path), "bf16", transformer = prebuilt) + assert pipe.transformer is prebuilt + + +# ── registry / trust / int8 exclusion wiring ───────────────────────────────── + + +def test_load_krea2_pipeline_requires_krea_capable_diffusers(monkeypatch): + # On diffusers < 0.39 (no Krea2Pipeline) the loader must fail fast with the upgrade + # hint instead of dying with a bare AttributeError mid-load. + import pytest + + fake = SimpleNamespace(__version__ = "0.38.0") + monkeypatch.setitem(sys.modules, "diffusers", fake) + with pytest.raises(RuntimeError, match = "0.39"): + load_krea2_pipeline("krea/Krea-2-Turbo", "bf16") + + +def test_krea2_family_wiring(): + from core.inference.diffusion import _is_trusted_diffusion_repo + from core.inference.diffusion_families import ( + default_generation_params, + detect_family, + family_sd_cpp_supported, + ) + from core.inference.diffusion_transformer_quant import TQ_INT8, exclude_tokens_for_scheme + + fam = detect_family("krea/Krea-2-Turbo") + assert fam is not None and fam.name == KREA2_FAMILY_NAME + # Both vendor repos are non-GGUF allowlisted (Turbo for inference, Raw for training); + # no sd.cpp mapping -> diffusers fallback. + assert _is_trusted_diffusion_repo("krea/Krea-2-Turbo") + assert _is_trusted_diffusion_repo("krea/Krea-2-Raw") + assert not family_sd_cpp_supported(fam) + # Krea2TimestepEmbedding runs at M = batch; int8 (torch._int_mm, M > 16) must skip it. + assert "time_embed" in exclude_tokens_for_scheme(TQ_INT8) + # Adapters train on Raw but run on Turbo, so the family carries a deploy override. + assert fam.deploy_base_repo == "krea/Krea-2-Turbo" + # The OpenAI /v1/images/generations route reads (steps, guidance) from this table; Krea + # Turbo is distilled (8 steps, no CFG), matching the Create UI seed instead of the + # generic (9, 0.0) fallback. Raw is the undistilled base (also inference-loadable) and runs + # its full 52-step / CFG 3.5 recipe, so its more specific key must win over the "krea" one. + assert default_generation_params("krea/Krea-2-Turbo") == (8, 0.0) + assert default_generation_params("krea/Krea-2-Raw") == (52, 3.5) + + +# ── training wiring ────────────────────────────────────────────────────────── + + +def test_krea2_training_registry(): + from core.inference.diffusion_families import trainable_family_names + from core.training.diffusion_train_common import ( + family_train_infos, + get_trainer, + train_defaults, + ) + from core.training.diffusion_dit_trainer import run_dit_lora_training + + assert "krea-2" in trainable_family_names() + assert get_trainer("krea-2") is run_dit_lora_training + # The Krea 2 authors' recommended starting point (their reference script defaults). + assert train_defaults("krea-2") == { + "lora_rank": 32, + "learning_rate": 3e-4, + "resolution": 512, + } + info = {i["name"]: i for i in family_train_infos()}["krea-2"] + # Krea's guidance: train LoRAs on the undistilled Raw model, run them on Turbo, so + # Raw leads the training bases while Turbo stays available. + assert info["default_base"] == "krea/Krea-2-Raw" + assert info["base_repos"] == ["krea/Krea-2-Raw", "krea/Krea-2-Turbo"] + assert info["supports_compile"] is True + # Deploy previews the adapter on Turbo, not the Raw checkpoint it trained on, so the UI + # loads the distilled inference recipe; other families leave this None. + assert info["deploy_base"] == "krea/Krea-2-Turbo" + assert {i["name"]: i for i in family_train_infos()}["flux.1"]["deploy_base"] is None + + +def test_krea2_spec_registered_with_authors_targets(): + from core.training.diffusion_dit_trainer import _KREA2_TARGETS, _SPECS + + spec = _SPECS["krea-2"] + assert spec.force_bf16 is True + assert spec.lora_targets == _KREA2_TARGETS + # The authors' full recommended set: attention + SwiGLU + text fusion + embedders. + for t in ("to_q", "to_gate", "ff.up", "text_fusion.projector", "time_mod_proj"): + assert t in _KREA2_TARGETS + + +def test_krea2_collate_and_forward_roundtrip(): + # spec.forward imports Krea2Pipeline (prepare_position_ids), so this needs a real + # diffusers install; CI hosts run the backend suite without one. + pytest.importorskip("diffusers") + import torch + from core.training.diffusion_dit_trainer import _SPECS + + spec = _SPECS["krea-2"] + # Two fixed-length embed entries collate to a plain concat with the mask batched. + entries = [ + (torch.randn(1, 8, 12, 16), torch.ones(1, 8, dtype = torch.int64)), + (torch.randn(1, 8, 12, 16), torch.ones(1, 8, dtype = torch.int64)), + ] + pe_b, mask_b = spec.collate(entries, "cpu", torch.float32) + assert pe_b.shape == (2, 8, 12, 16) + assert mask_b.shape == (2, 8) + + captured = {} + + class _FakeTransformer: + def __call__(self, **kwargs): + captured.update(kwargs) + # Echo the packed sequence: unpack(pack(x)) == x proves the inlined + # packing mirrors Krea2Pipeline exactly (they are mutual inverses). + return (kwargs["hidden_states"],) + + noisy = torch.randn(2, 16, 1, 8, 8) + timesteps = torch.tensor([250.0, 750.0]) + pred = spec.forward( + _FakeTransformer(), noisy, timesteps, None, (pe_b, mask_b), None, "cpu", torch.float32 + ) + assert torch.equal(pred, noisy) + # [B, (H/2)*(W/2), C*4] patches, one shared [(txt+img), 3] position grid, and the + # [0, 1] timestep convention. + assert captured["hidden_states"].shape == (2, 16, 64) + assert captured["position_ids"].shape == (8 + 16, 3) + assert torch.allclose(captured["timestep"], torch.tensor([0.25, 0.75])) + assert captured["encoder_attention_mask"] is mask_b diff --git a/studio/frontend/src/features/images/api.ts b/studio/frontend/src/features/images/api.ts index eb8ecd785e..610e5ec935 100644 --- a/studio/frontend/src/features/images/api.ts +++ b/studio/frontend/src/features/images/api.ts @@ -443,6 +443,10 @@ export interface DiffusionTrainableFamily { recommended_precision?: string; // Whether the family's transformer can be torch.compile'd (gates the Speed > Compile row). supports_compile?: boolean; + // When set, deploying a LoRA trained on this family previews it on this repo instead of + // the checkpoint it was trained on (Krea trains on Raw but runs adapters on Turbo). Null + // for families that deploy on the base they trained on. + deploy_base?: string | null; } // Where diffusion training reads/writes on this Studio, plus usable dataset folders. diff --git a/studio/frontend/src/features/images/images-page.tsx b/studio/frontend/src/features/images/images-page.tsx index c216ba1df4..800ba0a0d9 100644 --- a/studio/frontend/src/features/images/images-page.tsx +++ b/studio/frontend/src/features/images/images-page.tsx @@ -98,6 +98,8 @@ const editGguf = (id: string, name: string): ModelOption => ({ type SafetensorsSpec = { kind: "pipeline" | "single_file"; filename?: string }; const SAFETENSORS_MODELS: Record = { "unsloth/Z-Image-Turbo-unsloth-bnb-4bit": { kind: "pipeline" }, + // Krea 2 Turbo: official vendor repo (bf16 pipeline), on the backend allowlist. + "krea/Krea-2-Turbo": { kind: "pipeline" }, "unsloth/Qwen-Image-2512-unsloth-bnb-4bit": { kind: "pipeline" }, "unsloth/Qwen-Image-2512-FP8": { kind: "single_file", @@ -132,6 +134,7 @@ const MODELS: ModelOption[] = [ "Z-Image-Turbo (bnb-4bit)", "Safetensors · bnb-4bit", ), + safetensors("krea/Krea-2-Turbo", "Krea 2 Turbo", "Safetensors · bf16"), safetensors( "unsloth/Qwen-Image-2512-unsloth-bnb-4bit", "Qwen-Image 2512 (bnb-4bit)", @@ -208,6 +211,13 @@ const DEFAULT_GEN = { steps: 9, guidance: 0 }; const MODEL_DEFAULTS: Array<{ match: string; steps: number; guidance: number }> = [ { match: "z-image-turbo", steps: 9, guidance: 0 }, + // Krea 2 Raw is the undistilled base (also inference-loadable): its card runs 52 steps at + // guidance 3.5, so it must precede the distilled "krea-2" key below or a Raw load would run + // the 8-step recipe and produce garbage. + { match: "krea-2-raw", steps: 52, guidance: 3.5 }, + // Krea 2 Turbo is distilled (TDM): 8 steps, no CFG. "krea-2" then covers Turbo (and any + // other krea id) but Raw, which is matched more specifically above. + { match: "krea-2", steps: 8, guidance: 0 }, { match: "flux.1-schnell", steps: 4, guidance: 0 }, // Kontext (editing) before the generic flux.1: ~28 steps, lower guidance (~2.5). { match: "kontext", steps: 28, guidance: 2.5 }, diff --git a/studio/frontend/src/features/images/train/diffusion-train-panel.tsx b/studio/frontend/src/features/images/train/diffusion-train-panel.tsx index 0491e7cc6e..bf6e1b7abf 100644 --- a/studio/frontend/src/features/images/train/diffusion-train-panel.tsx +++ b/studio/frontend/src/features/images/train/diffusion-train-panel.tsx @@ -723,23 +723,38 @@ export function DiffusionTrainPanel({ [poll], ); + // Resolve the repo an adapter should be PREVIEWED on. Krea (and any family that trains on + // one checkpoint but runs adapters on another) declares a deploy_base: preview the adapter + // there instead of the training checkpoint, so the default Krea train-on-Raw flow does not + // load the adapter on Raw's non-distilled recipe. Only a recognised training base is + // overridden; a custom repo the user typed is respected as-is. + const deployBaseFor = useCallback( + (trainedBase: string, famName: string): string => { + const rec = info?.families?.find((f) => f.name === famName); + if (rec?.deploy_base && rec.base_repos.includes(trainedBase)) return rec.deploy_base; + return trainedBase; + }, + [info?.families], + ); + const onDeployClick = useCallback(() => { if (!status?.catalog_path) { toast.error("The trained adapter is not available yet."); return; } - const baseRepo = status.base_model || (effectiveBase === CUSTOM_BASE ? customBase : effectiveBase); - if (!baseRepo) { + const trainedBase = status.base_model || (effectiveBase === CUSTOM_BASE ? customBase : effectiveBase); + if (!trainedBase) { toast.error("Could not determine the base model to load for this adapter."); return; } + const famName = status.family || family?.name || ""; onDeploy?.({ - baseRepo, - family: status.family || family?.name || "", + baseRepo: deployBaseFor(trainedBase, famName), + family: famName, catalogPath: status.catalog_path, trigger: instancePrompt.trim(), }); - }, [status, baseChoice, customBase, family, instancePrompt, onDeploy]); + }, [status, effectiveBase, customBase, family, instancePrompt, onDeploy, deployBaseFor]); const numberField = ( label: string, @@ -1217,7 +1232,7 @@ export function DiffusionTrainPanel({ size="sm" onClick={() => onDeploy?.({ - baseRepo: viewRun.base_model || "", + baseRepo: deployBaseFor(viewRun.base_model || "", viewRun.family || ""), family: viewRun.family || "", catalogPath: viewRun.catalog_path || "", trigger: viewRun.instance_prompt || "",