mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 17:24:32 +02:00
fix(native-loader): handle arch-mismatched UNET/DiT override on cold start
This commit is contained in:
@@ -375,6 +375,7 @@ def test_transformer_spec_defaults():
|
||||
assert spec.siblings == {}
|
||||
assert spec.acceptable_missing == ('rope.', 'pos_embedder.', 'learnable_pos_embed.')
|
||||
assert spec.forbidden_markers == ()
|
||||
assert spec.required_markers == ()
|
||||
|
||||
|
||||
def test_sibling_spec_defaults():
|
||||
@@ -684,6 +685,185 @@ def test_load_rejects_non_safetensors():
|
||||
assert '.safetensors' in str(e)
|
||||
|
||||
|
||||
def test_load_raises_override_arch_mismatch():
|
||||
"""A file missing the spec's required markers raises OverrideArchMismatch
|
||||
from load() (the defense-in-depth contract), before the converter runs."""
|
||||
fd, path = tempfile.mkstemp(suffix='.safetensors')
|
||||
try:
|
||||
raw = {'model.diffusion_model.blocks.0.self_attn.q_proj.weight': torch.zeros(4)}
|
||||
write_fixture(raw, fd, path)
|
||||
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,
|
||||
required_markers=(('double_blocks.', 'Chroma double-stream blocks'),),
|
||||
)
|
||||
raised = False
|
||||
try:
|
||||
nt.load(local_file=path, repo_id='fake/repo', spec=spec, diffusers_cfg={})
|
||||
except nt.OverrideArchMismatch as e:
|
||||
raised = True
|
||||
assert 'double_blocks.' in str(e)
|
||||
assert raised, 'expected OverrideArchMismatch'
|
||||
finally:
|
||||
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)
|
||||
|
||||
|
||||
# ============================================================
|
||||
# detect_prefix
|
||||
# ============================================================
|
||||
|
||||
def test_detect_prefix_bare_returns_empty():
|
||||
keys = ['layers.0.weight', 'layers.0.bias']
|
||||
assert nt.detect_prefix(keys, nt.DEFAULT_PREFIXES, 'Test') == ''
|
||||
|
||||
|
||||
def test_detect_prefix_returns_dominant():
|
||||
keys = [f'model.diffusion_model.layers.{i}.weight' for i in range(5)]
|
||||
assert nt.detect_prefix(keys, nt.DEFAULT_PREFIXES, 'Test') == 'model.diffusion_model.'
|
||||
|
||||
|
||||
def test_detect_prefix_mixed_raises():
|
||||
keys = ['model.diffusion_model.layers.0.weight', 'net.layers.0.weight']
|
||||
try:
|
||||
nt.detect_prefix(keys, nt.DEFAULT_PREFIXES, 'Test')
|
||||
raise AssertionError('expected ValueError')
|
||||
except ValueError as e:
|
||||
assert 'mixed prefixes' in str(e)
|
||||
|
||||
|
||||
def test_detect_prefix_accepts_iterable_keys():
|
||||
"""Header-only key lists (no tensors) must work, not just dicts."""
|
||||
keys = ['diffusion_model.a.weight', 'diffusion_model.b.weight']
|
||||
assert nt.detect_prefix(keys, nt.DEFAULT_PREFIXES, 'Test') == 'diffusion_model.'
|
||||
|
||||
|
||||
# ============================================================
|
||||
# check_required_markers
|
||||
# ============================================================
|
||||
|
||||
def test_required_markers_passes_when_present():
|
||||
sd = {'double_blocks.0.x': 1, 'distilled_guidance_layer.in_proj.weight': 2}
|
||||
markers = (('double_blocks.', 'double blocks'), ('distilled_guidance_layer.', 'guidance'))
|
||||
nt.check_required_markers(sd, markers, 'Chroma', '/tmp/x.safetensors')
|
||||
# no exception = pass
|
||||
|
||||
|
||||
def test_required_markers_raises_when_missing():
|
||||
sd = {'blocks.0.self_attn.q_proj.weight': 1}
|
||||
markers = (('double_blocks.', 'Chroma double-stream blocks'),)
|
||||
try:
|
||||
nt.check_required_markers(sd, markers, 'Chroma', '/tmp/x.safetensors')
|
||||
raise AssertionError('expected OverrideArchMismatch')
|
||||
except nt.OverrideArchMismatch as e:
|
||||
msg = str(e)
|
||||
assert 'Chroma double-stream blocks' in msg
|
||||
assert 'double_blocks.' in msg
|
||||
|
||||
|
||||
def test_required_markers_requires_all():
|
||||
"""All listed markers must match; one missing raises."""
|
||||
sd = {'double_blocks.0.x': 1}
|
||||
markers = (('double_blocks.', 'double blocks'), ('distilled_guidance_layer.', 'guidance'))
|
||||
try:
|
||||
nt.check_required_markers(sd, markers, 'Chroma', '/tmp/x.safetensors')
|
||||
raise AssertionError('expected OverrideArchMismatch')
|
||||
except nt.OverrideArchMismatch as e:
|
||||
assert 'guidance' in str(e)
|
||||
|
||||
|
||||
def test_required_markers_empty_no_op():
|
||||
sd = {'anything.weight': 1}
|
||||
nt.check_required_markers(sd, (), 'Test', '/tmp/x.safetensors')
|
||||
|
||||
|
||||
# ============================================================
|
||||
# check_override_compatible (header-only)
|
||||
# ============================================================
|
||||
|
||||
def write_keys_fixture(keys: list) -> str:
|
||||
fd, path = tempfile.mkstemp(suffix='.safetensors')
|
||||
os.close(fd)
|
||||
safetensors.torch.save_file({k: torch.zeros(2) for k in keys}, path)
|
||||
return path
|
||||
|
||||
|
||||
def test_compat_accepts_matching_arch():
|
||||
path = write_keys_fixture([
|
||||
'model.diffusion_model.double_blocks.0.img_attn.qkv.weight',
|
||||
'model.diffusion_model.distilled_guidance_layer.in_proj.weight',
|
||||
])
|
||||
try:
|
||||
spec = nt.TransformerSpec(
|
||||
cls=MockMiniTransformer,
|
||||
converter=lambda sd: sd,
|
||||
required_markers=(('double_blocks.', 'double blocks'), ('distilled_guidance_layer.', 'guidance')),
|
||||
)
|
||||
ok, reason = nt.check_override_compatible(path, spec)
|
||||
assert ok is True, reason
|
||||
assert reason == ''
|
||||
finally:
|
||||
os.unlink(path)
|
||||
|
||||
|
||||
def test_compat_rejects_missing_required():
|
||||
path = write_keys_fixture([
|
||||
'model.diffusion_model.blocks.0.self_attn.q_proj.weight',
|
||||
'model.diffusion_model.blocks.0.adaln_modulation_self_attn.1.weight',
|
||||
])
|
||||
try:
|
||||
spec = nt.TransformerSpec(
|
||||
cls=MockMiniTransformer,
|
||||
converter=lambda sd: sd,
|
||||
required_markers=(('double_blocks.', 'Chroma double-stream blocks'),),
|
||||
)
|
||||
ok, reason = nt.check_override_compatible(path, spec)
|
||||
assert ok is False
|
||||
assert 'double_blocks.' in reason
|
||||
finally:
|
||||
os.unlink(path)
|
||||
|
||||
|
||||
def test_compat_rejects_forbidden_marker():
|
||||
path = write_keys_fixture(['double_blocks.0.x.weight', 'legacy.bad.key.weight'])
|
||||
try:
|
||||
spec = nt.TransformerSpec(
|
||||
cls=MockMiniTransformer,
|
||||
forbidden_markers=(('legacy.bad.key.weight', 'legacy format'),),
|
||||
required_markers=(('double_blocks.', 'double blocks'),),
|
||||
)
|
||||
ok, reason = nt.check_override_compatible(path, spec)
|
||||
assert ok is False
|
||||
assert 'legacy format' in reason
|
||||
finally:
|
||||
os.unlink(path)
|
||||
|
||||
|
||||
def test_compat_defers_on_unreadable_file():
|
||||
spec = nt.TransformerSpec(cls=MockMiniTransformer, required_markers=(('double_blocks.', 'd'),))
|
||||
ok, reason = nt.check_override_compatible('/nonexistent/path/model.safetensors', spec)
|
||||
assert ok is True
|
||||
assert reason == ''
|
||||
|
||||
|
||||
def test_compat_no_required_markers_accepts_any():
|
||||
path = write_keys_fixture(['whatever.weight'])
|
||||
try:
|
||||
spec = nt.TransformerSpec(cls=MockMiniTransformer)
|
||||
ok, reason = nt.check_override_compatible(path, spec)
|
||||
assert ok is True, reason
|
||||
finally:
|
||||
os.unlink(path)
|
||||
|
||||
|
||||
# ============================================================
|
||||
# Run
|
||||
# ============================================================
|
||||
@@ -767,6 +947,37 @@ def run_all():
|
||||
]:
|
||||
run_test(cat, fn)
|
||||
|
||||
log.warning('=== detect_prefix ===')
|
||||
cat = category('detect')
|
||||
for fn in [
|
||||
test_detect_prefix_bare_returns_empty,
|
||||
test_detect_prefix_returns_dominant,
|
||||
test_detect_prefix_mixed_raises,
|
||||
test_detect_prefix_accepts_iterable_keys,
|
||||
]:
|
||||
run_test(cat, fn)
|
||||
|
||||
log.warning('=== check_required_markers ===')
|
||||
cat = category('required')
|
||||
for fn in [
|
||||
test_required_markers_passes_when_present,
|
||||
test_required_markers_raises_when_missing,
|
||||
test_required_markers_requires_all,
|
||||
test_required_markers_empty_no_op,
|
||||
]:
|
||||
run_test(cat, fn)
|
||||
|
||||
log.warning('=== check_override_compatible ===')
|
||||
cat = category('compat')
|
||||
for fn in [
|
||||
test_compat_accepts_matching_arch,
|
||||
test_compat_rejects_missing_required,
|
||||
test_compat_rejects_forbidden_marker,
|
||||
test_compat_defers_on_unreadable_file,
|
||||
test_compat_no_required_markers_accepts_any,
|
||||
]:
|
||||
run_test(cat, fn)
|
||||
|
||||
log.warning('=== end-to-end load ===')
|
||||
cat = category('load')
|
||||
for fn in [
|
||||
@@ -775,6 +986,7 @@ def run_all():
|
||||
test_load_end_to_end_with_sibling_partition,
|
||||
test_load_raises_on_missing_sibling_class,
|
||||
test_load_rejects_non_safetensors,
|
||||
test_load_raises_override_arch_mismatch,
|
||||
]:
|
||||
run_test(cat, fn)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user