import torch import torch.nn as nn import torch.nn.functional as F from torch.cuda.amp import autocast import torchvision.transforms as T from transformers import ViTModel from PIL import Image import gradio as gr from pathlib import Path class EthiopianCropViT(nn.Module): def __init__(self, num_classes, dropout=0.3): super().__init__() self.vit = ViTModel.from_pretrained("google/vit-base-patch16-224-in21k") hidden = self.vit.config.hidden_size self.classifier = nn.Sequential( nn.LayerNorm(hidden), nn.Dropout(dropout), nn.Linear(hidden, 512), nn.GELU(), nn.BatchNorm1d(512), nn.Dropout(dropout*0.7), nn.Linear(512, 256), nn.GELU(), nn.BatchNorm1d(256), nn.Dropout(dropout*0.5), nn.Linear(256, num_classes) ) def forward(self, x): return self.classifier(self.vit(x).last_hidden_state[:, 0, :]) device = torch.device("cuda" if torch.cuda.is_available() else "cpu") print(f"Device: {device}") checkpoint = torch.load("ethiopian_crop_model.pth", map_location=device) classes = checkpoint['classes'] idx_to_class = {int(k): v for k, v in checkpoint['idx_to_class'].items()} cfg = checkpoint['config'] model = EthiopianCropViT(num_classes=len(classes)).to(device) model.load_state_dict(checkpoint['model_state_dict']) model.eval() print(f"Model loaded: {len(classes)} classes") transform = T.Compose([ T.Resize((224, 224)), T.ToTensor(), T.Normalize(mean=cfg['mean'], std=cfg['std']), ]) def predict_disease(image): if image is None: return {"Error": "Please upload an image"} try: img_tensor = transform(image).unsqueeze(0).to(device) with torch.no_grad(): if device.type == "cuda": with autocast(): logits = model(img_tensor) else: logits = model(img_tensor) probs = F.softmax(logits, dim=1)[0] top_probs, top_indices = torch.topk(probs, min(3, len(probs))) results = {} for prob, idx in zip(top_probs, top_indices): class_name = idx_to_class[idx.item()] crop = class_name.split("_")[0].title() disease = " ".join(class_name.split("_")[1:]).title() if "_" in class_name else "Unknown" emoji = "✅" if "healthy" in class_name.lower() else "⚠️" results[f"{emoji} {crop} - {disease}"] = float(prob) return results except Exception as e: return {"Error": str(e)} iface = gr.Interface( fn=predict_disease, inputs=gr.Image(type="pil", label="Upload Crop Photo"), outputs=gr.Label(num_top_classes=3, label="Predictions"), title="Ethiopian Crop Disease Detector", description="Detect diseases in Coffee, Teff, Maize, Wheat, Sorghum, Barley, Khat, Oilseeds, Pulses, Potato and Tomato", theme=gr.themes.Soft() ) if __name__ == "__main__": print("Starting Ethiopian Crop Disease Detector...") iface.launch(server_name="0.0.0.0", server_port=7860)