Merge nucleic/sleek-ember-seal-uady into dev
This commit is contained in:
@@ -0,0 +1,120 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Verify pinned Hugging Face ModernBERT -> MLX backbone parity."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import Sequence
|
||||
|
||||
import numpy as np
|
||||
|
||||
from deep_contract import DEEP_VARIANTS, validate_variant_config
|
||||
from purpose_data import DataError
|
||||
from train_deep_mlx import _configure_mlx_device, _resolve_source
|
||||
|
||||
|
||||
def verify(variant_name: str, local_model: Path | None, device: str) -> float:
|
||||
try:
|
||||
import mlx.core as mx
|
||||
import torch
|
||||
from transformers import ModernBertForMaskedLM
|
||||
|
||||
from deep_model_mlx import (
|
||||
ModernBertForPurposeClassification,
|
||||
ModernBertPurposeConfig,
|
||||
load_pretrained_weights,
|
||||
)
|
||||
except ImportError as exc:
|
||||
raise DataError(
|
||||
"deep parity requires PyTorch, Transformers, and MLX"
|
||||
) from exc
|
||||
|
||||
_configure_mlx_device(mx, device)
|
||||
variant = DEEP_VARIANTS[variant_name]
|
||||
source = _resolve_source(variant, local_model)
|
||||
try:
|
||||
config_json = json.loads(
|
||||
(source / "config.json").read_text(encoding="utf-8")
|
||||
)
|
||||
except (OSError, json.JSONDecodeError) as exc:
|
||||
raise DataError(f"cannot read ModernBERT config: {exc}") from exc
|
||||
validate_variant_config(config_json, variant)
|
||||
|
||||
torch.set_num_threads(1)
|
||||
reference = ModernBertForMaskedLM.from_pretrained(
|
||||
source,
|
||||
local_files_only=True,
|
||||
attn_implementation="eager",
|
||||
)
|
||||
reference.eval()
|
||||
candidate = ModernBertForPurposeClassification(
|
||||
ModernBertPurposeConfig.from_hugging_face(
|
||||
config_json, gradient_checkpointing=False
|
||||
)
|
||||
)
|
||||
report = load_pretrained_weights(candidate, source / "model.safetensors")
|
||||
candidate.eval()
|
||||
|
||||
# 96 tokens crosses the local layer's 64-token half-window, so this catches
|
||||
# both full-attention and sliding-window-mask parity without a slow 512-token
|
||||
# CPU reference pass.
|
||||
rng = np.random.default_rng(20260731)
|
||||
input_ids = rng.integers(
|
||||
3, int(config_json["vocab_size"]) - 1, size=(2, 96), dtype=np.int32
|
||||
)
|
||||
input_ids[:, 0] = int(config_json["bos_token_id"])
|
||||
input_ids[:, -1] = int(config_json["eos_token_id"])
|
||||
attention_mask = np.ones_like(input_ids, dtype=np.int32)
|
||||
attention_mask[0, -11:] = 0
|
||||
input_ids[0, -11:] = int(config_json["pad_token_id"])
|
||||
|
||||
with torch.inference_mode():
|
||||
sequence = reference.model(
|
||||
input_ids=torch.from_numpy(input_ids.astype(np.int64)),
|
||||
attention_mask=torch.from_numpy(attention_mask.astype(np.int64)),
|
||||
).last_hidden_state
|
||||
reference_pooled = reference.head(sequence[:, 0]).cpu().numpy()
|
||||
candidate_sequence = candidate.model(
|
||||
mx.array(input_ids), mx.array(attention_mask)
|
||||
)
|
||||
candidate_pooled = candidate.head(candidate_sequence[:, 0])
|
||||
mx.eval(candidate_pooled)
|
||||
error = float(
|
||||
np.max(np.abs(reference_pooled - np.asarray(candidate_pooled)))
|
||||
)
|
||||
print(
|
||||
f"{variant.model_id}@{variant.revision}: "
|
||||
f"loaded={report['loaded']} ignored={report['ignored']} "
|
||||
f"pooled_max_abs_error={error:.3g}",
|
||||
flush=True,
|
||||
)
|
||||
if not np.isfinite(error) or error > 5e-4:
|
||||
raise DataError(
|
||||
f"ModernBERT MLX parity failed: max abs error {error:.6g} > 0.0005"
|
||||
)
|
||||
return error
|
||||
|
||||
|
||||
def build_parser() -> argparse.ArgumentParser:
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("--variant", choices=tuple(DEEP_VARIANTS), default="base")
|
||||
parser.add_argument("--model", type=Path)
|
||||
parser.add_argument("--device", choices=("metal", "cpu"), default="metal")
|
||||
return parser
|
||||
|
||||
|
||||
def main(argv: Sequence[str] | None = None) -> int:
|
||||
args = build_parser().parse_args(argv)
|
||||
try:
|
||||
verify(args.variant, args.model, args.device)
|
||||
except (DataError, OSError, RuntimeError, ValueError) as exc:
|
||||
print(f"error: {exc}", file=sys.stderr)
|
||||
return 1
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
Reference in New Issue
Block a user