diff --git a/pipelines/generic_transformer.py b/pipelines/generic_transformer.py index 179c124fc..eec297b7f 100644 --- a/pipelines/generic_transformer.py +++ b/pipelines/generic_transformer.py @@ -32,9 +32,28 @@ def load_transformer(repo_id, cls_name, load_config=None, subfolder="transformer quant_type = model_quant.get_quant_type(quant_args) dtype = dtype or devices.dtype + def load_from_repo(): + nonlocal quant_args + log.debug(f'Load model: transformer="{repo_id}" cls={cls_name.__name__} subfolder={subfolder} quant="{quant_type}" loader={get_loader("diffusers")} args={load_args}') + if 'sdnq-' in repo_id.lower(): + quant_args = {} + if dtype is not None: + load_args['torch_dtype'] = dtype + if subfolder is not None: + load_args['subfolder'] = subfolder + if variant is not None: + load_args['variant'] = variant + return cls_name.from_pretrained( + repo_id, + cache_dir=shared.opts.hfcache_dir, + **load_args, + **quant_args, + **kwargs, + ) + local_file = None + from modules import sd_unet if shared.opts.sd_unet is not None and shared.opts.sd_unet != 'Default': - from modules import sd_unet if shared.opts.sd_unet not in list(sd_unet.unet_dict): log.error(f'Load module: type=transformer file="{shared.opts.sd_unet}" not found') elif os.path.exists(sd_unet.unet_dict[shared.opts.sd_unet]): @@ -51,19 +70,25 @@ def load_transformer(repo_id, cls_name, load_config=None, subfolder="transformer elif local_file is not None and local_file.lower().endswith('.safetensors') and native_spec is not None: from pipelines import native_transformer log.debug(f'Load model: transformer="{local_file}" cls={cls_name.__name__} quant="{quant_type}" loader=native args={load_args}') - transformer, _ = native_transformer.load( - local_file, - repo_id, - native_spec, - load_config, - allow_quant=allow_quant, - dtype=dtype, - modules_to_not_convert=modules_to_not_convert, - modules_dtype_dict=modules_dtype_dict, - quant_args=quant_args, - quant_type=quant_type, - **kwargs, - ) + try: + transformer, _ = native_transformer.load( + local_file, + repo_id, + native_spec, + load_config, + allow_quant=allow_quant, + dtype=dtype, + modules_to_not_convert=modules_to_not_convert, + modules_dtype_dict=modules_dtype_dict, + quant_args=quant_args, + quant_type=quant_type, + **kwargs, + ) + except native_transformer.OverrideArchMismatch as e: + log.warning(f'Load model: transformer override="{shared.opts.sd_unet}" incompatible with cls={cls_name.__name__} ({e}); ignoring override and loading base transformer') + shared.opts.data['sd_unet'] = 'Default' + sd_unet.loaded_unet = None + transformer = load_from_repo() # 3. load safetensors with diffusers loader elif local_file is not None and local_file.lower().endswith('.safetensors'): @@ -80,24 +105,10 @@ def load_transformer(repo_id, cls_name, load_config=None, subfolder="transformer **kwargs, ) - # 4. default loading from diffusers repo + # 4. default loading from diffusers repo (also the fallback when an + # incompatible override is dropped above) else: - log.debug(f'Load model: transformer="{repo_id}" cls={cls_name.__name__} subfolder={subfolder} quant="{quant_type}" loader={get_loader("diffusers")} args={load_args}') - if 'sdnq-' in repo_id.lower(): - quant_args = {} - if dtype is not None: - load_args['torch_dtype'] = dtype - if subfolder is not None: - load_args['subfolder'] = subfolder - if variant is not None: - load_args['variant'] = variant - transformer = cls_name.from_pretrained( - repo_id, - cache_dir=shared.opts.hfcache_dir, - **load_args, - **quant_args, - **kwargs, - ) + transformer = load_from_repo() sd_models.allow_post_quant = False # we already handled it if shared.opts.diffusers_offload_mode != 'none' and transformer is not None: diff --git a/pipelines/native_transformer.py b/pipelines/native_transformer.py index b52326e5a..a8ebe0629 100644 --- a/pipelines/native_transformer.py +++ b/pipelines/native_transformer.py @@ -60,6 +60,22 @@ DEFAULT_ACCEPTABLE_MISSING: tuple[str, ...] = ( ) +class OverrideArchMismatch(Exception): + """Raised when a user-selected UNET/DiT override cannot be loaded as the + spec's transformer class: the arch-specific converter rejects its keys, or + ``load_state_dict`` reports unexpected or missing keys. + + :func:`pipelines.generic_transformer.load_transformer` catches this to drop + the override and load the base repo transformer instead, so a stale or + wrong-arch UNET selection degrades to the base model rather than crashing + inside a converter. + + A tensor shape mismatch on otherwise-matching keys is not raised as this; it + surfaces as the native ``load_state_dict`` error, which already names the + conflicting shapes, so it stays a hard load error rather than a fall back. + """ + + @dataclass(frozen=True) class SiblingSpec: """Describes a non-transformer component that may ship inline in the same @@ -507,7 +523,14 @@ def build_component( try: if converter is not None: log.debug(f'Load model: native_transformer {component_name} converter={converter.__name__} keys={len(state_dict)}') - sd = converter(state_dict) + try: + sd = converter(state_dict) + except Exception as e: + raise OverrideArchMismatch( + f"Load model: type={cls.__name__} native_transformer converter " + f"{converter.__name__} rejected the override ({type(e).__name__}: {e}); " + f"file does not look like a {cls.__name__} checkpoint" + ) from e else: sd = state_dict @@ -535,6 +558,8 @@ def build_component( target_dtype = dtype if dtype is not None else devices.dtype log.debug(f'Load model: native_transformer {component_name} cast dtype={target_dtype}') component = component.to(dtype=target_dtype) + except OverrideArchMismatch: + raise except Exception as e: log.error(f"Load model: native_transformer {component_name} load failed: {e}") errors.display(e, "Load") @@ -565,13 +590,14 @@ def validate_state_dict_load( unexpected: list[str], acceptable_missing: tuple[str, ...], ) -> None: - """Raise ValueError if load_state_dict produced unexpected keys or - non-acceptable missing keys. Buffer-only missing keys matching the + """Raise :class:`OverrideArchMismatch` if load_state_dict produced + unexpected keys or non-acceptable missing keys, which means the override's + weights do not fit the target class. Buffer-only missing keys matching the ``acceptable_missing`` prefix list are logged at debug level and ignored. """ if unexpected: sample = ", ".join(unexpected[:5]) - raise ValueError( + raise OverrideArchMismatch( f"Load model: native_transformer {component_name} has {len(unexpected)} " f"unexpected keys (sample: {sample})" ) @@ -580,7 +606,7 @@ def validate_state_dict_load( ] if hard_missing: sample = ", ".join(hard_missing[:5]) - raise ValueError( + raise OverrideArchMismatch( f"Load model: native_transformer {component_name} missing " f"{len(hard_missing)} required keys (sample: {sample})" ) diff --git a/test/test-native-transformer.py b/test/test-native-transformer.py index 538ce84a5..3120b6ddd 100644 --- a/test/test-native-transformer.py +++ b/test/test-native-transformer.py @@ -273,8 +273,8 @@ def test_validate_rejects_unexpected(): unexpected=['some.junk.weight'], acceptable_missing=(), ) - raise AssertionError('expected ValueError') - except ValueError as e: + raise AssertionError('expected OverrideArchMismatch') + except nt.OverrideArchMismatch as e: assert 'unexpected' in str(e) assert 'some.junk.weight' in str(e) @@ -287,8 +287,8 @@ def test_validate_rejects_hard_missing(): unexpected=[], acceptable_missing=('rope.',), ) - raise AssertionError('expected ValueError') - except ValueError as e: + raise AssertionError('expected OverrideArchMismatch') + except nt.OverrideArchMismatch as e: msg = str(e) assert 'missing' in msg assert 'layers.0.weight' in msg @@ -684,6 +684,90 @@ def test_load_rejects_non_safetensors(): assert '.safetensors' in str(e) +def crashing_converter(sd): + """diffusers-style layer count that blows up when the block family is + absent, mirroring convert_chroma_..._to_diffusers on a wrong-arch file.""" + return list(set(int(k.split('.')[1]) for k in sd if 'double_blocks.' in k))[-1] + + +def test_build_component_converter_crash_raises_mismatch(): + """A converter that crashes on wrong-arch keys is wrapped as + OverrideArchMismatch (chaining the original), not the raw IndexError.""" + try: + nt.build_component( + component_name='transformer', + state_dict={'blocks.0.self_attn.weight': torch.zeros(2)}, + config={'dim': 8}, + cls=MockMiniTransformer, + converter=crashing_converter, + acceptable_missing=(), + quant_args={}, + quant_type=None, + ) + raise AssertionError('expected OverrideArchMismatch') + except nt.OverrideArchMismatch as e: + assert 'MockMiniTransformer' in str(e) + assert isinstance(e.__cause__, IndexError), 'original error must be chained' + + +def test_build_component_shape_mismatch_is_hard_error(): + """A tensor shape mismatch on otherwise-matching keys stays a hard + RuntimeError (the native size-mismatch message), it is NOT converted to + OverrideArchMismatch and so does not silently fall back to base.""" + # MockMiniTransformer(dim=8) expects (8, 8) projections; feed (4, 4). + sd = { + 'in_proj.weight': torch.randn(4, 4), 'in_proj.bias': torch.zeros(4), + 'out_proj.weight': torch.randn(4, 4), 'out_proj.bias': torch.zeros(4), + } + orig_display = nt.errors.display + nt.errors.display = lambda *a, **k: None # silence the expected traceback dump + try: + nt.build_component( + component_name='transformer', state_dict=sd, config={'dim': 8}, + cls=MockMiniTransformer, converter=None, acceptable_missing=(), + quant_args={}, quant_type=None, + ) + raise AssertionError('expected RuntimeError') + except nt.OverrideArchMismatch: + raise AssertionError('shape mismatch must not be OverrideArchMismatch') from None + except RuntimeError as e: + assert 'size mismatch' in str(e).lower() + finally: + nt.errors.display = orig_display + + +def test_load_converter_crash_raises_mismatch(): + """End-to-end: a crashing converter surfaces from load() as + OverrideArchMismatch so load_transformer can drop the override.""" + fd, path = tempfile.mkstemp(suffix='.safetensors') + try: + raw = {'model.diffusion_model.blocks.0.self_attn.weight': torch.zeros(8, 8)} + write_fixture(raw, fd, path) + orig_fetch = nt.fetch_component_config + nt.fetch_component_config = lambda repo, sub: {'dim': 8} + from modules import model_quant + orig_get_dit = model_quant.get_dit_args + orig_get_qtype = model_quant.get_quant_type + model_quant.get_dit_args = lambda *a, **k: ({}, {}) + model_quant.get_quant_type = lambda *a, **k: None + try: + spec = nt.TransformerSpec(cls=MockMiniTransformer, converter=crashing_converter) + raised = False + try: + nt.load(local_file=path, repo_id='fake/repo', spec=spec, diffusers_cfg={}) + except nt.OverrideArchMismatch as e: + raised = True + assert 'MockMiniTransformer' in str(e) + assert raised, 'expected OverrideArchMismatch' + finally: + nt.fetch_component_config = orig_fetch + model_quant.get_dit_args = orig_get_dit + model_quant.get_quant_type = orig_get_qtype + finally: + if os.path.exists(path): + os.unlink(path) + + # ============================================================ # Run # ============================================================ @@ -775,6 +859,9 @@ def run_all(): test_load_end_to_end_with_sibling_partition, test_load_raises_on_missing_sibling_class, test_load_rejects_non_safetensors, + test_build_component_converter_crash_raises_mismatch, + test_build_component_shape_mismatch_is_hard_error, + test_load_converter_crash_raises_mismatch, ]: run_test(cat, fn) diff --git a/test/test-override-mismatch-api.py b/test/test-override-mismatch-api.py new file mode 100644 index 000000000..da902cb60 --- /dev/null +++ b/test/test-override-mismatch-api.py @@ -0,0 +1,215 @@ +#!/usr/bin/env python +""" +API integration tests for UNET/DiT override architecture-mismatch self-heal. + +Verifies the reactive behavior end-to-end against a running SD.Next instance: +when a UNET/DiT override does not match the base model's architecture, the load +drops the override, resets the UNET dropdown to ``Default``, and the base model +still generates, instead of crashing inside an arch-specific converter. + +Covers: +- GET /sdapi/v1/unets, /sdapi/v1/sd-models (discovery / sanity) +- baseline: base model with no override loads and generates +- mismatch self-heal: base model + wrong-arch override -> override dropped, + sd_unet reset to Default, generation still succeeds +- match preserved (optional): a correct-arch override stays applied + +Requires a running SD.Next instance with the relevant models on disk. Model and +UNET names are environment-specific, so pass them explicitly. Run with no model +args to just list what the server has available. + +Usage: + python test/test-override-mismatch-api.py \ + --url http://127.0.0.1:7860 \ + --base-model "Diffusers/lodestones/Chroma1-HD [0e0c60ece1]" \ + --mismatch-unet "novaOrangeAM_v15" \ + [--match-unet ""] \ + [--steps 4] +""" + +import sys +import time +import argparse +import requests +import urllib3 + +urllib3.disable_warnings(urllib3.exceptions.InsecureRequestWarning) + + +class OverrideMismatchAPITest: + """Drives the override-mismatch self-heal scenarios over the HTTP API.""" + + def __init__(self, base_url, base_model=None, mismatch_unet=None, match_unet=None, steps=4): + self.base_url = base_url.rstrip('/') + self.base_model = base_model + self.mismatch_unet = mismatch_unet + self.match_unet = match_unet + self.steps = steps + self.load_timeout = 900 + self.gen_timeout = 600 + self.passed = 0 + self.failed = 0 + self.skipped = 0 + + # ---- low-level helpers ------------------------------------------------- + + def _get(self, endpoint): + r = requests.get(f'{self.base_url}{endpoint}', timeout=60, verify=False) + r.raise_for_status() + return r.json() + + def _post(self, endpoint, data=None, params=None, timeout=60): + r = requests.post(f'{self.base_url}{endpoint}', json=data, params=params, timeout=timeout, verify=False) + return r + + def record(self, ok, name, detail=''): + tag = 'PASS' if ok else 'FAIL' + self.passed += 1 if ok else 0 + self.failed += 0 if ok else 1 + line = f' {tag}: {name}' + if detail: + line += f' ({detail})' + print(line, flush=True) + + def skip(self, name, reason): + self.skipped += 1 + print(f' SKIP: {name} ({reason})', flush=True) + + # ---- mid-level operations ---------------------------------------------- + + def set_options(self, **kwargs): + r = self._post('/sdapi/v1/options', data=kwargs, timeout=self.load_timeout) + return r.status_code == 200, (r.text[:200] if r.status_code != 200 else '') + + def reload(self, force=True): + r = self._post('/sdapi/v1/reload-checkpoint', params={'force': str(force).lower()}, timeout=self.load_timeout) + return r.status_code == 200, (r.text[:200] if r.status_code != 200 else '') + + def get_sd_unet(self): + return self._get('/sdapi/v1/options').get('sd_unet') + + def generate(self): + payload = {'prompt': 'a photo of a cat', 'steps': self.steps, 'width': 512, 'height': 512, 'save_images': False} + t0 = time.time() + r = self._post('/sdapi/v1/txt2img', data=payload, timeout=self.gen_timeout) + elapsed = time.time() - t0 + if r.status_code != 200: + return False, f'http {r.status_code}: {r.text[:160]}' + body = r.json() + images = body.get('images') or [] + if not images: + return False, f'no images returned ({elapsed:.1f}s)' + return True, f'{elapsed:.1f}s' + + # ---- scenarios --------------------------------------------------------- + + def test_discovery(self): + print('=== discovery ===', flush=True) + try: + unets = self._get('/sdapi/v1/unets') + self.record(isinstance(unets, list), 'GET /sdapi/v1/unets', f'{len(unets)} unets') + models = self._get('/sdapi/v1/sd-models') + self.record(isinstance(models, list), 'GET /sdapi/v1/sd-models', f'{len(models)} models') + if not (self.base_model and self.mismatch_unet): + print(' available UNET names:', flush=True) + for u in unets: + print(f' - {u.get("name")}', flush=True) + print(' available model titles:', flush=True) + for m in models[:40]: + print(f' - {m.get("title")}', flush=True) + except Exception as e: + self.record(False, 'discovery', f'exception: {e}') + + def test_baseline(self): + print('=== baseline (base model, no override) ===', flush=True) + if not self.base_model: + self.skip('baseline', 'no --base-model') + return False + ok, err = self.set_options(sd_unet='Default', sd_model_checkpoint=self.base_model) + if not ok: + self.record(False, 'set base model + Default unet', err) + return False + ok, err = self.reload(force=True) + if not ok: + self.record(False, 'reload base model', err) + return False + ok, detail = self.generate() + self.record(ok, 'generate with base model', detail) + return ok + + def test_mismatch_self_heal(self): + print('=== mismatch self-heal ===', flush=True) + if not (self.base_model and self.mismatch_unet): + self.skip('mismatch self-heal', 'needs --base-model and --mismatch-unet') + return + # configure base model with the wrong-arch override, then force a clean reload + ok, err = self.set_options(sd_model_checkpoint=self.base_model, sd_unet=self.mismatch_unet) + if not ok: + self.record(False, 'set base model + mismatch override', err) + return + ok, err = self.reload(force=True) + # the reload itself must not error out (the whole point of the fix) + self.record(ok, 'reload does not error on mismatched override', err) + # override must self-heal back to Default + healed = self.get_sd_unet() + self.record(healed == 'Default', 'sd_unet reset to Default', f'sd_unet={healed!r}') + # base model must still be usable + gen_ok, detail = self.generate() + self.record(gen_ok, 'generate after self-heal', detail) + + def test_match_preserved(self): + print('=== match preserved (no false drop) ===', flush=True) + if not (self.base_model and self.match_unet): + self.skip('match preserved', 'no --match-unet') + return + ok, err = self.set_options(sd_model_checkpoint=self.base_model, sd_unet=self.match_unet) + if not ok: + self.record(False, 'set base model + matching override', err) + return + ok, err = self.reload(force=True) + self.record(ok, 'reload with matching override', err) + kept = self.get_sd_unet() + self.record(kept == self.match_unet, 'matching override kept (not dropped)', f'sd_unet={kept!r}') + gen_ok, detail = self.generate() + self.record(gen_ok, 'generate with matching override', detail) + + def cleanup(self): + self.set_options(sd_unet='Default') + + def run(self): + try: + self.test_discovery() + self.test_baseline() + self.test_mismatch_self_heal() + self.test_match_preserved() + finally: + self.cleanup() + print('=== results ===', flush=True) + print(f' passed={self.passed} failed={self.failed} skipped={self.skipped}', flush=True) + return self.failed == 0 + + +def main(): + ap = argparse.ArgumentParser(description='Override arch-mismatch self-heal API tests') + ap.add_argument('--url', default='http://127.0.0.1:7860', help='SD.Next base URL') + ap.add_argument('--base-model', default=None, help='checkpoint title to load as the base (DiT arch)') + ap.add_argument('--mismatch-unet', default=None, help='UNET name whose arch does NOT match the base') + ap.add_argument('--match-unet', default=None, help='optional UNET name whose arch DOES match the base') + ap.add_argument('--steps', type=int, default=4) + args = ap.parse_args() + + try: + requests.get(f'{args.url.rstrip("/")}/sdapi/v1/sd-models', timeout=10, verify=False) + except Exception as e: + print(f'cannot reach SD.Next at {args.url}: {e}', flush=True) + return 2 + + ok = OverrideMismatchAPITest( + args.url, base_model=args.base_model, mismatch_unet=args.mismatch_unet, + match_unet=args.match_unet, steps=args.steps, + ).run() + return 0 if ok else 1 + + +if __name__ == '__main__': + sys.exit(main())