Merge remote-tracking branch 'origin/diffusion-krea2' into diffusion-krea2

This commit is contained in:
Daniel Han 2026-07-03 14:55:22 +00:00
commit 7f9ec4f73e
2 changed files with 7 additions and 7 deletions

View file

@ -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

View file

@ -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}