Views
No views yet
1import torch
2from model_drbitron_b32 import DRBiTronB32
3
4# Load segmentation model
5seg_model = torch.load('MM-UNet-checkpoint.pth')
6model = DRBiTronB32(seg_model=seg_model, embedding_dim=256)
7
8# Load best checkpoint
9checkpoint = torch.load('best_smoothl1_model.pth')
10model.load_state_dict(checkpoint['model_state_dict'], strict=False)
11model.eval()
12
13# Inference
14with torch.no_grad():
15 logits = model(img_224, img_1024)
16 probs = torch.softmax(logits, dim=1)
17 pred = torch.argmax(logits, dim=1)