diff --git a/studio/backend/core/inference/diffusion.py b/studio/backend/core/inference/diffusion.py index 6cbf6d569c..e06c52b132 100644 --- a/studio/backend/core/inference/diffusion.py +++ b/studio/backend/core/inference/diffusion.py @@ -1716,7 +1716,8 @@ class DiffusionBackend: ) if transformer is not None: pipe = self._assemble_pipe( - pipeline_cls, base, transformer, dtype, hf_token, device, base_local_dir + pipeline_cls, base, transformer, dtype, hf_token, device, base_local_dir, + fam = fam, ) return pipe, scheme @@ -1731,7 +1732,7 @@ class DiffusionBackend: base, subfolder = "transformer", torch_dtype = dtype, token = hf_token ) pipe = self._assemble_pipe( - pipeline_cls, base, transformer, dtype, hf_token, device, base_local_dir + pipeline_cls, base, transformer, dtype, hf_token, device, base_local_dir, fam = fam ) scheme = quantize_transformer( pipe, @@ -1754,9 +1755,19 @@ class DiffusionBackend: hf_token: Optional[str], device: str, base_local_dir: Optional[str] = None, + fam: Optional[DiffusionFamily] = None, ) -> Any: """Assemble the diffusers pipeline around ``transformer`` and place it on ``device`` (a no-op for an already-placed pre-quantized transformer; it moves the companions).""" + if getattr(fam, "name", None) == KREA2_FAMILY_NAME: + # krea ships transformers-5.x configs and no top-level tokenizer files, so + # Pipeline.from_pretrained dies in the tokenizer (vocab_file = None); assemble + # per-component like every other krea load path (see diffusion_krea2.py). + pipe = load_krea2_pipeline( + base_local_dir or base, dtype, hf_token = hf_token, transformer = transformer + ) + pipe.to(device) + return pipe pipe_kwargs: dict[str, Any] = {"torch_dtype": dtype, "transformer": transformer} if hf_token: pipe_kwargs["token"] = hf_token diff --git a/studio/backend/tests/test_diffusion_backend.py b/studio/backend/tests/test_diffusion_backend.py index ad791f80dc..916ac28382 100644 --- a/studio/backend/tests/test_diffusion_backend.py +++ b/studio/backend/tests/test_diffusion_backend.py @@ -2681,6 +2681,45 @@ def test_dense_quant_prequant_proceeds_but_forbids_dense_fallback(fake_runtime, assert attempted == [False] # ...fast path still attempted, dense fallback forbidden +def test_assemble_pipe_routes_krea2_per_component(monkeypatch): + # krea's repo ships transformers-5.x configs and no top-level tokenizer files, so + # Pipeline.from_pretrained dies in the tokenizer (vocab_file = None). The quant fast + # path must assemble per-component via load_krea2_pipeline like every other krea load. + from core.inference import diffusion as dmod + + calls: dict = {} + + class Pipe: + def to(self, device): + calls["device"] = device + return self + + def fake_loader(base, dtype, hf_token = None, transformer = None): + calls["base"] = base + calls["transformer"] = transformer + return Pipe() + + monkeypatch.setattr(dmod, "load_krea2_pipeline", fake_loader) + + class ExplodingPipeline: + @staticmethod + def from_pretrained(*a, **k): + raise AssertionError("krea-2 must not go through Pipeline.from_pretrained") + + marker = object() + pipe = dmod.DiffusionBackend._assemble_pipe( + ExplodingPipeline, + "krea/Krea-2-Turbo", + marker, + "bf16", + None, + "cuda:0", + fam = types.SimpleNamespace(name = "krea-2"), + ) + assert isinstance(pipe, Pipe) + assert calls == {"base": "krea/Krea-2-Turbo", "transformer": marker, "device": "cuda:0"} + + def test_dense_quant_unusable_prequant_path_runs_dense_refit(fake_runtime, tmp_path, monkeypatch): # A request-supplied transformer_prequant_path the loader refuses (missing, or outside # UNSLOTH_ALLOW_LOCAL_PREQUANT_PATH) resolves to NO usable prequant source, so the