AX-Unlimited-OCR-3B-MoE-MLX-MXFP8 / quantization /mixed_precision_convert.py
AutomatosX's picture
Publish verified OCR-aware MXFP8 checkpoint
6708e4d verified
Raw
History Blame
14 kB
"""Mixed-precision conversion script for OCR-aware quantization.
Applies a precision map to produce enhanced model weights where OCR-sensitive
layers remain at BF16 while less sensitive layers use MXFP8.
Usage:
python quantization/mixed_precision_convert.py \
--model-path baidu/Unlimited-OCR \
--precision-map quantization/precision_map.json \
--output-dir ./enhanced_model/
Requires: mlx, mlx-vlm
"""
from __future__ import annotations
import argparse
import hashlib
import json
import re
import shutil
import tempfile
from pathlib import Path
SUPPORTED_PRECISIONS = {"bfloat16", "mxfp8", "affine8"}
QUANTIZATION_CONFIGS = {
"mxfp8": {"group_size": 32, "bits": 8, "mode": "mxfp8"},
"affine8": {"group_size": 32, "bits": 8, "mode": "affine"},
}
def _json_digest(value: dict) -> str:
payload = json.dumps(value, sort_keys=True, separators=(",", ":")).encode("utf-8")
return hashlib.sha256(payload).hexdigest()
def normalize_unlimited_ocr_metadata(model_dir: Path) -> None:
"""Select mlx-vlm's native backend and validate MXFP8/R-SWA metadata.
``mlx-vlm convert`` always emits ``config.json`` (and usually
``processor_config.json``). Missing configs mean conversion did not finish
and must not be published.
"""
config_path = model_dir / "config.json"
processor_path = model_dir / "processor_config.json"
if not config_path.is_file():
raise ValueError(
"Converted model is missing config.json — mlx-vlm conversion incomplete"
)
try:
config = json.loads(config_path.read_text(encoding="utf-8"))
except (OSError, UnicodeError, json.JSONDecodeError) as exc:
raise ValueError("Converted model has invalid config.json") from exc
if not isinstance(config, dict):
raise ValueError("Converted config.json must contain a JSON object")
architectures = config.get("architectures")
if not isinstance(architectures, list) or "UnlimitedOCRForCausalLM" not in architectures:
raise ValueError("Converted checkpoint is not UnlimitedOCRForCausalLM")
quantization = config.get("quantization") or config.get("quantization_config")
if not isinstance(quantization, dict) or quantization.get("mode") != "mxfp8":
raise ValueError("Converted checkpoint does not declare MXFP8 quantization")
text_config = config.get("language_config") or config.get("text_config") or config
window_size = None
if isinstance(text_config, dict):
window_size = text_config.get("sliding_window_size", text_config.get("sliding_window"))
if window_size is None:
window_size = config.get("sliding_window_size", config.get("sliding_window"))
if not isinstance(window_size, int) or isinstance(window_size, bool) or window_size < 1:
raise ValueError("Converted checkpoint is missing a positive sliding-window size")
config["model_type"] = "unlimited-ocr"
config.pop("auto_map", None)
# Keep R-SWA fields consistent at the top level and under language_config.
config["sliding_window"] = int(window_size)
config["sliding_window_size"] = int(window_size)
language_config = config.get("language_config")
if isinstance(language_config, dict):
language_config.pop("auto_map", None)
language_config["sliding_window"] = int(window_size)
language_config["sliding_window_size"] = int(window_size)
if processor_path.is_file():
try:
processor = json.loads(processor_path.read_text(encoding="utf-8"))
except (OSError, UnicodeError, json.JSONDecodeError) as exc:
raise ValueError("Converted model has invalid processor_config.json") from exc
if not isinstance(processor, dict):
raise ValueError("Converted processor_config.json must contain a JSON object")
else:
# Some convert paths only emit tokenizer assets; still mark the processor
# class so mlx-vlm loads the Unlimited-OCR handler.
processor = {}
processor["processor_class"] = "UnlimitedOCRHFProcessor"
processor["sft_format"] = "unlimitedocr"
config_path.write_text(
json.dumps(config, indent=2, ensure_ascii=False) + "\n",
encoding="utf-8",
)
processor_path.write_text(
json.dumps(processor, indent=2, ensure_ascii=False) + "\n",
encoding="utf-8",
)
# mlx-vlm convert may stamp a different processor_class on tokenizer_config.
# Keep it aligned so Hub/transformers-style loaders and mlx-vlm agree.
tokenizer_config_path = model_dir / "tokenizer_config.json"
if tokenizer_config_path.is_file():
try:
tokenizer_config = json.loads(
tokenizer_config_path.read_text(encoding="utf-8")
)
except (OSError, UnicodeError, json.JSONDecodeError) as exc:
raise ValueError(
"Converted model has invalid tokenizer_config.json"
) from exc
if not isinstance(tokenizer_config, dict):
raise ValueError(
"Converted tokenizer_config.json must contain a JSON object"
)
tokenizer_config["processor_class"] = "UnlimitedOCRHFProcessor"
tokenizer_config_path.write_text(
json.dumps(tokenizer_config, indent=2, ensure_ascii=False) + "\n",
encoding="utf-8",
)
def load_precision_map(path: Path) -> dict:
"""Load the precision map JSON."""
with open(path, encoding="utf-8") as f:
precision_map = json.load(f)
validate_precision_map(precision_map)
return precision_map
def validate_precision_map(precision_map: dict) -> None:
"""Validate precision-map structure and supported precision values."""
if not isinstance(precision_map, dict):
raise ValueError("precision map must be a JSON object")
rules = {
pattern: precision
for pattern, precision in precision_map.items()
if not pattern.startswith("_")
}
if not rules:
raise ValueError("precision map has no module rules")
invalid = {
pattern: precision
for pattern, precision in rules.items()
if not isinstance(precision, str) or precision not in SUPPORTED_PRECISIONS
}
if invalid:
details = ", ".join(f"{pattern}={precision!r}" for pattern, precision in invalid.items())
raise ValueError(f"unsupported precision-map values: {details}")
if "mxfp8" not in rules.values():
raise ValueError("precision map does not select any modules for MXFP8")
def matches_pattern(param_name: str, pattern: str) -> bool:
"""Check if a parameter name matches a precision map pattern.
Supports wildcard '*' for layer indices.
Example: ``model.layers.*.self_attn.q_proj`` matches layer 5's q_proj.
"""
regex = re.escape(pattern).replace(r"\*", r"\d+")
# Precision-map entries name complete module-path segments but parameter
# names may contain a model prefix and a trailing ``.weight``.
return bool(re.search(rf"(?:^|\.){regex}(?=\.|$)", param_name))
def get_precision_for_param(param_name: str, precision_map: dict) -> str:
"""Determine the target precision for a given parameter name."""
matching_rules = [
(index, pattern, precision)
for index, (pattern, precision) in enumerate(precision_map.items())
if not pattern.startswith("_") and matches_pattern(param_name, pattern)
]
if matching_rules:
# Exact per-layer sensitivity overrides must beat an earlier wildcard
# group rule. More literal path segments are more specific; insertion
# order is only a tie breaker.
_, _, precision = max(
matching_rules,
key=lambda item: (
len([segment for segment in item[1].split(".") if segment != "*"]),
len(item[1].split(".")),
item[0],
),
)
return precision
# Unlisted modules are preserved. Quantizing by a loose name heuristic can
# accidentally include the vision tower or projector.
return "bfloat16"
def convert_model(
model_path: str,
precision_map: dict,
output_dir: Path,
verbose: bool = True,
*,
source_revision: str | None = None,
) -> Path:
"""Apply mixed-precision quantization according to the precision map.
The conversion is staged in a temporary sibling directory and moved into
place only after quantized weights and all requested rule matches exist.
"""
validate_precision_map(precision_map)
output_dir = Path(output_dir)
if output_dir.exists():
raise FileExistsError(
f"Output directory already exists; choose a new path: {output_dir}"
)
output_dir.parent.mkdir(parents=True, exist_ok=True)
from mlx_vlm.convert import convert
if verbose:
print(f"Loading and converting model from: {model_path}")
quantized_modules: list[str] = []
quantized_precisions: dict[str, str] = {}
preserved_modules: list[str] = []
def quantization_predicate(path, module):
precision = get_precision_for_param(path, precision_map)
if precision in QUANTIZATION_CONFIGS:
quantized_modules.append(path)
quantized_precisions[path] = precision
return dict(QUANTIZATION_CONFIGS[precision])
preserved_modules.append(path)
return False
if verbose:
print("\nPrecision map summary:")
for pattern, precision in precision_map.items():
if not pattern.startswith("_"):
print(f" {pattern}: {precision}")
staging_dir = Path(tempfile.mkdtemp(
prefix=f".{output_dir.name}_staging_",
dir=output_dir.parent,
))
try:
convert(
hf_path=model_path,
mlx_path=str(staging_dir),
revision=source_revision,
quantize=True,
q_group_size=32,
q_bits=8,
q_mode="mxfp8",
quant_predicate=quantization_predicate,
)
if not quantized_modules:
raise RuntimeError(
"Precision map matched no quantizable MLX modules; refusing to save a BF16-only model"
)
requested_patterns = [
pattern
for pattern, precision in precision_map.items()
if not pattern.startswith("_") and precision in QUANTIZATION_CONFIGS
]
unmatched_patterns = [
pattern
for pattern in requested_patterns
if not any(matches_pattern(path, pattern) for path in quantized_modules)
]
if unmatched_patterns:
raise RuntimeError(
"MXFP8 precision rules matched no modules: " + ", ".join(unmatched_patterns)
)
if not list(staging_dir.glob("*.safetensors")):
raise RuntimeError("mlx-vlm conversion produced no safetensors weights")
normalize_unlimited_ocr_metadata(staging_dir)
(staging_dir / "precision_map.json").write_text(
json.dumps(precision_map, indent=2, ensure_ascii=False) + "\n",
encoding="utf-8",
)
(staging_dir / "quantization_summary.json").write_text(
json.dumps({
"method": "mxfp8",
"group_size": 32,
"bits": 8,
"source_model": (
Path(model_path).name if Path(model_path).is_dir() else model_path
),
"source_revision": source_revision,
"precision_map_sha256": _json_digest(precision_map),
"quantized_module_count": len(set(quantized_modules)),
"quantized_precision_counts": {
precision: sum(
selected == precision
for selected in quantized_precisions.values()
)
for precision in sorted(set(quantized_precisions.values()))
},
"preserved_quantizable_module_count": len(set(preserved_modules)),
"quantized_modules": sorted(set(quantized_modules)),
"quantized_module_precisions": dict(sorted(quantized_precisions.items())),
}, indent=2),
encoding="utf-8",
)
staging_dir.replace(output_dir)
finally:
if staging_dir.exists():
shutil.rmtree(staging_dir)
if verbose:
print(f"\nQuantized {len(set(quantized_modules))} module(s).")
print(f"Model saved to: {output_dir}")
return output_dir
def main():
parser = argparse.ArgumentParser(
description="Apply OCR-aware mixed-precision quantization"
)
parser.add_argument("--model-path", required=True,
help="Source model (baidu/Unlimited-OCR or local path)")
parser.add_argument(
"--source-revision",
default=None,
help="Immutable source commit used for remote loading and provenance",
)
parser.add_argument("--precision-map", type=Path,
default=Path(__file__).parent / "precision_map.json",
help="Path to precision_map.json")
parser.add_argument("--output-dir", type=Path, default=Path("./enhanced_model"),
help="Output directory for enhanced model")
parser.add_argument(
"--verbose",
action=argparse.BooleanOptionalAction,
default=True,
help="Show conversion progress (use --no-verbose to suppress)",
)
args = parser.parse_args()
print("=" * 60)
print("Unlimited-OCR Mixed-Precision Conversion")
print("=" * 60)
precision_map = load_precision_map(args.precision_map)
convert_model(
args.model_path,
precision_map,
args.output_dir,
args.verbose,
source_revision=args.source_revision,
)
if __name__ == "__main__":
main()