Views
No views yet
1import torch
2from improved_model import ImprovedPix2Pix3D
3
4# Load checkpoint
5checkpoint = torch.load("last.ckpt", map_location='cpu')
6model = ImprovedPix2Pix3D(**checkpoint['hyper_parameters'])
7model.load_state_dict(checkpoint['state_dict'])
8model.eval()
9
10# Inference
11with torch.no_grad():
12 output = model(voided_image, mask)last.ckpt: Complete model checkpoint with weights and hyperparametersimproved_model.py: Model architecture code