Views
No views yet
1import torch
2from torchvision import transforms
3
4# Load model
5model = torch.load("model.pth")
6model.eval()
7
8# Prepare image
9transform = transforms.Compose([
10 transforms.Resize((224, 224)),
11 transforms.ToTensor(),
12 transforms.Normalize(mean=[0.485, 0.456, 0.406],
13 std=[0.229, 0.224, 0.225])
14])
15
16# Inference (example for a single image tensor 'input_tensor')
17with torch.no_grad():
18 # input_tensor = transform(image).unsqueeze(0)
19 output = model(input_tensor)
20 prediction = torch.argmax(output, dim=1)1@article{nourbakhsh2025kd,
2 title={KD-OCT: Efficient Knowledge Distillation for Clinical-Grade Retinal OCT Classification},
3 author={Nourbakhsh, Erfan and Sanjari, Nasrin and Nourbakhsh, Ali},
4 journal={arXiv preprint arXiv:2512.09069},
5 year={2025}
6}