File size: 2,009 Bytes
d2d7586
 
 
2f501a4
d2d7586
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2f501a4
 
 
 
 
d2d7586
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
60
61
62
63
64
65
66
67
"""Convert a project ``.pt`` checkpoint into Hugging Face format.

Example:
    python convert_checkpoint.py --checkpoint ../../checkpoints/policy/on-policy/step-0006000.pt
"""

from __future__ import annotations

import argparse
from pathlib import Path
import sys

import torch

HERE = Path(__file__).resolve().parent
ROOT = HERE.parents[1]
if str(HERE) not in sys.path:
    sys.path.insert(0, str(HERE))

from configuration_chess_policy import ChessPolicyConfig  # noqa: E402
from modeling_chess_policy import ChessTransitionPolicy  # noqa: E402


def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument(
        "--checkpoint",
        type=Path,
        default=ROOT
        / "checkpoints"
        / "policy"
        / "on-policy"
        / "step-0006000.pt",
    )
    parser.add_argument("--output-dir", type=Path, default=HERE)
    parser.add_argument(
        "--safe-serialization",
        action=argparse.BooleanOptionalAction,
        default=True,
        help="Write model.safetensors instead of pytorch_model.bin.",
    )
    args = parser.parse_args()

    checkpoint = torch.load(args.checkpoint, map_location="cpu", weights_only=False)
    if "model" not in checkpoint or "model_config" not in checkpoint:
        raise ValueError("Checkpoint must contain model and model_config entries")

    config = ChessPolicyConfig(**checkpoint["model_config"])
    config.architectures = ["ChessTransitionPolicy"]
    config.auto_map = {
        "AutoConfig": "configuration_chess_policy.ChessPolicyConfig",
        "AutoModel": "modeling_chess_policy.ChessTransitionPolicy",
    }
    model = ChessTransitionPolicy(config)
    model.load_state_dict(checkpoint["model"])
    model.eval()
    args.output_dir.mkdir(parents=True, exist_ok=True)
    model.save_pretrained(
        args.output_dir,
        safe_serialization=args.safe_serialization,
    )
    print(f"Saved Hugging Face model to {args.output_dir}")


if __name__ == "__main__":
    main()