Download grn_convert.py from drift-generator/grn-safetensor: direct link, hf CLI and curl.
- Browser
- Download file 1.2 kB
-
https://huggingface.co/drift-generator/grn-safetensor/resolve/main/grn_convert.py
- Command line
-
hf download hf://drift-generator/grn-safetensor/grn_convert.py
-
curl -L -o grn_convert.py https://huggingface.co/drift-generator/grn-safetensor/resolve/main/grn_convert.py
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) | |