| import unittest |
| from unittest import mock |
|
|
| from lora_utils import ensure_loras_loaded, parse_adapter_specs |
|
|
|
|
| class DummyAdapterHost: |
| def __init__(self): |
| self.loaded = [] |
| self.active_names = None |
| self.active_weights = None |
|
|
| def load_lora_adapter(self, state_dict, adapter_name, prefix=None): |
| self.loaded.append((state_dict, adapter_name, prefix)) |
|
|
| def set_adapters(self, adapter_names, adapter_weights=None): |
| self.active_names = adapter_names |
| self.active_weights = adapter_weights |
|
|
| def get_list_adapters(self): |
| return [adapter_name for _, adapter_name, _ in self.loaded] |
|
|
| def delete_adapters(self, adapter_name): |
| self.loaded = [entry for entry in self.loaded if entry[1] != adapter_name] |
|
|
|
|
| class DummyIdeogramPipeline: |
| def __init__(self): |
| self.transformer = DummyAdapterHost() |
| self.unconditional_transformer = DummyAdapterHost() |
|
|
|
|
| class DummyPeftLayer: |
| def __init__(self): |
| self.lora_A = {} |
| self.active_names = None |
| self.scales = [] |
| self.enabled = None |
| self.deleted = [] |
|
|
| def set_adapter(self, adapter_names): |
| self.active_names = adapter_names |
|
|
| def set_scale(self, adapter_name, scale): |
| self.scales.append((adapter_name, scale)) |
|
|
| def enable_adapters(self, enabled): |
| self.enabled = enabled |
|
|
| def delete_adapter(self, adapter_name): |
| self.deleted.append(adapter_name) |
| self.lora_A.pop(adapter_name, None) |
|
|
|
|
| class DummyPeftHost: |
| def __init__(self): |
| self.peft_config = {} |
| self.layer = DummyPeftLayer() |
|
|
| def modules(self): |
| return [self, self.layer] |
|
|
| def named_modules(self): |
| return [("", self), ("layer", self.layer)] |
|
|
|
|
| class DummyPeftOnlyIdeogramPipeline: |
| def __init__(self): |
| self.transformer = DummyPeftHost() |
| self.unconditional_transformer = DummyPeftHost() |
|
|
|
|
| class ParseAdapterSpecsTest(unittest.TestCase): |
| def test_blank_spec_returns_no_entries(self): |
| self.assertEqual(parse_adapter_specs("\n \n", 1.0), []) |
|
|
| def test_parses_repo_weight_and_default_scale(self): |
| entries = parse_adapter_specs("vladi/real-loras:ema4_flux2_klein_9b_000002000.safetensors", 1.0) |
|
|
| self.assertEqual(entries[0]["repo_id"], "vladi/real-loras") |
| self.assertEqual(entries[0]["weight_name"], "ema4_flux2_klein_9b_000002000.safetensors") |
| self.assertEqual(entries[0]["scale"], 1.0) |
|
|
| def test_parses_multiple_entries_and_inline_scale(self): |
| entries = parse_adapter_specs( |
| "vladi/real-loras:ema4_flux2_klein_9b_000002000.safetensors\n" |
| "vladi/loras:klein9b/klein_snofs_v1_4.safetensors@0.8", |
| 0.5, |
| ) |
|
|
| self.assertEqual(len(entries), 2) |
| self.assertEqual(entries[1]["repo_id"], "vladi/loras") |
| self.assertEqual(entries[1]["weight_name"], "klein9b/klein_snofs_v1_4.safetensors") |
| self.assertEqual(entries[1]["scale"], 0.4) |
|
|
| def test_parses_huggingface_url_with_subfolder(self): |
| entries = parse_adapter_specs( |
| "https://huggingface.co/vladi/loras/blob/dev/klein9b/klein_snofs_v1_4.safetensors@0.8", |
| 1.0, |
| ) |
|
|
| self.assertEqual(entries[0]["repo_id"], "vladi/loras") |
| self.assertEqual(entries[0]["weight_name"], "klein9b/klein_snofs_v1_4.safetensors") |
| self.assertEqual(entries[0]["revision"], "dev") |
| self.assertEqual(entries[0]["scale"], 0.8) |
|
|
| def test_rejects_missing_weight_name(self): |
| with self.assertRaises(ValueError): |
| parse_adapter_specs("vladi/real-loras", 1.0) |
|
|
| def test_rejects_duplicate_entries(self): |
| with self.assertRaises(ValueError): |
| parse_adapter_specs( |
| "vladi/real-loras:ema4_flux2_klein_9b_000002000.safetensors\n" |
| "vladi/real-loras:ema4_flux2_klein_9b_000002000.safetensors@0.8", |
| 1.0, |
| ) |
|
|
| def test_loads_lora_without_pipeline_loader(self): |
| pipe = DummyIdeogramPipeline() |
| active_by_key = {} |
| state_dict = {"transformer_blocks.0.attn.to_q.lora_A.weight": object()} |
|
|
| with mock.patch("lora_utils._download_lora_weight", return_value="/tmp/adapter.safetensors"), mock.patch( |
| "lora_utils._load_adapter_state_dict", return_value=state_dict |
| ): |
| entries = ensure_loras_loaded( |
| pipe, |
| "vladi/real-loras:ema4_flux2_klein_9b_000002000.safetensors", |
| 0.8, |
| active_by_key, |
| token=None, |
| ) |
|
|
| adapter_name = entries[0]["adapter_name"] |
| self.assertEqual(pipe.transformer.loaded[0], (state_dict, adapter_name, None)) |
| self.assertEqual(pipe.unconditional_transformer.loaded[0], (state_dict, adapter_name, None)) |
| self.assertEqual(pipe.transformer.active_names, [adapter_name]) |
| self.assertEqual(pipe.unconditional_transformer.active_names, [adapter_name]) |
| self.assertEqual(pipe.transformer.active_weights, [0.8]) |
| self.assertEqual(pipe.unconditional_transformer.active_weights, [0.8]) |
|
|
| def test_loads_lora_with_peft_fallback_when_transformers_have_no_loader(self): |
| pipe = DummyPeftOnlyIdeogramPipeline() |
| active_by_key = {} |
| state_dict = {"transformer.transformer_blocks.0.attn.to_q.lora_A.weight": object()} |
| loaded = [] |
|
|
| def fake_load_with_peft(host, host_state_dict, adapter_name): |
| loaded.append((host, host_state_dict, adapter_name)) |
| host.peft_config[adapter_name] = object() |
| host.layer.lora_A[adapter_name] = object() |
|
|
| with mock.patch("lora_utils._download_lora_weight", return_value="/tmp/adapter.safetensors"), mock.patch( |
| "lora_utils._load_adapter_state_dict", return_value=state_dict |
| ), mock.patch("lora_utils._load_lora_with_peft", side_effect=fake_load_with_peft): |
| entries = ensure_loras_loaded( |
| pipe, |
| "https://huggingface.co/vladi/loras/blob/dev/klein9b/klein_snofs_v1_4.safetensors", |
| 0.7, |
| active_by_key, |
| token="secret", |
| ) |
|
|
| adapter_name = entries[0]["adapter_name"] |
| expected_state_dict = {"transformer_blocks.0.attn.to_q.lora_A.weight": state_dict[next(iter(state_dict))]} |
| self.assertEqual(loaded[0], (pipe.transformer, expected_state_dict, adapter_name)) |
| self.assertEqual(loaded[1], (pipe.unconditional_transformer, expected_state_dict, adapter_name)) |
| self.assertEqual(pipe.transformer.layer.active_names, [adapter_name]) |
| self.assertEqual(pipe.unconditional_transformer.layer.active_names, [adapter_name]) |
| self.assertEqual(pipe.transformer.layer.scales, [(adapter_name, 0.7)]) |
| self.assertEqual(pipe.unconditional_transformer.layer.scales, [(adapter_name, 0.7)]) |
|
|
|
|
| if __name__ == "__main__": |
| unittest.main() |