"""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 [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)