(args, model_without_ddp, optimizer, loss_scaler)
| 314 | |
| 315 | |
| 316 | def load_model(args, model_without_ddp, optimizer, loss_scaler): |
| 317 | if args.resume: |
| 318 | if args.resume.startswith('https'): |
| 319 | checkpoint = torch.hub.load_state_dict_from_url( |
| 320 | args.resume, map_location='cpu', check_hash=True) |
| 321 | else: |
| 322 | checkpoint = torch.load(args.resume, map_location='cpu', weights_only=False) |
| 323 | model_without_ddp.load_state_dict(checkpoint["state_dict"]) |
| 324 | print("Resume checkpoint %s" % args.resume) |
| 325 | if 'optimizer' in checkpoint and 'epoch' in checkpoint and not (hasattr(args, 'eval') and args.eval): |
| 326 | optimizer.load_state_dict(checkpoint['optimizer']) |
| 327 | args.start_epoch = checkpoint['epoch'] + 1 |
| 328 | if loss_scaler is not None and loss_scaler.state_dict_key in checkpoint: |
| 329 | loss_scaler.load_state_dict(checkpoint[loss_scaler.state_dict_key]) |
| 330 | print("With optim & sched!") |
| 331 | |
| 332 | def load_old_ckpt(model, checkpoint_path, verbose=True): |
| 333 | """ |
nothing calls this directly
no test coverage detected