Merge nucleic/sleek-ember-seal-uady into dev
This commit is contained in:
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user