Skip to content

Commit 4b575f6

Browse files
authored
Merge pull request #3186 from bghira/bugfix/3184
Fix adapter-only validation adapter runs
2 parents 096534d + fc89166 commit 4b575f6

2 files changed

Lines changed: 54 additions & 28 deletions

File tree

simpletuner/helpers/training/validation_adapters.py

Lines changed: 8 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -235,20 +235,19 @@ def _make_run(label: str | None, specs: Sequence[ValidationAdapterSpec]) -> Vali
235235
is_base=False,
236236
)
237237

238-
if adapter_path and mode != "none":
238+
if mode != "none" and adapter_path:
239239
specs = [_build_adapter_spec(adapter_path, adapter_strength, adapter_name)]
240240
preferred_label = adapter_name or _stem_from_path(adapter_path)
241241
runs.append(_make_run(preferred_label, specs))
242242

243-
for entry in _iter_config_entries(adapter_config):
244-
label, specs = _normalize_run_entry(entry)
245-
if not specs:
246-
continue
247-
runs.append(_make_run(label, specs))
243+
if mode != "none":
244+
for entry in _iter_config_entries(adapter_config):
245+
label, specs = _normalize_run_entry(entry)
246+
if not specs:
247+
continue
248+
runs.append(_make_run(label, specs))
248249

249-
include_base = True
250-
if adapter_path and mode == "adapter_only" and adapter_config in (None, [], {}):
251-
include_base = False
250+
include_base = mode == "comparison" or not runs
252251

253252
ordered_runs: List[ValidationAdapterRun] = []
254253
if include_base:

tests/test_validation_adapters.py

Lines changed: 46 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -48,22 +48,20 @@ def test_config_supports_multiple_runs(self):
4848
},
4949
]
5050
runs = build_validation_adapter_runs(None, config)
51-
self.assertEqual(3, len(runs))
51+
self.assertEqual(2, len(runs))
5252

5353
first = runs[0]
54-
self.assertTrue(first.is_base)
54+
self.assertFalse(first.is_base)
55+
self.assertEqual("base_adapter", first.slug)
56+
self.assertEqual("org/base_adapter", first.adapters[0].repo_id)
5557

5658
second = runs[1]
57-
self.assertEqual("base_adapter", second.slug)
58-
self.assertEqual("org/base_adapter", second.adapters[0].repo_id)
59-
60-
third = runs[2]
61-
self.assertEqual("combo", third.label)
62-
self.assertEqual(2, len(third.adapters))
63-
strengths = [adapter.strength for adapter in third.adapters]
59+
self.assertEqual("combo", second.label)
60+
self.assertEqual(2, len(second.adapters))
61+
strengths = [adapter.strength for adapter in second.adapters]
6462
self.assertEqual([0.5, 1.0], strengths)
65-
self.assertEqual("repo/style", third.adapters[1].repo_id)
66-
self.assertEqual("style.safetensors", third.adapters[1].weight_name)
63+
self.assertEqual("repo/style", second.adapters[1].repo_id)
64+
self.assertEqual("style.safetensors", second.adapters[1].weight_name)
6765

