Views
No views yet
1{
2 'step': step,
3 'config': asdict(model.config),
4 'model_state_dict': model.state_dict(),
5},1checkpoint = torch.load(path, weights_only=True)
2
3config = GPTConfig(**checkpoint['config'])
4model = GPT(config)
5
6any_key = next(iter(checkpoint['model_state_dict'].keys()))
7if any_key.startswith("_orig_mod."):
8 # strip "_orig_mod." if the model was compiled
9 model_state_dict = {k[10:]: v for k, v in checkpoint['model_state_dict'].items()}
10else:
11 model_state_dict = checkpoint['model_state_dict']
12
13model.load_state_dict(model_state_dict)