ideogram4 / tests /test_lora_utils.py
Vlad Iliescu
better lora
b990348
Raw
History Blame Contribute Delete
6.93 kB
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()