Views
No views yet
1import torch
2import timm
3from PIL import Image
4from torchvision import transforms
5
6# Load model
7model = timm.create_model('vit_base_patch16_224', pretrained=False, num_classes=9)
8checkpoint = torch.load('best_model.pth', map_location='cpu')
9model.load_state_dict(checkpoint['model_state_dict'])
10model.eval()
11
12# Preprocessing
13transform = transforms.Compose([
14 transforms.Resize((224, 224)),
15 transforms.ToTensor(),
16 transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
17])
18
19# Inference
20image = Image.open('waste_item.jpg').convert('RGB')
21input_tensor = transform(image).unsqueeze(0)
22
23with torch.no_grad():
24 outputs = model(input_tensor)
25 probabilities = torch.nn.functional.softmax(outputs, dim=1)
26 predicted_class = torch.argmax(probabilities, dim=1).item()
27
28categories = ['Cardboard', 'Food Organics', 'Glass', 'Metal', 'Miscellaneous Trash', 'Paper', 'Plastic', 'Textile Trash', 'Vegetation']
29print(f"Predicted: {categories[predicted_class]}")