Dolphin3.0-CoreML / scripts /validate_release.py
ales27pm's picture
Make the Core ML release self-contained and runnable
95671cf verified
Raw
History Blame Contribute Delete
18.3 kB
#!/usr/bin/env python3
# /// script
# requires-python = ">=3.11,<3.12"
# dependencies = [
# "coremltools==8.0",
# "numpy==1.26.4",
# ]
# ///
"""Validate the repository contract and a stateful Dolphin Core ML package."""
from __future__ import annotations
import argparse
import hashlib
import json
import platform
import shutil
import subprocess
import tempfile
from pathlib import Path
from typing import Any
import coremltools as ct
SOURCE_REVISION = "392a6f57223e7ccfe6ef4ebdb2ff101a42d57364"
EXPORT_REPORT = (
"validation/"
"Dolphin3.0-Llama3.2-3B-stateful-int4.export-report.json"
)
EXPECTED_TOKENIZER_HASHES = {
"config.json": "e21ff53ea39726f972362beba869807216775d5e308bc2f531784846c06a0249",
"generation_config.json": "e627b5a8b2dc371f90388947ada64fa6e71de0f991c04c835f0c0bc97e305a4f",
"special_tokens_map.json": "2df2c4620bb1a9eb877bc7c90c7fa04608bda9fa7c0cf2cdcc0a17b849649683",
"tokenizer.json": "e40b93124a3e29f62d5f4ff41be56cb2af34ecacf9239acd9da53a98860380b5",
"tokenizer_config.json": "51ad9580aba8d00016efda43357185a0d8ff9884584dcc82ab58ca552afd14e1",
}
REQUIRED_REPOSITORY_FILES = (
"README.md",
"LICENSE",
"USE_POLICY.md",
"NOTICE",
"coreml_artifacts.json",
EXPORT_REPORT,
"validation/tiny-stateful-runtime-smoke.json",
"requirements.txt",
"scripts/export_stateful_coreml.py",
"scripts/generate.py",
"scripts/validate_release.py",
"examples/swift/DolphinCoreMLCLI/Package.swift",
"examples/swift/DolphinCoreMLCLI/Package.resolved",
"examples/swift/DolphinCoreMLCLI/Sources/DolphinCoreMLCLI/main.swift",
)
RECOMMENDED_ARTIFACT = "Dolphin3.0-Llama3.2-3B-stateful-int4.mlpackage"
EXPECTED_LEGACY_ARTIFACT_BYTES = {
"Dolphin3.0-Llama3.2-3B-fp16.mlpackage": 6_455_796_893,
"Dolphin3.0-Llama3.2-3B-int8.mlpackage": 3_230_380_696,
"Dolphin3.0-Llama3.2-3B-int4-lut.mlpackage": 1_614_959_144,
}
def sha256(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as handle:
for chunk in iter(lambda: handle.read(8 * 1024 * 1024), b""):
digest.update(chunk)
return digest.hexdigest()
def package_inventory(path: Path) -> dict[str, Any]:
files = []
total = 0
for item in sorted(candidate for candidate in path.rglob("*") if candidate.is_file()):
size = item.stat().st_size
total += size
files.append(
{
"path": str(item.relative_to(path)),
"bytes": size,
"sha256": sha256(item),
}
)
return {"bytes": total, "files": files}
def validate_repository(root: Path) -> dict[str, Any]:
missing = [name for name in REQUIRED_REPOSITORY_FILES if not (root / name).is_file()]
if missing:
raise RuntimeError(f"Missing repository files: {missing}")
readme = (root / "README.md").read_text()
if "license: llama3.2" not in readme.split("---", 2)[1]:
raise RuntimeError("README front matter must declare license: llama3.2")
if SOURCE_REVISION not in readme:
raise RuntimeError("README does not pin the source revision")
if "iOS 18" not in readme or "macOS 15" not in readme:
raise RuntimeError("README does not state the actual deployment targets")
requirements = (root / "requirements.txt").read_text().splitlines()
for dependency in ("accelerate==1.2.1", "jinja2==3.1.5"):
if dependency not in requirements:
raise RuntimeError(f"Missing pinned runtime dependency: {dependency}")
actual_hashes = {}
for name, expected in EXPECTED_TOKENIZER_HASHES.items():
path = root / name
if not path.is_file():
raise RuntimeError(f"Missing tokenizer/config asset: {name}")
actual = sha256(path)
actual_hashes[name] = actual
if actual != expected:
raise RuntimeError(f"Hash mismatch for {name}: {actual} != {expected}")
config = json.loads((root / "config.json").read_text())
generation = json.loads((root / "generation_config.json").read_text())
tokenizer_config = json.loads((root / "tokenizer_config.json").read_text())
if config.get("vocab_size") != 128258:
raise RuntimeError("Unexpected tokenizer/model vocabulary size")
expected_stops = [128256, 128001, 128008, 128009]
if generation.get("eos_token_id") != expected_stops:
raise RuntimeError("generation_config.json has an unexpected stop-token contract")
chat_template = tokenizer_config.get("chat_template", "")
if "<|im_start|>assistant" not in chat_template:
raise RuntimeError("tokenizer_config.json lacks the Dolphin assistant template")
return {
"required_files": list(REQUIRED_REPOSITORY_FILES),
"tokenizer_hashes": actual_hashes,
"vocab_size": config["vocab_size"],
"stop_token_ids": expected_stops,
}
def validate_artifact_manifest(
root: Path, artifact_path: Path, artifact_inventory: dict[str, Any]
) -> dict[str, Any]:
manifest_path = root / "coreml_artifacts.json"
manifest = json.loads(manifest_path.read_text())
if manifest.get("schema_version") != 2:
raise RuntimeError("coreml_artifacts.json must use schema_version 2")
if manifest.get("recommended") != RECOMMENDED_ARTIFACT:
raise RuntimeError(
f"Manifest must recommend {RECOMMENDED_ARTIFACT!r}"
)
source = manifest.get("source") or {}
if source.get("revision") != SOURCE_REVISION:
raise RuntimeError("Manifest does not pin the expected source revision")
artifacts = manifest.get("artifacts")
if not isinstance(artifacts, list) or not artifacts:
raise RuntimeError("Manifest must contain a non-empty artifacts list")
by_name: dict[str, dict[str, Any]] = {}
for entry in artifacts:
if not isinstance(entry, dict) or not isinstance(entry.get("file"), str):
raise RuntimeError("Manifest artifact entries must be named objects")
name = entry["file"]
if name in by_name:
raise RuntimeError(f"Duplicate manifest artifact: {name}")
by_name[name] = entry
recommended = by_name.get(RECOMMENDED_ARTIFACT)
if recommended is None:
raise RuntimeError("Manifest does not inventory the recommended artifact")
if artifact_path.name != RECOMMENDED_ARTIFACT:
raise RuntimeError(
f"Validated artifact must be named {RECOMMENDED_ARTIFACT!r}"
)
if recommended.get("status") != "recommended":
raise RuntimeError("Recommended artifact has an unexpected status")
if recommended.get("stateful") is not True:
raise RuntimeError("Recommended artifact must declare stateful: true")
if recommended.get("quantization") != "int4-per-block-linear":
raise RuntimeError("Recommended artifact has an unexpected quantization")
manifest_inventory = {
"bytes": recommended.get("bytes"),
"files": recommended.get("files"),
}
if manifest_inventory != artifact_inventory:
raise RuntimeError(
"Recommended artifact inventory does not match the validated package"
)
for name, expected_bytes in EXPECTED_LEGACY_ARTIFACT_BYTES.items():
entry = by_name.get(name)
if entry is None:
raise RuntimeError(f"Manifest is missing legacy artifact {name}")
if entry.get("status") != "legacy" or entry.get("bytes") != expected_bytes:
raise RuntimeError(f"Legacy artifact metadata mismatch for {name}")
return {
"schema_version": manifest["schema_version"],
"recommended": manifest["recommended"],
"source": source,
"artifact_count": len(artifacts),
"recommended_inventory": manifest_inventory,
"legacy_artifact_bytes": EXPECTED_LEGACY_ARTIFACT_BYTES,
}
def validate_export_report(
root: Path, artifact_inventory: dict[str, Any]
) -> dict[str, Any]:
report = json.loads((root / EXPORT_REPORT).read_text())
expected = {
"schema_version": 1,
"artifact": RECOMMENDED_ARTIFACT,
"tiny_test": False,
"quantization": "int4",
"max_context_length": 2048,
"max_query_length": 512,
"state_names": ["keyCache", "valueCache"],
}
mismatches = {
key: {"expected": value, "actual": report.get(key)}
for key, value in expected.items()
if report.get(key) != value
}
source = report.get("source") or {}
if source.get("revision") != SOURCE_REVISION:
mismatches["source.revision"] = {
"expected": SOURCE_REVISION,
"actual": source.get("revision"),
}
if report.get("inventory") != artifact_inventory:
mismatches["inventory"] = "does not match the validated package"
parity = report.get("torch_kv_cache_parity") or {}
for metric in ("max_abs_error", "mean_abs_error"):
value = parity.get(metric)
if not isinstance(value, (int, float)) or value < 0 or value > 0.005:
mismatches[f"torch_kv_cache_parity.{metric}"] = {
"expected": "finite value between 0 and 0.005",
"actual": value,
}
if mismatches:
raise RuntimeError(f"Export report mismatches: {mismatches}")
return {
"path": EXPORT_REPORT,
"source": source,
"artifact": report["artifact"],
"tiny_test": report["tiny_test"],
"quantization": report["quantization"],
"torch_kv_cache_parity": parity,
"inventory": report["inventory"],
}
def range_bounds(feature: Any, dimension: int) -> tuple[int, int]:
ranges = feature.type.multiArrayType.shapeRange.sizeRanges
return int(ranges[dimension].lowerBound), int(ranges[dimension].upperBound)
def mil_tensor_shape(tensor_type: Any) -> list[int | None]:
shape: list[int | None] = []
for dimension in tensor_type.dimensions:
shape.append(int(dimension.constant.size) if dimension.HasField("constant") else None)
return shape
def program_output_type(spec: Any, name: str) -> tuple[Any, int]:
function = spec.mlProgram.functions["main"]
block = function.block_specializations[function.opset]
if name not in block.outputs:
raise RuntimeError(f"ML Program does not declare {name!r} as an output")
quantized_weight_ops = sum(
operation.type == "constexpr_blockwise_shift_scale"
for operation in block.operations
)
for operation in reversed(block.operations):
for output in operation.outputs:
if output.name == name:
return output.type.tensorType, quantized_weight_ops
raise RuntimeError(f"ML Program has no typed value for output {name!r}")
def validate_package(path: Path) -> dict[str, Any]:
model = ct.models.MLModel(str(path), skip_model_load=True)
spec = model.get_spec()
description = spec.description
inputs = {item.name: item for item in description.input}
outputs = {item.name: item for item in description.output}
states = {item.name: item for item in description.state}
if set(inputs) != {"inputIds", "causalMask"}:
raise RuntimeError(f"Unexpected input schema: {sorted(inputs)}")
if set(outputs) != {"logits"}:
raise RuntimeError(f"Unexpected output schema: {sorted(outputs)}")
if set(states) != {"keyCache", "valueCache"}:
raise RuntimeError(f"Unexpected state schema: {sorted(states)}")
if spec.specificationVersion != 9:
raise RuntimeError(
f"Expected Core ML specification version 9, got {spec.specificationVersion}"
)
feature_types = ct.proto.FeatureTypes_pb2.ArrayFeatureType
if inputs["inputIds"].type.multiArrayType.dataType != feature_types.INT32:
raise RuntimeError("inputIds must use Int32 values")
if inputs["causalMask"].type.multiArrayType.dataType != feature_types.FLOAT16:
raise RuntimeError("causalMask must use Float16 values")
if outputs["logits"].type.multiArrayType.dataType != feature_types.FLOAT16:
raise RuntimeError("logits must use Float16 values")
query_bounds = range_bounds(inputs["inputIds"], 1)
mask_query_bounds = range_bounds(inputs["causalMask"], 2)
context_bounds = range_bounds(inputs["causalMask"], 3)
if query_bounds != (1, 512) or mask_query_bounds != (1, 512):
raise RuntimeError(f"Unexpected query ranges: {query_bounds}, {mask_query_bounds}")
if context_bounds != (1, 2048):
raise RuntimeError(f"Unexpected context range: {context_bounds}")
expected_state_shape = [28, 1, 8, 2048, 128]
state_shapes = {}
for name, state in states.items():
array_type = state.type.stateType.arrayType
shape = [int(dimension) for dimension in array_type.shape]
state_shapes[name] = shape
if shape != expected_state_shape:
raise RuntimeError(
f"Unexpected {name} shape: {shape} != {expected_state_shape}"
)
if array_type.dataType != feature_types.FLOAT16:
raise RuntimeError(f"{name} must use Float16 values")
logits_type, quantized_weight_ops = program_output_type(spec, "logits")
logits_shape = mil_tensor_shape(logits_type)
if logits_type.dataType != ct.proto.MIL_pb2.DataType.FLOAT16:
raise RuntimeError("ML Program logits must use Float16 values")
if logits_shape != [1, None, 128258]:
raise RuntimeError(
f"Unexpected ML Program logits shape: {logits_shape} != [1, *, 128258]"
)
if quantized_weight_ops < 1:
raise RuntimeError("ML Program contains no blockwise quantized weight operations")
metadata = dict(description.metadata.userDefined)
expected_metadata = {
"co.huggingface.exporters.name": "ales27pm/Dolphin3.0-CoreML",
"com.ales27pm.dolphin.source_revision": SOURCE_REVISION,
"com.ales27pm.dolphin.max_context_length": "2048",
"com.ales27pm.dolphin.max_query_length": "512",
"com.ales27pm.dolphin.cache": "stateful-key-value",
"com.ales27pm.dolphin.quantization": "int4",
}
mismatches = {
key: {"expected": value, "actual": metadata.get(key)}
for key, value in expected_metadata.items()
if metadata.get(key) != value
}
if mismatches:
raise RuntimeError(f"Core ML metadata mismatches: {mismatches}")
return {
"specification_version": spec.specificationVersion,
"inputs": sorted(inputs),
"outputs": sorted(outputs),
"states": sorted(states),
"state_shapes": state_shapes,
"input_dtypes": {"inputIds": "int32", "causalMask": "float16"},
"logits": {"dtype": "float16", "shape": logits_shape},
"blockwise_quantized_weight_ops": quantized_weight_ops,
"query_range": query_bounds,
"context_range": context_bounds,
"metadata": expected_metadata,
"inventory": package_inventory(path),
}
def run_compiler(path: Path) -> dict[str, Any]:
compiler = shutil.which("xcrun")
if compiler is None:
raise RuntimeError("xcrun is unavailable; compiler validation requires macOS/Xcode")
with tempfile.TemporaryDirectory(prefix="dolphin-coreml-compile-") as temporary:
destination = Path(temporary)
generated = destination / "generated"
compiled = destination / "compiled"
generated.mkdir()
compiled.mkdir()
commands = [
["xcrun", "coremlcompiler", "metadata", str(path)],
[
"xcrun",
"coremlcompiler",
"generate",
str(path),
str(generated),
"--language",
"Swift",
"--platform",
"macos",
"--deployment-target",
"15.0",
],
[
"xcrun",
"coremlcompiler",
"compile",
str(path),
str(compiled),
"--platform",
"macOS",
"--deployment-target",
"15.0",
],
]
results = []
for command in commands:
completed = subprocess.run(
command, check=False, capture_output=True, text=True
)
results.append(
{
"command": command,
"exit_code": completed.returncode,
"stdout": completed.stdout,
"stderr": completed.stderr,
}
)
if completed.returncode != 0:
raise RuntimeError(
f"coremlcompiler failed ({completed.returncode}): "
f"{completed.stderr or completed.stdout}"
)
return {"commands": results}
def main() -> int:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("artifact", type=Path)
parser.add_argument("--repo-root", type=Path, default=Path(__file__).parents[1])
parser.add_argument("--compile", action="store_true")
parser.add_argument("--output", type=Path)
args = parser.parse_args()
repository = validate_repository(args.repo_root.resolve())
artifact = validate_package(args.artifact.resolve())
artifact_manifest = validate_artifact_manifest(
args.repo_root.resolve(), args.artifact.resolve(), artifact["inventory"]
)
export_report = validate_export_report(
args.repo_root.resolve(), artifact["inventory"]
)
report = {
"schema_version": 1,
"status": "passed",
"repository": repository,
"artifact_manifest": artifact_manifest,
"export_report": export_report,
"artifact": artifact,
"compiler": run_compiler(args.artifact.resolve()) if args.compile else None,
"environment": {
"platform": platform.platform(),
"machine": platform.machine(),
"coremltools": ct.__version__,
},
}
rendered = json.dumps(report, indent=2, sort_keys=True) + "\n"
if args.output:
args.output.write_text(rendered)
print(rendered, end="")
return 0
if __name__ == "__main__":
raise SystemExit(main())