model-a-scratch / mla /checkpoint.py
karthik-2905's picture
Upload folder using huggingface_hub
ae9e4fe verified
Raw
History Blame Contribute Delete
1.38 kB
import json
from dataclasses import asdict
import numpy as np
from .backend import to_numpy, to_device
from .model import Config, Model
from .optim import AdamW
def save_checkpoint(path, model, opt, step):
arrays = {}
params = model.parameters()
for i, p in enumerate(params):
arrays[f"p{i}"] = to_numpy(p.data)
for i in range(len(opt.params)):
arrays[f"m{i}"] = to_numpy(opt.m[i])
arrays[f"v{i}"] = to_numpy(opt.v[i])
meta = {
"config": asdict(model.cfg),
"step": step,
"n_params": len(params),
"opt": {
"lr": opt.lr, "b1": opt.b1, "b2": opt.b2,
"eps": opt.eps, "wd": opt.wd, "t": opt.t,
},
}
arrays["meta"] = np.array(json.dumps(meta))
np.savez(path, **arrays)
def load_checkpoint(path):
data = np.load(path, allow_pickle=False)
meta = json.loads(data["meta"].item())
cfg = Config(**meta["config"])
model = Model(cfg)
params = model.parameters()
for i, p in enumerate(params):
p.data = to_device(data[f"p{i}"])
o = meta["opt"]
opt = AdamW(params, lr=o["lr"], betas=(o["b1"], o["b2"]),
eps=o["eps"], weight_decay=o["wd"])
opt.t = o["t"]
for i in range(len(params)):
opt.m[i] = to_device(data[f"m{i}"])
opt.v[i] = to_device(data[f"v{i}"])
return model, opt, meta["step"]