File size: 6,934 Bytes
a30449a 4511070 a30449a 4511070 a30449a b990348 a30449a b990348 a30449a 4511070 b990348 a30449a | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 | 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() |