File size: 2,017 Bytes
63a1291
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Export the canonical model.png to model.safetensors (standard format).

model.png is the canonical model. This produces the exact same weights in a
standard container, with verifiable parameter-count metadata baked into the
safetensors header - so INFERENCE.py can run from safetensors alone and match
main.py output bit for bit.

    python convert_to_safetensors.py            # model.png -> model.safetensors
"""

from __future__ import annotations

import argparse
import json

import torch
from safetensors.torch import save_file

from model import load_config, load_model_png, param_breakdown, PAD_ID


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--png", default="model.png")
    ap.add_argument("--config", default="config.json")
    ap.add_argument("--out", default="model.safetensors")
    args = ap.parse_args()

    cfg = load_config(args.config)
    model = load_model_png(args.png, cfg, map_location="cpu")

    info = param_breakdown(model)
    total = info["total_parameters"]
    n_bias = sum(int(t.numel()) for n, t in model.named_parameters() if n.endswith("bias"))

    state = {k: v.contiguous().cpu() for k, v in model.state_dict().items()}

    metadata = {
        "format": "pt",
        "model": "PixelModel-v3",
        "total_parameters": str(total),
        "param_breakdown": json.dumps(info["param_breakdown"]),
        "has_bias": "true" if n_bias > 0 else "false",
        "bias_parameters": str(n_bias),
        "text_encoder_parameters": str(
            sum(int(t.numel()) for n, t in model.named_parameters()
                if n.startswith("embed") or n.startswith("text_"))),
        "vae_parameters": "0",
        "config": json.dumps(cfg.to_json()),
    }

    save_file(state, args.out, metadata=metadata)
    print(f"[convert] {args.png} -> {args.out}")
    print(f"[convert] total_parameters = {total:,}  (bias={n_bias:,})")
    print(f"[convert] config + param_breakdown embedded in safetensors metadata")


if __name__ == "__main__":
    main()