{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":25563,"databundleVersionId":2094376,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":272173710,"sourceType":"kernelVersion"},{"sourceId":613792,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":461186,"modelId":476952}],"dockerImageVersionId":31153,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"🌿 Plant Disease Classification — Kaggle Plant Pathology 2021\n\nMulti-label image classification using PyTorch & EfficientNet-B0.\nNotebook trains an EfficientNet-B0 model on 18K+ leaf images, including full data pipeline, training logs, and an inference demo (upload an image to test predictions).\n\nDataset: https://www.kaggle.com/c/plant-pathology-2021\n","metadata":{}},{"cell_type":"code","source":"# Environment & imports\nimport os\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nfrom sklearn.preprocessing import MultiLabelBinarizer\nfrom sklearn.metrics import f1_score, accuracy_score\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as T\nimport torchvision.models as models\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nprint(\"✅ imports ok — torch:\", torch.__version__)","metadata":{"_uuid":"880b022d-1048-49c1-a46d-c7ca400b91f6","_cell_guid":"fa001f5d-51a0-4b9b-aad7-304cd92af960","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-10-30T19:08:21.704309Z","iopub.execute_input":"2025-10-30T19:08:21.704786Z","iopub.status.idle":"2025-10-30T19:08:21.714375Z","shell.execute_reply.started":"2025-10-30T19:08:21.704701Z","shell.execute_reply":"2025-10-30T19:08:21.713221Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df = pd.read_csv(\"/kaggle/input/plant-pathology-2021-fgvc8/train.csv\")\ntrain_df['labels'] = train_df['labels'].apply(lambda x: x.split())\nmlb = MultiLabelBinarizer()\ny = mlb.fit_transform(train_df['labels'])\nprint(\"Classes:\", mlb.classes_)\ntrain_df.head()","metadata":{"_uuid":"ea478dd1-744b-44be-b736-bfefdf06676f","_cell_guid":"b8a3ae81-1d99-42db-b2c5-5d01bcef4507","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-10-30T19:08:24.312893Z","iopub.execute_input":"2025-10-30T19:08:24.313548Z","iopub.status.idle":"2025-10-30T19:08:24.719779Z","shell.execute_reply.started":"2025-10-30T19:08:24.313525Z","shell.execute_reply":"2025-10-30T19:08:24.718814Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# class counts (multi-label -> count each label presence)\nfrom collections import Counter\ncnt = Counter([label for labels in train_df['labels'] for label in labels])\npd.Series(dict(cnt)).sort_values(ascending=False).plot.bar(figsize=(8,4))\nplt.title(\"Label counts\")\n\n# show a grid of sample images\ndef show_samples(df, img_dir, n=8):\n    fig, axes = plt.subplots(2, n//2, figsize=(12,5))\n    for ax, (_, row) in zip(axes.flatten(), df.sample(n).iterrows()):\n        img = Image.open(os.path.join(img_dir, row.image)).convert(\"RGB\")\n        ax.imshow(img)\n        ax.set_title(\", \".join(row.labels))\n        ax.axis('off')\n    plt.tight_layout()\n\nshow_samples(train_df, \"/kaggle/input/plant-pathology-2021-fgvc8/train_images/\")","metadata":{"_uuid":"c5296302-a430-43d5-8f03-1daad8aebdf4","_cell_guid":"5af004a5-f9d3-4820-a3d3-2fad3f164e1b","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-10-30T19:08:27.160875Z","iopub.execute_input":"2025-10-30T19:08:27.161154Z","iopub.status.idle":"2025-10-30T19:08:37.142783Z","shell.execute_reply.started":"2025-10-30T19:08:27.161139Z","shell.execute_reply":"2025-10-30T19:08:37.141604Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class PlantDataset(Dataset):\n    def __init__(self, df, labels, img_dir, transform=None):\n        self.df = df.reset_index(drop=True)\n        self.labels = labels\n        self.img_dir = img_dir\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        img_path = os.path.join(self.img_dir, self.df.iloc[idx, 0])\n        image = np.array(Image.open(img_path).convert(\"RGB\"))\n        label = torch.tensor(self.labels[idx]).float()\n        if self.transform:\n            image = self.transform(image=image)['image']\n        return image, label","metadata":{"_uuid":"17ee8504-c2a1-4d8b-9c49-8e173bf64a0d","_cell_guid":"de70b31d-87bc-42d3-a035-daf2b4d429f9","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-10-30T19:08:46.243166Z","iopub.execute_input":"2025-10-30T19:08:46.243401Z","iopub.status.idle":"2025-10-30T19:08:46.248624Z","shell.execute_reply.started":"2025-10-30T19:08:46.243389Z","shell.execute_reply":"2025-10-30T19:08:46.247889Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_transform = A.Compose([\n    A.Resize(224, 224),\n    A.HorizontalFlip(p=0.5),\n    A.RandomBrightnessContrast(p=0.2),\n    A.Normalize(),\n    ToTensorV2(),\n])\n\ntrain_dataset = PlantDataset(\n    train_df, y,\n    \"/kaggle/input/plant-pathology-2021-fgvc8/train_images/\",\n    transform=train_transform\n)\ntrain_loader = DataLoader(train_dataset, batch_size=16, shuffle=True, num_workers=2)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-30T19:08:56.450646Z","iopub.execute_input":"2025-10-30T19:08:56.450907Z","iopub.status.idle":"2025-10-30T19:08:56.461121Z","shell.execute_reply.started":"2025-10-30T19:08:56.450890Z","shell.execute_reply":"2025-10-30T19:08:56.460166Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class_names = mlb.classes_\nimport matplotlib.pyplot as plt\n\n# Sample 4 random rows\nsample_idx = train_df.sample(4, random_state=42).index\nsample = PlantDataset(train_df.loc[sample_idx], y[sample_idx], \"/kaggle/input/plant-pathology-2021-fgvc8/train_images/\", transform=train_transform)\n\n# Plot\nfor i, (img, lbl) in enumerate(sample):\n    plt.figure(figsize=(3,3))\n    plt.imshow(img.permute(1, 2, 0).numpy())\n    \n    # Convert multi-hot label vector to readable class names\n    active_classes = [class_names[j] for j, val in enumerate(lbl.numpy()) if val == 1.0]\n    plt.title(\", \".join(active_classes))\n    plt.axis('off')\n    plt.show()","metadata":{"_uuid":"df6f5264-c77b-4798-bd70-f06cde4da315","_cell_guid":"431875f8-23b9-4b09-8975-7fe8144c993e","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-10-30T19:08:49.292379Z","iopub.execute_input":"2025-10-30T19:08:49.292630Z","iopub.status.idle":"2025-10-30T19:08:49.800761Z","shell.execute_reply.started":"2025-10-30T19:08:49.292613Z","shell.execute_reply":"2025-10-30T19:08:49.799881Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# instantiate model\nmodel = models.efficientnet_b0(weights=None)\nmodel.classifier[1] = nn.Linear(model.classifier[1].in_features, len(mlb.classes_))\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel = model.to(device)\n\ncriterion = nn.BCEWithLogitsLoss()\noptimizer = torch.optim.Adam(model.parameters(), lr=1e-4)\n\n# simple training loop (show only final loss)\nfor epoch in range(3):\n    model.train()\n    running_loss = 0.0\n    for images, labels in train_loader:\n        images, labels = images.to(device), labels.to(device)\n        optimizer.zero_grad()\n        outputs = model(images)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        running_loss += loss.item()\n    print(f\"Epoch [{epoch+1}/3], Loss: {running_loss/len(train_loader):.4f}\")","metadata":{"_uuid":"ebf8d13d-4499-44e4-8292-5b68e9286680","_cell_guid":"cc7d51d9-e27e-4067-be8a-f62e0d4cf78e","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# load weights if saved\n# model.load_state_dict(torch.load(\"plant_model.pth\", map_location=device))\nmodel.eval()\n\ndef predict_pil(img_pil, threshold=0.5):\n    img = np.array(img_pil.convert(\"RGB\"))\n    transform = A.Compose([A.Resize(224,224), A.Normalize(), ToTensorV2()])\n    x = transform(image=img)['image'].unsqueeze(0).to(device)\n    with torch.no_grad():\n        logits = model(x)\n        probs = torch.sigmoid(logits).cpu().numpy()[0]\n    labels = [mlb.classes_[i] for i,v in enumerate(probs) if v >= threshold]\n    return labels, probs\n\n# show predictions on random samples\nfor idx in train_df.sample(4).index:\n    img = Image.open(os.path.join(\"/kaggle/input/plant-pathology-2021-fgvc8/train_images/\", train_df.loc[idx,'image']))\n    labels, probs = predict_pil(img)\n    display(img.resize((256,256)))\n    print(\"True:\", train_df.loc[idx,'labels'], \"Predicted:\", labels)","metadata":{"_uuid":"2931260b-795a-41a3-a2eb-d6a53c722b7b","_cell_guid":"741e8969-8fa2-44c6-9e7b-a12af811dbcc","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-10-30T19:24:52.482041Z","iopub.status.idle":"2025-10-30T19:24:52.482284Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.save(model.state_dict(), \"plant_model_demo.pth\")\n# Save a small CSV of a few sample predictions to display in notebook\nsample_preds = []\nfor idx in train_df.sample(200).index:\n    img_path = os.path.join(\"/kaggle/input/plant-pathology-2021-fgvc8/train_images/\", train_df.loc[idx,'image'])\n    labels, probs = predict_pil(Image.open(img_path))\n    sample_preds.append((train_df.loc[idx,'image'], train_df.loc[idx,'labels'], labels))\npd.DataFrame(sample_preds, columns=[\"image\",\"true\",\"pred\"]).to_csv(\"sample_predictions.csv\", index=False)","metadata":{"_uuid":"dee69c0a-35f0-4809-99fa-ce4efa3409a4","_cell_guid":"e00e099d-1395-4430-adb8-380a9db4e6c2","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-10-30T19:24:52.485754Z","iopub.status.idle":"2025-10-30T19:24:52.485937Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 🧠 Inference Demo\nThis section shows how to use the trained EfficientNet-B0 model to predict diseases on unseen leaf images.","metadata":{"_uuid":"1222b9cd-067b-4f75-84e4-e6c05c6263ca","_cell_guid":"fe852140-f065-4f1b-b2d6-77e817ab7ea3","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"import torch\nfrom PIL import Image\nimport numpy as np\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport matplotlib.pyplot as plt\n\n# Load model (must match your training setup)\nmodel = models.efficientnet_b0(weights=None)\nmodel.classifier[1] = nn.Linear(model.classifier[1].in_features, 6)  # 6 = number of classes\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nmodel.load_state_dict(torch.load(\"/kaggle/input/plant-model/pytorch/default/1/plant_model.pth\", map_location=device))\nmodel = model.to(device)\nmodel.eval()\n\n# Label names (same order used during training)\nlabels = ['complex', 'frog_eye_leaf_spot', 'healthy', 'powdery_mildew', 'rust', 'scab']\n\n# Transform for inference\ninference_transform = A.Compose([\n    A.Resize(224, 224),\n    A.Normalize(),\n    ToTensorV2(),\n])\n\ndef predict_image(image_path, threshold=0.5):\n    img = np.array(Image.open(image_path).convert(\"RGB\"))\n    img_tensor = inference_transform(image=img)[\"image\"].unsqueeze(0).to(device)\n    with torch.no_grad():\n        logits = model(img_tensor)\n        probs = torch.sigmoid(logits).cpu().numpy()[0]\n    predicted_labels = [labels[i] for i, p in enumerate(probs) if p >= threshold]\n    return predicted_labels, probs\n\n# Pick a few random test images\ntest_dir = \"/kaggle/input/plant-pathology-2021-fgvc8/train_images/\"\nsample_images = [\n    \"fffe472a0001bd25.jpg\",  # healthy\n    \"fffe105cf6808292.jpg\",  # scab frog_eye_leaf_spot\n    \"fffc94e092a59086.jpg\",  # rust\n]\n\nfor img_name in sample_images:\n    path = test_dir + img_name\n    preds, probs = predict_image(path)\n    plt.imshow(Image.open(path))\n    plt.axis('off')\n    plt.title(f\"Predicted: {preds}\")\n    plt.show()","metadata":{"_uuid":"6caa42e2-ede7-495b-809b-0d4e0d3f83ef","_cell_guid":"f47dfde1-3a48-49c8-802c-4c1dc7f8813b","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-10-30T19:24:56.606940Z","iopub.execute_input":"2025-10-30T19:24:56.607253Z","iopub.status.idle":"2025-10-30T19:24:59.375701Z","shell.execute_reply.started":"2025-10-30T19:24:56.607237Z","shell.execute_reply":"2025-10-30T19:24:59.374813Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# For testing uploaded image\n\n# from PIL import Image\n# import torchvision.transforms as transforms\n\n# # Example: test on an uploaded image\n# img_path = \"/kaggle/input/test-image.jpg\"  # update this with your test image\n# img = Image.open(img_path).convert(\"RGB\")\n\n# transform = transforms.Compose([\n#     transforms.Resize((224, 224)),\n#     transforms.ToTensor(),\n#     transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])\n# ])\n\n# input_tensor = transform(img).unsqueeze(0).to(device)\n# with torch.no_grad():\n#     outputs = model(input_tensor)\n#     _, predicted = torch.max(outputs, 1)\n\n# print(\"Predicted class index:\", predicted.item())","metadata":{"_uuid":"de19d144-aa77-47b2-a899-c21e40f821a0","_cell_guid":"1c18026a-5464-47f0-bffa-d7215b7f060c","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-10-30T19:07:33.983934Z","iopub.status.idle":"2025-10-30T19:07:33.984227Z","shell.execute_reply.started":"2025-10-30T19:07:33.984079Z","shell.execute_reply":"2025-10-30T19:07:33.984094Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null}]}