danielfein's picture
Add training support package
a4019dd verified
Raw
History Blame Contribute Delete
1.64 kB
from __future__ import annotations
from dataclasses import dataclass
from pathlib import Path
import torch
@dataclass(slots=True)
class TokenCheckpoint:
new_token: str
token_id: int | None
embedding: torch.Tensor
loss_history: list[float]
secondary_embeddings: list[torch.Tensor] | None = None
def save_token_checkpoint(
*,
token: str,
token_id: int,
embedding: torch.Tensor,
loss_history: list[float],
path: Path,
secondary_embeddings: list[torch.Tensor] | None = None,
) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
torch.save(
{
"new_token": token,
"token_id": token_id,
"embedding": embedding.detach().cpu(),
"loss_history": list(loss_history),
"secondary_embeddings": (
[row.detach().cpu() for row in secondary_embeddings]
if secondary_embeddings
else None
),
},
path,
)
def load_token_checkpoint(path: Path) -> TokenCheckpoint:
payload = torch.load(path, map_location="cpu", weights_only=True)
token = payload.get("new_token", payload.get("token"))
if token is None:
raise KeyError(f"Checkpoint {path} has no token metadata")
return TokenCheckpoint(
new_token=str(token),
token_id=payload.get("token_id"),
embedding=payload["embedding"].detach().cpu(),
loss_history=list(payload.get("loss_history", [])),
secondary_embeddings=[
row.detach().cpu()
for row in (payload.get("secondary_embeddings") or [])
] or None,
)