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()