Download src/convert_bf16.py from Reizxn/makeitwork: direct link, hf CLI and curl.
- Browser
- Download file 876 Bytes
-
https://huggingface.co/Reizxn/makeitwork/resolve/main/src/convert_bf16.py
- Command line
-
hf download hf://Reizxn/makeitwork/src/convert_bf16.py
-
curl -L -o convert_bf16.py https://huggingface.co/Reizxn/makeitwork/resolve/main/src/convert_bf16.py
876 Bytes
| """Convert the latest checkpoint to bf16 to halve file size for Colab upload.""" | |
| import torch | |
| import json | |
| import os | |
| def main(): | |
| ckpt = torch.load("checkpoints/latest.pt", map_location="cpu", weights_only=False) | |
| print(f"Original checkpoint: step={ckpt['step']}, loss={ckpt['loss']}") | |
| print(f" Embedding shape: {ckpt['model_state_dict']['token_embedding.weight'].shape}") | |
| # Convert all float32 tensors to bf16 | |
| state = ckpt["model_state_dict"] | |
| for key in state: | |
| if state[key].dtype == torch.float32: | |
| state[key] = state[key].to(torch.bfloat16) | |
| ckpt["model_state_dict"] = state | |
| out_path = "checkpoints/latest_bf16.pt" | |
| torch.save(ckpt, out_path) | |
| size_mb = os.path.getsize(out_path) / 1e6 | |
| print(f"bf16 checkpoint saved: {out_path} ({size_mb:.1f} MB)") | |
| if __name__ == "__main__": | |
| main() | |