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()
|