Views
No views yet
1import torch
2import timm
3from torchvision import transforms
4from PIL import Image
5
6# Load model
7model = timm.create_model('mobilevit_xs', num_classes=2, pretrained=False)
8model.load_state_dict(torch.load('pytorch_model.bin'))
9model.eval()
10
11# Prepare image
12transform = transforms.Compose([
13 transforms.Resize((256, 256)),
14 transforms.ToTensor(),
15 transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
16])
17
18img = transform(Image.open('frame.webp').convert('RGB')).unsqueeze(0).cuda()
19
20# Predict
21with torch.no_grad():
22 logits = model(img)
23 probs = torch.softmax(logits, dim=1)
24 garbage_prob = probs[0, 0].item() # Class 0 = garbage
25
26# Decision
27is_garbage = garbage_prob > 0.6315 # Use optimal threshold