6866
def test_adapter_mode_comparison_includes_base(self):
6967
runs = build_validation_adapter_runs(
@@ -76,11 +74,40 @@ def test_adapter_mode_comparison_includes_base(self):
7674
self.assertEqual("sample", second.adapters[0].adapter_name)
7775
self.assertEqual(0.8, second.adapters[0].strength)
7876

77+
def test_adapter_mode_comparison_with_config_includes_base(self):
78+
runs = build_validation_adapter_runs(
79+
None,
80+
[{"label": "custom", "path": "repo/hero"}],
81+
adapter_mode="comparison",
82+
)
83+
self.assertEqual(2, len(runs))
84+
self.assertTrue(runs[0].is_base)
85+
self.assertEqual("custom", runs[1].slug)
86+
87+
def test_adapter_mode_adapter_only_with_config_excludes_base(self):
88+
runs = build_validation_adapter_runs(
89+
None,
90+
[{"label": "custom", "path": "repo/hero"}],
91+
adapter_mode="adapter_only",
92+
)
93+
self.assertEqual(1, len(runs))
94+
self.assertFalse(runs[0].is_base)
95+
self.assertEqual("custom", runs[0].slug)
96+
7997
def test_adapter_mode_none_skips_loading(self):
8098
runs = build_validation_adapter_runs("foo/bar", None, adapter_mode="none")
8199
self.assertEqual(1, len(runs))
82100
self.assertTrue(runs[0].is_base)
83101

102+
def test_adapter_mode_none_skips_config_loading(self):
103+
runs = build_validation_adapter_runs(
104+
None,
105+
[{"label": "custom", "path": "repo/hero"}],
106+
adapter_mode="none",
107+
)
108+
self.assertEqual(1, len(runs))
109+
self.assertTrue(runs[0].is_base)
110+
84111
def test_config_entry_with_strength_and_name(self):
85112
config = [
86113
{
@@ -91,8 +118,8 @@ def test_config_entry_with_strength_and_name(self):
91118
}
92119
]
93120
runs = build_validation_adapter_runs(None, config)
94-
self.assertEqual(2, len(runs))
95-
run = runs[1]
121+
self.assertEqual(1, len(runs))
122+
run = runs[0]
96123
self.assertEqual("custom", run.label)
97124
adapter = run.adapters[0]
98125
self.assertEqual("hero_adapter", adapter.adapter_name)
@@ -111,7 +138,7 @@ def test_config_entry_supports_run_level_target_stage(self):
111138
]
112139
runs = build_validation_adapter_runs(None, config)
113140

114-
run = runs[1]
141+
run = runs[0]
115142
self.assertEqual("repo/detail", run.adapters[0].repo_id)
116143
self.assertEqual("two", run.adapters[0].target_stage)
117144
self.assertEqual("one", run.adapters[1].target_stage)
@@ -246,7 +273,7 @@ def test_targeted_adapter_loads_only_for_matching_stage(self):
246273
run = build_validation_adapter_runs(
247274
None,
248275
[{"label": "refiner", "path": "repo/refiner", "target_stage": "two", "strength": 0.4}],
249-
)[1]
276+
)[0]
250277

251278
with validator._temporary_validation_adapters(run):
252279
self.assertEqual([], base_pipeline.load_calls)
@@ -266,7 +293,7 @@ def test_wan_low_stage_uses_transformer_2_component(self):
266293
run = build_validation_adapter_runs(
267294
None,
268295
[{"label": "low-noise", "path": "repo/low", "target_stage": "low", "strength": 0.6}],
269-
)[1]
296+
)[0]
270297

271298
with validator._temporary_validation_adapters(run):
272299
with validator._temporary_validation_stage_adapters(pipeline, ("high", "low")):
@@ -291,7 +318,7 @@ def test_stage_adapter_preserves_global_adapter_on_component_target(self):
291318
],
292319
}
293320
],
294-
)[1]
321+
)[0]
295322

296323
with validator._temporary_validation_adapters(run):
297324
self.assertEqual(("combo", 0.25), pipeline.set_calls[0])
@@ -314,7 +341,7 @@ def test_validation_adapter_restore_ignores_loaded_temporary_adapter(self):
314341
run = build_validation_adapter_runs(
315342
None,
316343
[{"label": "Krea2 Turbo adapter", "path": "repo/krea2", "adapter_name": "krea2_turbo_adapter"}],
317-
)[1]
344+
)[0]
318345

319346
with validator._temporary_validation_adapters(run):
320347
self.assertEqual("repo/krea2", pipeline.load_calls[0][0])
@@ -332,7 +359,7 @@ def test_validation_adapter_restore_preserves_existing_active_adapter(self):
332359
run = build_validation_adapter_runs(
333360
None,
334361
[{"label": "Krea2 Turbo adapter", "path": "repo/krea2", "adapter_name": "krea2_turbo_adapter"}],
335-
)[1]
362+
)[0]
336363

337364
with validator._temporary_validation_adapters(run):
338365
self.assertEqual(

0 commit comments

Comments
 (0)