| 526 | |
| 527 | |
| 528 | def load_model(filepath, args=None, device="cpu", **kwargs): |
| 529 | ckpt = torch.load(filepath, map_location="cpu") |
| 530 | if args is None: |
| 531 | args = ckpt["hyper_parameters"] |
| 532 | |
| 533 | for key, value in kwargs.items(): |
| 534 | if not key in args: |
| 535 | rank_zero_warn(f"Unknown hyperparameter: {key}={value}") |
| 536 | args[key] = value |
| 537 | |
| 538 | model = create_model(args) |
| 539 | state_dict = {re.sub(r"^model\.", "", k): v for k, v in ckpt["state_dict"].items()} |
| 540 | model.load_state_dict(state_dict) |
| 541 | |
| 542 | return model.to(device) |
| 543 | |
| 544 | |
| 545 | class ViSNet(nn.Module): |