Merge nucleic/sleek-ember-seal-uady into dev

This commit is contained in:
2026-07-31 03:05:29 -07:00
parent 9a1228efbb
commit 4ed4763557
4 changed files with 131 additions and 10 deletions
+41
View File
@@ -351,6 +351,47 @@ def load_pretrained_weights(
}
def load_checkpoint_weights(
model: ModernBertForPurposeClassification,
checkpoint: Path,
) -> dict[str, int]:
"""Strictly restore a trained purpose-deep checkpoint, including task heads."""
if not checkpoint.is_file():
raise DataError(f"{checkpoint}: purpose-deep checkpoint is missing")
weights = mx.load(str(checkpoint))
parameters = dict(tree_flatten(model.parameters()))
missing = sorted(set(parameters) - set(weights))
unexpected = sorted(set(weights) - set(parameters))
if missing or unexpected:
details = []
if missing:
details.append(
f"missing {len(missing)} tensors ({', '.join(missing[:3])})"
)
if unexpected:
details.append(
f"has {len(unexpected)} unexpected tensors "
f"({', '.join(unexpected[:3])})"
)
raise DataError(
f"{checkpoint}: trained checkpoint " + " and ".join(details)
)
for key, parameter in parameters.items():
if tuple(weights[key].shape) != tuple(parameter.shape):
raise DataError(
f"purpose-deep tensor {key} has shape {weights[key].shape}; "
f"expected {parameter.shape}"
)
model.load_weights(list(weights.items()), strict=True)
mx.eval(model.parameters())
return {
"loaded": len(weights),
"ignored": 0,
"freshTaskHeads": 0,
}
def save_weights(model: ModernBertForPurposeClassification, checkpoint: Path) -> None:
mx.eval(model.parameters())
mx.save_safetensors(