File size: 1,644 Bytes
a4019dd
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
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,
    )