grn-safetensor / grn_convert.py
flaroche's picture
Upload grn_convert.py
1fbadbb verified
Raw History Blame Contribute Delete
1.2 kB
"""GRN's pickle checkpoints to safetensors, on the CPU, tensors only (weights_only):
the transformer to BF16 (the released pipeline runs its blocks in BF16), the
tokenizer as stored. Usage: grn_convert.py <in.pth|ckpt> <out.safetensors> [bf16]"""
import sys, torch
from safetensors.torch import save_file
src, dst = sys.argv[1], sys.argv[2]
state = torch.load(src, map_location="cpu", weights_only=True)
for key in ("trainer", "gpt_fsdp", "ema", "vae", "state_dict"):
while isinstance(state, dict) and key in state and isinstance(state[key], dict):
print("descending into", key, "of", list(state.keys())[:6])
state = state[key]
tensors = {k: v for k, v in state.items() if isinstance(v, torch.Tensor)}
print(len(tensors), "tensors;", [k for k in state if k not in tensors][:5])
if len(sys.argv) > 3:
tensors = {k: (v.to(torch.bfloat16) if v.is_floating_point() else v) for k, v in tensors.items()}
import re, collections
seen = collections.OrderedDict()
for k, v in tensors.items():
seen.setdefault(re.sub(r"\.\d+\.", ".N.", k), (tuple(v.shape), str(v.dtype)))
for k, v in seen.items(): print(" ", k, v)
save_file({k: v.contiguous() for k, v in tensors.items()}, dst)