diff --git a/modules/lora/lora_stack.py b/modules/lora/lora_stack.py index b0fa08910..544035e36 100644 --- a/modules/lora/lora_stack.py +++ b/modules/lora/lora_stack.py @@ -33,6 +33,7 @@ SAMPLE_CAP = 1 << 22 # strided subsample bound for magnitude quantiles (full-siz state: dict = {'entries': {}, 'flips': {}, 'gamma': 1.0, 'gamma_e': 1.0, 'total_steps': 0, 'finalized': False, 'reported': None} warned: set = set() +warned_context = None def mode(): @@ -60,7 +61,17 @@ def signature(): return 'sum' +def warn_context(): + """Settings the degradation warnings below speak about.""" + return (signature(), getattr(shared.opts, 'diffusers_offload_mode', ''), int(getattr(shared.opts, 'lora_sdnq_host_rank', 0) or 0), getattr(shared.opts, 'sd_model_checkpoint', '')) + + def warn_once(key, message): + global warned_context # pylint: disable=global-statement + context = warn_context() + if context != warned_context: + warned.clear() # what was reported under the old settings says nothing about the new ones + warned_context = context if key not in warned: warned.add(key) log.warning(message) diff --git a/test/test-sdnq-lora-factors.py b/test/test-sdnq-lora-factors.py index d4bbb9a2c..191429890 100644 --- a/test/test-sdnq-lora-factors.py +++ b/test/test-sdnq-lora-factors.py @@ -2071,6 +2071,40 @@ def test_select_host_disabled_falls_back_to_sum(): return True +def test_degradation_warning_rearms_on_settings_change(): + class CountingLog: + def __init__(self): + self.warnings = 0 + + def warning(self, _message): + self.warnings += 1 + + counter = CountingLog() + real_log = lora_stack.log + saved = {k: getattr(shared.opts, k, None) for k in ('lora_stack_mode', 'lora_sdnq_host_rank')} + lora_stack.log = counter + lora_stack.warned.clear() + lora_stack.warned_context = None + try: + shared.opts.lora_stack_mode = 'klora' + lora_stack.warn_once('probe', 'Network stack: probe') + lora_stack.warn_once('probe', 'Network stack: probe') + assert counter.warnings == 1, f'one degradation under one settings context says it once, got {counter.warnings}' + shared.opts.lora_stack_mode = 'estlora' + lora_stack.warn_once('probe', 'Network stack: probe') + assert counter.warnings == 2, 'changing the stack mode must let the degradation be said again' + shared.opts.lora_sdnq_host_rank = 0 + lora_stack.warn_once('probe', 'Network stack: probe') + assert counter.warnings == 3, 'the host rank belongs to that context too' + finally: + lora_stack.log = real_log + for k, v in saved.items(): + setattr(shared.opts, k, v) + lora_stack.warned.clear() + lora_stack.warned_context = None + return True + + def test_flip_lands_before_crossover_step(): layer = build_layer('uint4') n1, n2, D1, D2 = select_pair(layer, seed0=55, seed1=56) @@ -2907,7 +2941,7 @@ def run_tests(): test_select_deactivate_from_midflip, test_select_requires_exactly_two_nets, test_select_gated_off_when_compiled, test_select_finalize_drops_dead_module, test_select_int8_pair_rides_segments, test_select_gate_dormant_without_pair, test_stale_schedule_dropped_on_reapply, test_est_energy_matches_full_frobenius, test_select_weight_kind_plain_layer, - test_select_gamma_tracks_live_entries, test_select_host_disabled_falls_back_to_sum, test_flip_lands_before_crossover_step, + test_select_gamma_tracks_live_entries, test_select_host_disabled_falls_back_to_sum, test_degradation_warning_rearms_on_settings_change, test_flip_lands_before_crossover_step, test_score_pair_chunked_precision, test_select_replay_from_cache_skips_calc, test_select_weight_replay_from_cache_skips_calc, test_select_reset_reports_timing, test_select_weight_flip_calcs_on_accelerator]: run_test(CAT_SELECT, fn)