#!/usr/bin/env python """ Offline unit tests for pipelines.native_transformer. Covers the pure helpers that own per-arch knob handling: - ``strip_prefix`` for single/multi prefix detection and mixed-prefix rejection - ``partition_siblings`` for inline-sibling key partitioning - ``check_forbidden_markers`` for structural-mismatch rejection - ``is_noop_converter`` for diffusers no-op lambda detection - ``validate_state_dict_load`` for unexpected / missing key handling - ``make_default_spec`` default-spec synthesis with diffusers converter pickup - ``auto_pickup_converter`` for diffusers ``SINGLE_FILE_LOADABLE_CLASSES`` integration - ``TransformerSpec`` / ``SiblingSpec`` defaults Plus one end-to-end ``load`` test against a tiny mock module that exercises the read -> strip -> convert -> from_config -> load_state_dict -> validate pipeline without needing a real diffusers transformer or hf_hub_download. No running server required. Usage: python test/test-native-transformer.py """ import os import sys import tempfile import torch import safetensors.torch script_dir = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) sys.path.insert(0, script_dir) os.chdir(script_dir) os.environ['SD_INSTALL_QUIET'] = '1' # Bootstrap cmd_args before any module that pulls in shared.py. import modules.cmd_args # pylint: disable=wrong-import-position import installer # pylint: disable=wrong-import-position _orig_argv = sys.argv sys.argv = [sys.argv[0]] try: modules.cmd_args.parse_args() finally: sys.argv = _orig_argv installer.add_args(modules.cmd_args.parser) modules.cmd_args.parsed, _ = modules.cmd_args.parser.parse_known_args([]) from modules.errors import log # pylint: disable=wrong-import-position from pipelines import native_transformer as nt # pylint: disable=wrong-import-position # ============================================================ # Test infrastructure # ============================================================ results: dict[str, dict] = {} def category(name: str): if name not in results: results[name] = {'passed': 0, 'failed': 0, 'tests': []} return name def record(cat: str, passed: bool, name: str, detail: str = ''): status = 'PASS' if passed else 'FAIL' results[cat]['passed' if passed else 'failed'] += 1 results[cat]['tests'].append((status, name)) msg = f' {status}: {name}' if detail: msg += f' ({detail})' if passed: log.info(msg) else: log.error(msg) def run_test(cat: str, fn): name = fn.__name__ try: ok = fn() if ok is False: record(cat, False, name) else: record(cat, True, name) except AssertionError as e: record(cat, False, name, str(e)) except Exception as e: # pylint: disable=broad-except record(cat, False, name, f'exception: {e}') import traceback traceback.print_exc() # ============================================================ # strip_prefix # ============================================================ def test_strip_prefix_bare_keys_pass_through(): sd = {'layers.0.weight': 1, 'layers.0.bias': 2} out = nt.strip_prefix(sd, nt.DEFAULT_PREFIXES, 'Test') assert out == sd, 'bare keys must pass through unchanged' def test_strip_prefix_dominant_single_variant(): sd = {f'model.diffusion_model.layers.{i}.weight': i for i in range(10)} out = nt.strip_prefix(sd, nt.DEFAULT_PREFIXES, 'Test') assert all(k.startswith('layers.') for k in out) assert len(out) == 10 def test_strip_prefix_picks_longest_match_first(): """``model.diffusion_model.`` must beat ``diffusion_model.`` when both match.""" sd = { 'model.diffusion_model.layers.0.weight': 1, 'model.diffusion_model.layers.1.weight': 2, } out = nt.strip_prefix(sd, nt.DEFAULT_PREFIXES, 'Test') # If shorter prefix matched, keys would start with 'model.' assert 'layers.0.weight' in out assert 'layers.1.weight' in out assert not any(k.startswith('model.') for k in out) def test_strip_prefix_mixed_prefixes_raises(): sd = { 'model.diffusion_model.layers.0.weight': 1, 'net.layers.0.weight': 2, } try: nt.strip_prefix(sd, nt.DEFAULT_PREFIXES, 'Test') raise AssertionError('expected ValueError') except ValueError as e: assert 'mixed prefixes' in str(e) def test_strip_prefix_net_variant(): sd = {'net.layers.0.weight': 1, 'net.layers.1.bias': 2} out = nt.strip_prefix(sd, nt.DEFAULT_PREFIXES, 'Test') assert set(out.keys()) == {'layers.0.weight', 'layers.1.bias'} def test_strip_prefix_diffusion_model_variant(): sd = {'diffusion_model.layers.0.weight': 1, 'diffusion_model.layers.1.bias': 2} out = nt.strip_prefix(sd, nt.DEFAULT_PREFIXES, 'Test') assert set(out.keys()) == {'layers.0.weight', 'layers.1.bias'} def test_strip_prefix_custom_prefix_set(): sd = {'lora_unet_blocks_0.weight': 1, 'lora_unet_blocks_1.weight': 2} out = nt.strip_prefix(sd, ('lora_unet_',), 'Test') assert set(out.keys()) == {'blocks_0.weight', 'blocks_1.weight'} # ============================================================ # partition_siblings # ============================================================ def test_partition_siblings_empty_spec_returns_state_dict_unchanged(): sd = {'a': 1, 'b': 2} transformer_sd, siblings = nt.partition_siblings(sd, {}) assert transformer_sd == sd assert siblings == {} def test_partition_siblings_no_matches_keeps_all_in_transformer(): sd = {'layers.0.weight': 1, 'layers.1.weight': 2} siblings_spec = {'llm_adapter': nt.SiblingSpec(subfolder='llm_adapter', inline_prefix='llm_adapter.')} transformer_sd, siblings = nt.partition_siblings(sd, siblings_spec) assert transformer_sd == sd assert siblings == {'llm_adapter': {}} def test_partition_siblings_single_sibling_split(): sd = { 'layers.0.weight': 'tx0', 'layers.1.weight': 'tx1', 'llm_adapter.input_proj.weight': 'ad0', 'llm_adapter.output_proj.weight': 'ad1', } siblings_spec = {'llm_adapter': nt.SiblingSpec(subfolder='llm_adapter', inline_prefix='llm_adapter.')} transformer_sd, siblings = nt.partition_siblings(sd, siblings_spec) assert set(transformer_sd.keys()) == {'layers.0.weight', 'layers.1.weight'} assert set(siblings['llm_adapter'].keys()) == {'input_proj.weight', 'output_proj.weight'} assert siblings['llm_adapter']['input_proj.weight'] == 'ad0' def test_partition_siblings_multiple_siblings(): sd = { 'layers.0.weight': 'tx', 'sibling_a.x.weight': 'a0', 'sibling_b.y.weight': 'b0', 'sibling_b.z.weight': 'b1', } siblings_spec = { 'sibling_a': nt.SiblingSpec(subfolder='a', inline_prefix='sibling_a.'), 'sibling_b': nt.SiblingSpec(subfolder='b', inline_prefix='sibling_b.'), } transformer_sd, siblings = nt.partition_siblings(sd, siblings_spec) assert list(transformer_sd.keys()) == ['layers.0.weight'] assert set(siblings['sibling_a'].keys()) == {'x.weight'} assert set(siblings['sibling_b'].keys()) == {'y.weight', 'z.weight'} # ============================================================ # check_forbidden_markers # ============================================================ def test_forbidden_markers_passes_when_absent(): sd = {'layers.0.weight': 1} markers = (('legacy.marker.weight', 'old format'),) nt.check_forbidden_markers(sd, markers, 'Test', '/tmp/x.safetensors') # no exception = pass def test_forbidden_markers_raises_when_present(): sd = {'layers.0.weight': 1, 'legacy.marker.weight': 2} markers = (('legacy.marker.weight', 'old Cosmos 1.0 structure'),) try: nt.check_forbidden_markers(sd, markers, 'Test', '/tmp/x.safetensors') raise AssertionError('expected ValueError') except ValueError as e: msg = str(e) assert 'old Cosmos 1.0 structure' in msg assert 'legacy.marker.weight' in msg def test_forbidden_markers_empty_tuple_no_op(): sd = {'layers.0.weight': 1} nt.check_forbidden_markers(sd, (), 'Test', '/tmp/x.safetensors') # ============================================================ # is_noop_converter # ============================================================ def test_noop_converter_identity_lambda(): fn = lambda checkpoint, **kwargs: checkpoint # pylint: disable=unnecessary-lambda-assignment assert nt.is_noop_converter(fn) is True def test_noop_converter_real_function(): def real(checkpoint, **kwargs): # pylint: disable=unused-argument return {k.replace('a.', 'b.'): v for k, v in checkpoint.items()} assert nt.is_noop_converter(real) is False def test_noop_converter_lambda_with_modification(): fn = lambda checkpoint, **kwargs: {k: v.float() for k, v in checkpoint.items()} # pylint: disable=unnecessary-lambda-assignment assert nt.is_noop_converter(fn) is False # ============================================================ # validate_state_dict_load # ============================================================ def test_validate_accepts_buffer_only_missing(): nt.validate_state_dict_load( component_name='transformer', missing=['rope.freqs', 'pos_embedder.pos'], unexpected=[], acceptable_missing=('rope.', 'pos_embedder.'), ) def test_validate_rejects_unexpected(): try: nt.validate_state_dict_load( component_name='transformer', missing=[], unexpected=['some.junk.weight'], acceptable_missing=(), ) raise AssertionError('expected OverrideArchMismatch') except nt.OverrideArchMismatch as e: assert 'unexpected' in str(e) assert 'some.junk.weight' in str(e) def test_validate_rejects_hard_missing(): try: nt.validate_state_dict_load( component_name='transformer', missing=['layers.0.weight', 'rope.freqs'], unexpected=[], acceptable_missing=('rope.',), ) raise AssertionError('expected OverrideArchMismatch') except nt.OverrideArchMismatch as e: msg = str(e) assert 'missing' in msg assert 'layers.0.weight' in msg # Buffer-only missing must not show up in the hard-missing list assert msg.count('rope.freqs') == 0 def test_validate_empty_passes(): nt.validate_state_dict_load( component_name='transformer', missing=[], unexpected=[], acceptable_missing=(), ) # ============================================================ # make_default_spec # ============================================================ class FakeTransformer: """Minimal stand-in for a diffusers transformer class.""" def test_make_default_spec_for_unknown_class(): spec = nt.make_default_spec(FakeTransformer) assert spec.cls is FakeTransformer assert spec.subfolder == 'transformer' assert spec.prefixes == nt.DEFAULT_PREFIXES assert spec.converter is None # no diffusers entry for FakeTransformer assert spec.siblings == {} assert spec.forbidden_markers == () def test_make_default_spec_picks_up_real_diffusers_converter(): import diffusers spec = nt.make_default_spec(diffusers.FluxTransformer2DModel) assert spec.converter is not None assert spec.converter.__name__ == 'convert_flux_transformer_checkpoint_to_diffusers' def test_make_default_spec_skips_qwen_image_noop(): """QwenImageTransformer2DModel's diffusers entry is a no-op lambda; the default spec must NOT pick it up, leaving converter=None so the caller sees only their own (potentially absent) override.""" import diffusers spec = nt.make_default_spec(diffusers.QwenImageTransformer2DModel) assert spec.converter is None # ============================================================ # auto_pickup_converter # ============================================================ def test_auto_pickup_returns_none_for_unknown_class(): assert nt.auto_pickup_converter(FakeTransformer) is None def test_auto_pickup_returns_real_diffusers_converter(): import diffusers fn = nt.auto_pickup_converter(diffusers.FluxTransformer2DModel) assert fn is not None assert callable(fn) assert fn.__name__ == 'convert_flux_transformer_checkpoint_to_diffusers' def test_auto_pickup_skips_noop_converter_qwen(): """QwenImageTransformer2DModel registers a no-op lambda in diffusers; auto_pickup_converter must return None so the spec falls back to no converter (the user-registered spec can override with a real converter).""" import diffusers assert nt.auto_pickup_converter(diffusers.QwenImageTransformer2DModel) is None # ============================================================ # TransformerSpec / SiblingSpec defaults # ============================================================ def test_transformer_spec_defaults(): spec = nt.TransformerSpec(cls=FakeTransformer) assert spec.subfolder == 'transformer' assert spec.prefixes == ('model.diffusion_model.', 'diffusion_model.', 'net.') assert spec.converter is None assert spec.siblings == {} assert spec.acceptable_missing == ('rope.', 'pos_embedder.', 'learnable_pos_embed.') assert spec.forbidden_markers == () def test_sibling_spec_defaults(): spec = nt.SiblingSpec(subfolder='llm_adapter', inline_prefix='llm_adapter.') assert spec.subfolder == 'llm_adapter' assert spec.inline_prefix == 'llm_adapter.' assert spec.acceptable_missing == () def test_transformer_spec_is_frozen(): spec = nt.TransformerSpec(cls=FakeTransformer) try: spec.subfolder = 'changed' # type: ignore[misc] except Exception as e: # pylint: disable=broad-except assert 'FrozenInstanceError' in type(e).__name__ or 'frozen' in str(e).lower() return raise AssertionError('expected FrozenInstanceError') # ============================================================ # Integration: end-to-end load() with a tiny mock module # ============================================================ # We sidestep diffusers + hf_hub_download by patching: # - ``fetch_component_config`` to return a hand-rolled config dict # - ``model_quant.get_dit_args`` / ``model_quant.get_quant_type`` / quant # application to no-ops (we only want to test the load path itself). # The mock cls is a torch.nn.Module subclass whose ``from_config`` constructs # a fresh module of the expected shape; ``load_state_dict`` is the standard # PyTorch method. class MockMiniTransformer(torch.nn.Module): """Tiny stand-in: linear in -> linear out, plus a nested rope sub-module holding a buffer the trainer state dict won't carry. Nested mirrors how real DiTs structure rope / pos_embedder buffers.""" @classmethod def from_config(cls, config: dict) -> 'MockMiniTransformer': return cls(dim=config['dim']) def __init__(self, dim: int): super().__init__() self.in_proj = torch.nn.Linear(dim, dim) self.out_proj = torch.nn.Linear(dim, dim) self.rope = torch.nn.Module() self.rope.register_buffer('freqs', torch.zeros(dim)) class MockKwargsTransformer(MockMiniTransformer): """Records the kwargs from_config received, so a test can assert the native path forwards caller kwargs to construction. Mirrors diffusers from_config, which accepts **kwargs.""" last_kwargs: dict = {} @classmethod def from_config(cls, config: dict, **kwargs) -> 'MockKwargsTransformer': cls.last_kwargs = dict(kwargs) return cls(dim=config['dim']) def write_fixture(state_dict_keys: dict, fd: int, path: str) -> str: os.close(fd) safetensors.torch.save_file(state_dict_keys, path) return path def test_load_end_to_end_with_bfl_prefix_no_converter(): """Exercise the full load pipeline: read .safetensors, strip prefix, no converter, instantiate via from_config, load weights, validate. """ fd, path = tempfile.mkstemp(suffix='.safetensors') try: dim = 8 # Save with model.diffusion_model. prefix; in_proj.* and out_proj.* # are the real weights the mock cls expects after the strip. raw = { 'model.diffusion_model.in_proj.weight': torch.randn(dim, dim), 'model.diffusion_model.in_proj.bias': torch.zeros(dim), 'model.diffusion_model.out_proj.weight': torch.randn(dim, dim), 'model.diffusion_model.out_proj.bias': torch.zeros(dim), } write_fixture(raw, fd, path) # Patch fetch_component_config to return our hand-rolled config. orig_fetch = nt.fetch_component_config nt.fetch_component_config = lambda repo, sub: {'dim': dim} # Patch quant helpers (we only care about the load path). from modules import model_quant orig_get_dit = model_quant.get_dit_args orig_get_qtype = model_quant.get_quant_type orig_do_post = model_quant.do_post_load_quant model_quant.get_dit_args = lambda *a, **k: ({}, {}) model_quant.get_quant_type = lambda *a, **k: None model_quant.do_post_load_quant = lambda *a, **k: None try: spec = nt.TransformerSpec(cls=MockMiniTransformer) transformer, siblings = nt.load( local_file=path, repo_id='fake/repo', spec=spec, diffusers_cfg={}, ) finally: nt.fetch_component_config = orig_fetch model_quant.get_dit_args = orig_get_dit model_quant.get_quant_type = orig_get_qtype model_quant.do_post_load_quant = orig_do_post assert isinstance(transformer, MockMiniTransformer) assert transformer.in_proj.weight.shape == (dim, dim) # Weights from the fixture should match what was loaded. loaded_in_w = transformer.in_proj.weight.detach().cpu() fixture_in_w = raw['model.diffusion_model.in_proj.weight'].to(loaded_in_w.dtype) assert torch.allclose(loaded_in_w, fixture_in_w) assert siblings == {} finally: if os.path.exists(path): os.unlink(path) def test_load_forwards_kwargs_to_from_config(): """Caller **kwargs reach cls.from_config through the native load path rather than being dropped.""" fd, path = tempfile.mkstemp(suffix='.safetensors') try: dim = 8 raw = { 'model.diffusion_model.in_proj.weight': torch.randn(dim, dim), 'model.diffusion_model.in_proj.bias': torch.zeros(dim), 'model.diffusion_model.out_proj.weight': torch.randn(dim, dim), 'model.diffusion_model.out_proj.bias': torch.zeros(dim), } write_fixture(raw, fd, path) orig_fetch = nt.fetch_component_config nt.fetch_component_config = lambda repo, sub: {'dim': dim} from modules import model_quant orig_get_dit = model_quant.get_dit_args orig_get_qtype = model_quant.get_quant_type orig_do_post = model_quant.do_post_load_quant model_quant.get_dit_args = lambda *a, **k: ({}, {}) model_quant.get_quant_type = lambda *a, **k: None model_quant.do_post_load_quant = lambda *a, **k: None MockKwargsTransformer.last_kwargs = {} try: spec = nt.TransformerSpec(cls=MockKwargsTransformer) transformer, _ = nt.load( local_file=path, repo_id='fake/repo', spec=spec, diffusers_cfg={}, low_cpu_mem_usage=True, ) finally: nt.fetch_component_config = orig_fetch model_quant.get_dit_args = orig_get_dit model_quant.get_quant_type = orig_get_qtype model_quant.do_post_load_quant = orig_do_post assert MockKwargsTransformer.last_kwargs == {'low_cpu_mem_usage': True} assert isinstance(transformer, MockKwargsTransformer) finally: if os.path.exists(path): os.unlink(path) def test_load_end_to_end_with_sibling_partition(): """Bundled-sibling case: file carries both transformer and sibling weights, sibling_classes supplies the runtime sibling class, partition routes each half into its target.""" fd, path = tempfile.mkstemp(suffix='.safetensors') try: dim = 8 sibling_dim = 4 raw = { # Transformer half (after strip). 'model.diffusion_model.in_proj.weight': torch.randn(dim, dim), 'model.diffusion_model.in_proj.bias': torch.zeros(dim), 'model.diffusion_model.out_proj.weight': torch.randn(dim, dim), 'model.diffusion_model.out_proj.bias': torch.zeros(dim), # Sibling half (after strip + sibling partition). 'model.diffusion_model.sibling.in_proj.weight': torch.randn(sibling_dim, sibling_dim), 'model.diffusion_model.sibling.in_proj.bias': torch.zeros(sibling_dim), 'model.diffusion_model.sibling.out_proj.weight': torch.randn(sibling_dim, sibling_dim), 'model.diffusion_model.sibling.out_proj.bias': torch.zeros(sibling_dim), } write_fixture(raw, fd, path) orig_fetch = nt.fetch_component_config def patched_fetch(_repo, sub): return {'dim': dim if sub == 'transformer' else sibling_dim} nt.fetch_component_config = patched_fetch from modules import model_quant orig_get_dit = model_quant.get_dit_args orig_get_qtype = model_quant.get_quant_type orig_do_post = model_quant.do_post_load_quant model_quant.get_dit_args = lambda *a, **k: ({}, {}) model_quant.get_quant_type = lambda *a, **k: None model_quant.do_post_load_quant = lambda *a, **k: None try: spec = nt.TransformerSpec( cls=MockMiniTransformer, siblings={ 'sibling': nt.SiblingSpec( subfolder='sibling', inline_prefix='sibling.', acceptable_missing=('rope.',), ), }, ) transformer, siblings = nt.load( local_file=path, repo_id='fake/repo', spec=spec, diffusers_cfg={}, sibling_classes={'sibling': MockMiniTransformer}, ) finally: nt.fetch_component_config = orig_fetch model_quant.get_dit_args = orig_get_dit model_quant.get_quant_type = orig_get_qtype model_quant.do_post_load_quant = orig_do_post assert isinstance(transformer, MockMiniTransformer) assert transformer.in_proj.weight.shape == (dim, dim) assert 'sibling' in siblings assert isinstance(siblings['sibling'], MockMiniTransformer) assert siblings['sibling'].in_proj.weight.shape == (sibling_dim, sibling_dim) finally: if os.path.exists(path): os.unlink(path) def test_load_raises_on_missing_sibling_class(): """Sibling keys present in file but caller forgot to supply the class.""" fd, path = tempfile.mkstemp(suffix='.safetensors') try: dim = 8 sibling_dim = 4 raw = { 'model.diffusion_model.in_proj.weight': torch.randn(dim, dim), 'model.diffusion_model.in_proj.bias': torch.zeros(dim), 'model.diffusion_model.out_proj.weight': torch.randn(dim, dim), 'model.diffusion_model.out_proj.bias': torch.zeros(dim), 'model.diffusion_model.sibling.in_proj.weight': torch.randn(sibling_dim, sibling_dim), } write_fixture(raw, fd, path) orig_fetch = nt.fetch_component_config nt.fetch_component_config = lambda repo, sub: {'dim': dim} from modules import model_quant orig_get_dit = model_quant.get_dit_args orig_get_qtype = model_quant.get_quant_type orig_do_post = model_quant.do_post_load_quant model_quant.get_dit_args = lambda *a, **k: ({}, {}) model_quant.get_quant_type = lambda *a, **k: None model_quant.do_post_load_quant = lambda *a, **k: None try: spec = nt.TransformerSpec( cls=MockMiniTransformer, siblings={'sibling': nt.SiblingSpec(subfolder='s', inline_prefix='sibling.')}, ) raised = False try: nt.load( local_file=path, repo_id='fake/repo', spec=spec, diffusers_cfg={}, sibling_classes={}, # missing! ) except ValueError as e: raised = True assert "'sibling'" in str(e) assert 'sibling_classes' in str(e) assert raised, 'expected ValueError' finally: nt.fetch_component_config = orig_fetch model_quant.get_dit_args = orig_get_dit model_quant.get_quant_type = orig_get_qtype model_quant.do_post_load_quant = orig_do_post finally: if os.path.exists(path): os.unlink(path) def test_load_rejects_non_safetensors(): spec = nt.TransformerSpec(cls=MockMiniTransformer) try: nt.load( local_file='/tmp/some.gguf', repo_id='fake/repo', spec=spec, diffusers_cfg={}, ) raise AssertionError('expected ValueError') except ValueError as e: 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 # ============================================================ def run_all(): log.warning('=== strip_prefix ===') cat = category('strip') for fn in [ test_strip_prefix_bare_keys_pass_through, test_strip_prefix_dominant_single_variant, test_strip_prefix_picks_longest_match_first, test_strip_prefix_mixed_prefixes_raises, test_strip_prefix_net_variant, test_strip_prefix_diffusion_model_variant, test_strip_prefix_custom_prefix_set, ]: run_test(cat, fn) log.warning('=== partition_siblings ===') cat = category('partition') for fn in [ test_partition_siblings_empty_spec_returns_state_dict_unchanged, test_partition_siblings_no_matches_keeps_all_in_transformer, test_partition_siblings_single_sibling_split, test_partition_siblings_multiple_siblings, ]: run_test(cat, fn) log.warning('=== forbidden_markers ===') cat = category('forbidden') for fn in [ test_forbidden_markers_passes_when_absent, test_forbidden_markers_raises_when_present, test_forbidden_markers_empty_tuple_no_op, ]: run_test(cat, fn) log.warning('=== noop_converter detection ===') cat = category('noop') for fn in [ test_noop_converter_identity_lambda, test_noop_converter_real_function, test_noop_converter_lambda_with_modification, ]: run_test(cat, fn) log.warning('=== validate_state_dict_load ===') cat = category('validate') for fn in [ test_validate_accepts_buffer_only_missing, test_validate_rejects_unexpected, test_validate_rejects_hard_missing, test_validate_empty_passes, ]: run_test(cat, fn) log.warning('=== make_default_spec ===') cat = category('default_spec') for fn in [ test_make_default_spec_for_unknown_class, test_make_default_spec_picks_up_real_diffusers_converter, test_make_default_spec_skips_qwen_image_noop, ]: run_test(cat, fn) log.warning('=== auto_pickup_converter ===') cat = category('autopickup') for fn in [ test_auto_pickup_returns_none_for_unknown_class, test_auto_pickup_returns_real_diffusers_converter, test_auto_pickup_skips_noop_converter_qwen, ]: run_test(cat, fn) log.warning('=== TransformerSpec / SiblingSpec ===') cat = category('specs') for fn in [ test_transformer_spec_defaults, test_sibling_spec_defaults, test_transformer_spec_is_frozen, ]: run_test(cat, fn) log.warning('=== end-to-end load ===') cat = category('load') for fn in [ test_load_end_to_end_with_bfl_prefix_no_converter, test_load_forwards_kwargs_to_from_config, 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) log.warning('=== Results ===') total_passed = 0 total_failed = 0 for cat_name, info in results.items(): ok = info['failed'] == 0 status = 'PASS' if ok else 'FAIL' log.info(f" {cat_name}: {info['passed']} passed, {info['failed']} failed [{status}]") total_passed += info['passed'] total_failed += info['failed'] log.warning(f'Total: {total_passed} passed, {total_failed} failed') return total_failed == 0 if __name__ == '__main__': import time t0 = time.time() ok = run_all() log.warning(f'Total time: {time.time() - t0:.2f}s') sys.exit(0 if ok else 1)