PixelModel-v3 / convert_to_safetensors.py
TobiasLogic's picture
PixelModel v3: SIREN+FiLM CPPN, 919K params, beats v1 FID
63a1291 verified
Raw
History Blame
2.02 kB
"""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()