diff --git a/modules/model_probe.py b/modules/model_probe.py index f6b1dabca..7972c7fdd 100644 --- a/modules/model_probe.py +++ b/modules/model_probe.py @@ -10,7 +10,7 @@ from collections import Counter from dataclasses import dataclass -SCHEMA_VERSION = 4 +SCHEMA_VERSION = 5 MAX_HEADER_BYTES = 16 * 1024 * 1024 STRIP_PREFIXES = ('model.diffusion_model.', 'diffusion_model.', 'net.') # companion families bundled alongside the diffusion core in full checkpoints @@ -209,20 +209,33 @@ def match_arch(keys: list, metadata: dict | None, shapes: dict | None = None) -> } -def detect_quant(keys: list, dtypes: Counter, container: str) -> dict: +# storage dtype of a marked layer's weight identifies its comfy_quant format +COMFY_QUANT_DTYPE_FORMATS = {'I8': 'int8_tensorwise', 'F8_E4M3': 'float8_e4m3fn', 'U8': 'nvfp4'} + + +def detect_quant(keys: list, dtypes: Counter, container: str, metadata: dict | None = None, key_dtypes: dict | None = None) -> dict: if container == 'gguf': quantized = {d: n for d, n in dtypes.items() if d not in ('F32', 'F16', 'BF16')} dominant = max(quantized, key=quantized.get) if quantized else None return {'scheme': 'gguf' if dominant else None, 'format': dominant, 'marked_layers': None, 'source': 'gguf-qtype'} + quant_metadata = (metadata or {}).get('_quantization_metadata') + if quant_metadata: + try: + layers = (json.loads(quant_metadata) if isinstance(quant_metadata, str) else quant_metadata).get('layers') or {} + except Exception: + layers = {} + formats = Counter(entry.get('format') for entry in layers.values() if isinstance(entry, dict) and entry.get('format')) + if formats: + return {'scheme': 'comfy_quant', 'format': formats.most_common(1)[0][0], 'marked_layers': len(layers), 'source': 'header'} comfy = [k for k in keys if k.endswith('.comfy_quant')] if comfy: - if dtypes.get('U8', 0) > dtypes.get('F8_E4M3', 0): - fmt = 'int8_tensorwise' - elif dtypes.get('F8_E4M3', 0): - fmt = 'float8_e4m3fn' - else: - fmt = None - return {'scheme': 'comfy_quant', 'format': fmt, 'marked_layers': len(comfy), 'source': 'dtype-inferred'} + weight_dtypes = Counter() + for marker in comfy: + dtype = (key_dtypes or {}).get(f"{marker[: -len('.comfy_quant')]}.weight") + if dtype: + weight_dtypes[dtype] += 1 + dominant = weight_dtypes.most_common(1)[0][0] if weight_dtypes else None + return {'scheme': 'comfy_quant', 'format': COMFY_QUANT_DTYPE_FORMATS.get(dominant), 'marked_layers': len(comfy), 'source': 'weight-dtype'} has_fp8 = dtypes.get('F8_E4M3', 0) + dtypes.get('F8_E5M2', 0) > 0 has_scales = any(k.endswith(('scaled_fp8', '.scale_weight', '.scale_input', '.weight_scale')) for k in keys) if has_fp8 and has_scales: @@ -254,7 +267,8 @@ def analyze_header(header: dict, container: str = 'safetensors', arch_metadata: inner_keys = [strip_key(k) for k in keys] shapes = {k: v.get('shape') for k, v in entries.items()} arch = match_arch(keys, metadata, shapes) - quant = detect_quant(inner_keys, dtypes, container) + key_dtypes = {strip_key(k): v.get('dtype') for k, v in entries.items()} + quant = detect_quant(inner_keys, dtypes, container, metadata=metadata, key_dtypes=key_dtypes) flags = [] if not metadata: flags.append('no_metadata') diff --git a/test/test-native-transformer.py b/test/test-native-transformer.py index eb63e92fb..aee7fda93 100644 --- a/test/test-native-transformer.py +++ b/test/test-native-transformer.py @@ -637,6 +637,68 @@ def test_detect_comfy_nvfp4_convrot_raises(): assert 'convrot' in str(e) +# ============================================================ +# model_probe quant detection (header-only, mirrors the loader's coverage) +# ============================================================ + +def test_probe_detects_header_quant_metadata(): + import json + from modules import model_probe + header = { + '__metadata__': {'_quantization_metadata': json.dumps({'layers': {'blocks.0.wq': {'format': 'float8_e4m3fn'}}})}, + 'blocks.0.wq.weight': {'dtype': 'F8_E4M3', 'shape': [4, 4]}, + 'blocks.0.wq.weight_scale': {'dtype': 'F32', 'shape': []}, + } + quant = model_probe.analyze_header(header)['quant'] + assert quant == {'scheme': 'comfy_quant', 'format': 'float8_e4m3fn', 'marked_layers': 1, 'source': 'header'} + + +def test_probe_resolves_nvfp4_from_weight_dtype(): + from modules import model_probe + header = { + 'model.diffusion_model.blocks.0.wq.weight': {'dtype': 'U8', 'shape': [32, 16]}, + 'model.diffusion_model.blocks.0.wq.weight_scale': {'dtype': 'F8_E4M3', 'shape': [128, 4]}, + 'model.diffusion_model.blocks.0.wq.weight_scale_2': {'dtype': 'F32', 'shape': []}, + 'model.diffusion_model.blocks.0.wq.comfy_quant': {'dtype': 'U8', 'shape': [60]}, + } + quant = model_probe.analyze_header(header)['quant'] + assert quant['format'] == 'nvfp4' + assert quant['source'] == 'weight-dtype' + + +def test_probe_fp8_markers_despite_u8_noise(): + """The marked layer's own weight dtype decides the format; unrelated U8 + tensors must not outvote it.""" + from modules import model_probe + header = { + 'blocks.0.wq.weight': {'dtype': 'F8_E4M3', 'shape': [8, 8]}, + 'blocks.0.wq.weight_scale': {'dtype': 'F32', 'shape': []}, + 'blocks.0.wq.comfy_quant': {'dtype': 'U8', 'shape': [27]}, + 'extra.blob1': {'dtype': 'U8', 'shape': [64]}, + 'extra.blob2': {'dtype': 'U8', 'shape': [64]}, + } + quant = model_probe.analyze_header(header)['quant'] + assert quant['format'] == 'float8_e4m3fn' + + +def test_probe_int8_markers(): + from modules import model_probe + header = { + 'blocks.0.wq.weight': {'dtype': 'I8', 'shape': [8, 8]}, + 'blocks.0.wq.weight_scale': {'dtype': 'F32', 'shape': []}, + 'blocks.0.wq.comfy_quant': {'dtype': 'U8', 'shape': [29]}, + } + quant = model_probe.analyze_header(header)['quant'] + assert quant['format'] == 'int8_tensorwise' + assert quant['marked_layers'] == 1 + + +def test_probe_plain_file_no_quant(): + from modules import model_probe + header = {'blocks.0.wq.weight': {'dtype': 'BF16', 'shape': [8, 8]}} + assert model_probe.analyze_header(header)['quant']['scheme'] is None + + # ============================================================ # is_noop_converter # ============================================================ @@ -2184,6 +2246,17 @@ def run_all(): ]: run_test(cat, fn) + log.warning('=== model_probe quant detection ===') + cat = category('probe') + for fn in [ + test_probe_detects_header_quant_metadata, + test_probe_resolves_nvfp4_from_weight_dtype, + test_probe_fp8_markers_despite_u8_noise, + test_probe_int8_markers, + test_probe_plain_file_no_quant, + ]: + run_test(cat, fn) + log.warning('=== noop_converter detection ===') cat = category('noop') for fn in [