diff --git a/studio/backend/core/inference/diffusion_krea2.py b/studio/backend/core/inference/diffusion_krea2.py index 312a4c8aa2..65619ccd47 100644 --- a/studio/backend/core/inference/diffusion_krea2.py +++ b/studio/backend/core/inference/diffusion_krea2.py @@ -62,14 +62,16 @@ def remap_rope_parameters(text_config) -> None: 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" - } + 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): +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 diff --git a/studio/backend/tests/test_diffusion_krea2.py b/studio/backend/tests/test_diffusion_krea2.py index b360ae43f9..4cface12cd 100644 --- a/studio/backend/tests/test_diffusion_krea2.py +++ b/studio/backend/tests/test_diffusion_krea2.py @@ -55,9 +55,7 @@ def test_remap_rope_parameters_noop_on_5x_runtime_or_plain_4x_config(): 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}) - ) + (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}