{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":33679,"databundleVersionId":3212216,"sourceType":"competition"},{"sourceId":14113087,"sourceType":"datasetVersion","datasetId":8990132}],"dockerImageVersionId":31193,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport json\nimport time\nfrom math import inf\n\nimport numpy as np\nimport pandas as pd\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import confusion_matrix, precision_recall_fscore_support, roc_auc_score, average_precision_score\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\n\nfrom PIL import Image\nimport torchvision.transforms as transforms\nimport torchvision.models as models\n\n#import torch.cuda.amp as amp\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-12-04T00:22:28.557123Z","iopub.execute_input":"2025-12-04T00:22:28.557955Z","iopub.status.idle":"2025-12-04T00:22:28.563293Z","shell.execute_reply.started":"2025-12-04T00:22:28.557924Z","shell.execute_reply":"2025-12-04T00:22:28.562468Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DATA_DIR = \"/kaggle/input/herbarium-2022-fgvc9\"\ntrain_meta_path = os.path.join(DATA_DIR, \"train_metadata.json\")\ntest_meta_path  = os.path.join(DATA_DIR, \"test_metadata.json\")\nprint(\"DATA_DIR:\", DATA_DIR)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-04T00:22:28.564751Z","iopub.execute_input":"2025-12-04T00:22:28.565263Z","iopub.status.idle":"2025-12-04T00:22:28.578718Z","shell.execute_reply.started":"2025-12-04T00:22:28.565245Z","shell.execute_reply":"2025-12-04T00:22:28.578192Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Prepare","metadata":{}},{"cell_type":"code","source":"with open(train_meta_path, \"r\") as f:\n    train_meta = json.load(f)\n\nann_df    = pd.DataFrame(train_meta[\"annotations\"])\nimg_df    = pd.DataFrame(train_meta[\"images\"])\ncat_df    = pd.DataFrame(train_meta[\"categories\"])\n\nmerged = ann_df.merge(img_df, on=\"image_id\", how=\"left\")\nmerged = merged.merge(cat_df, on=\"category_id\", how=\"left\")\nmerged[\"image_path\"] = merged[\"file_name\"].apply(lambda fn: os.path.join(DATA_DIR, \"train_images\", fn))\n\ncore_cols = [\n    \"image_id\",\n    \"image_path\",\n    \"category_id\",\n    \"genus_id\",\n    \"family\",\n    \"genus\",\n    \"species\",\n    \"scientificName\",\n]\ncore_df = merged[core_cols].copy()\nprint(\"core_df shape:\", core_df.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-04T00:22:28.579769Z","iopub.execute_input":"2025-12-04T00:22:28.579958Z","iopub.status.idle":"2025-12-04T00:22:39.475015Z","shell.execute_reply.started":"2025-12-04T00:22:28.579943Z","shell.execute_reply":"2025-12-04T00:22:39.474308Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Building dataFrames ","metadata":{}},{"cell_type":"code","source":"    toxic_genus_list = [\n        \"Toxicodendron\",  # poison ivy, oak, sumac\n        \"Euphorbia\",      # spurges, skin irritants\n        \"Urtica\",         # nettles\n        \"Cicuta\",         # water hemlock\n        \"Conium\",         # poison hemlock\n        \"Heracleum\",      # hogweed, burns\n    ]\n    \n    toxic_csv_path = os.path.join(DATA_DIR, \"toxic_species.csv\")\n    \n    if os.path.exists(toxic_csv_path):\n        print(\"Found toxic_species.csv — using curated list to label examples.\")\n        tox_df = pd.read_csv(toxic_csv_path)\n        toxic_scientific = set()\n        toxic_species = set()\n        toxic_genus = set()\n    \n        if \"scientificName\" in tox_df.columns:\n            toxic_scientific.update(list(tox_df[\"scientificName\"].dropna().astype(str).str.strip()))\n        if \"species\" in tox_df.columns:\n            toxic_species.update(list(tox_df[\"species\"].dropna().astype(str).str.strip()))\n        if \"genus\" in tox_df.columns:\n            toxic_genus.update(list(tox_df[\"genus\"].dropna().astype(str).str.strip()))\n    \n        def is_toxic_row(row):\n            if row[\"scientificName\"] in toxic_scientific:\n                return True\n            if str(row[\"species\"]) in toxic_species:\n                return True\n            if str(row[\"genus\"]) in toxic_genus:\n                return True\n            return False\n    \n        core_df[\"do_not_touch\"] = core_df.apply(is_toxic_row, axis=1).astype(int)\n    else:\n        print(\"No toxic_species.csv found — falling back to genus-level list.\")\n        core_df[\"do_not_touch\"] = core_df[\"genus\"].isin(toxic_genus_list).astype(int)\n    \n    print(\"do_not_touch value counts (overall):\")\n    print(core_df[\"do_not_touch\"].value_counts())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-04T00:22:39.475736Z","iopub.execute_input":"2025-12-04T00:22:39.475942Z","iopub.status.idle":"2025-12-04T00:22:39.529686Z","shell.execute_reply.started":"2025-12-04T00:22:39.475926Z","shell.execute_reply":"2025-12-04T00:22:39.528857Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 5% Stratisfied test set ","metadata":{}},{"cell_type":"code","source":"TEST_FRACTION = 0.05\nstrat_col = core_df[\"do_not_touch\"]\nrest_df, test_df = train_test_split(\n    core_df,\n    test_size=TEST_FRACTION,\n    random_state=42,\n    stratify=strat_col,\n)\nprint(\"Realistic test size:\", len(test_df))\nprint(test_df[\"do_not_touch\"].value_counts())\n\ntoxic_df = rest_df[rest_df[\"do_not_touch\"] == 1].reset_index(drop=True)\nsafe_df  = rest_df[rest_df[\"do_not_touch\"] == 0].reset_index(drop=True)\n\nprint(\"Toxic (train pool) count:\", len(toxic_df))\nprint(\"Safe  (train pool) count:\", len(safe_df))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-04T00:22:39.53168Z","iopub.execute_input":"2025-12-04T00:22:39.532071Z","iopub.status.idle":"2025-12-04T00:22:40.373067Z","shell.execute_reply.started":"2025-12-04T00:22:39.532044Z","shell.execute_reply":"2025-12-04T00:22:40.372209Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Img transformations + Dataset loader ","metadata":{}},{"cell_type":"code","source":"image_transforms = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225],\n    ),\n])\n\nstrong_transforms = transforms.Compose([\n    transforms.RandomResizedCrop(224, scale=(0.75, 1.0)),\n    transforms.RandomHorizontalFlip(),\n    transforms.ColorJitter(\n        brightness=0.25,\n        contrast=0.25,\n        saturation=0.25,\n        hue=0.03,\n    ),\n    transforms.RandomRotation(10),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225],\n    ),\n])\n\nclass PlantDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df.reset_index(drop=True)\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img_path = row[\"image_path\"]\n        try:\n            img = Image.open(img_path).convert(\"RGB\")\n        except Exception:\n            img = Image.new(\"RGB\", (224, 224), (0, 0, 0))\n        if self.transform:\n            img = self.transform(img)\n        toxic_label = int(row[\"do_not_touch\"])\n        return img, torch.tensor(toxic_label, dtype=torch.long)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-04T00:22:40.374537Z","iopub.execute_input":"2025-12-04T00:22:40.374744Z","iopub.status.idle":"2025-12-04T00:22:40.400577Z","shell.execute_reply.started":"2025-12-04T00:22:40.374728Z","shell.execute_reply":"2025-12-04T00:22:40.399819Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Model Precision, Recall, F1, and AUC ","metadata":{}},{"cell_type":"code","source":"from sklearn.metrics import precision_recall_fscore_support, confusion_matrix, roc_auc_score, average_precision_score\n\ndef compute_metrics(y_true, y_probs, threshold=0.5):\n    y_pred = (y_probs >= threshold).astype(int)\n    precision, recall, f1, _ = precision_recall_fscore_support(y_true, y_pred, average='binary', zero_division=0)\n    cm = confusion_matrix(y_true, y_pred)\n    metrics = {}\n    metrics[\"precision\"] = float(precision)\n    metrics[\"recall\"] = float(recall)\n    metrics[\"f1\"] = float(f1)\n    metrics[\"confusion_matrix\"] = cm.tolist()\n    try:\n        metrics[\"roc_auc\"] = float(roc_auc_score(y_true, y_probs))\n    except Exception:\n        metrics[\"roc_auc\"] = None\n    try:\n        metrics[\"pr_auc\"] = float(average_precision_score(y_true, y_probs))\n    except Exception:\n        metrics[\"pr_auc\"] = None\n    return metrics","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-04T00:22:40.40157Z","iopub.execute_input":"2025-12-04T00:22:40.40184Z","iopub.status.idle":"2025-12-04T00:22:40.409574Z","shell.execute_reply.started":"2025-12-04T00:22:40.401816Z","shell.execute_reply":"2025-12-04T00:22:40.408994Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### EfficientNet and Training config","metadata":{}},{"cell_type":"code","source":"device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nprint(\"Using device:\", device)\n\n# base with weight\nmodel = models.efficientnet_b0(weights=models.EfficientNet_B0_Weights.IMAGENET1K_V1)\nin_features = model.classifier[1].in_features\nmodel.classifier = nn.Sequential(\n    nn.Dropout(0.2),\n    nn.Linear(in_features, 1)  # binary output logit\n)\nmodel = model.to(device)\n\ncriterion = nn.BCEWithLogitsLoss()\noptimizer = torch.optim.Adam(model.parameters(), lr=1e-4)  # can tweak later if needed\n\n# learning rate scheduler and early stopping\nscheduler = ReduceLROnPlateau(optimizer, mode='max', factor=0.5, patience=2)\nPATIENCE = 20  # how many rotations without improvement I'm okay with\nbest_val_f1 = 0.0\npatience_counter = 0\n\n# load previously trained checkpoint (the one i gave my presentation on) to continue training \nPRETRAINED_CKPT_PATH = \"/kaggle/input/presentedtoclasscheckpoints/final_rich_checkpoint.pth\"\n\nif os.path.exists(PRETRAINED_CKPT_PATH):\n    print(f\"Loading pretrained checkpoint from: {PRETRAINED_CKPT_PATH}\")\n    ckpt = torch.load(PRETRAINED_CKPT_PATH, map_location=device)\n\n    if isinstance(ckpt, dict) and \"model_state_dict\" in ckpt:\n        model.load_state_dict(ckpt[\"model_state_dict\"])\n        try:\n            optimizer.load_state_dict(ckpt[\"optimizer_state_dict\"])\n        except Exception as e:\n            print(\"sadly could not load optimizer state, using fresh instead! \", e)\n        print(\"Checkpoint loaded (model + optimizer if compatible).\")\n    else:\n        model.load_state_dict(ckpt)\n        print(\"loaded raw state_dict checkpoint.\")\nelse:\n    print(\"no pretrained checkpoint found, starting from ImageNet weights.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-04T00:22:40.410277Z","iopub.execute_input":"2025-12-04T00:22:40.410469Z","iopub.status.idle":"2025-12-04T00:22:40.424872Z","shell.execute_reply.started":"2025-12-04T00:22:40.410454Z","shell.execute_reply":"2025-12-04T00:22:40.424345Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Building the training / validation / test loaders","metadata":{}},{"cell_type":"code","source":"BATCH_SIZE = 64        \nNUM_WORKERS = 4\nNUM_ROTATIONS = 40 # did 20 last run, double it.\nSAFE_PER_ROTATION = len(toxic_df)\n\n# build initial balanced_df and train/val split\nsafe_sample_initial = safe_df.sample(n=len(toxic_df), random_state=42)\nbalanced_df = pd.concat([toxic_df, safe_sample_initial]).sample(frac=1.0, random_state=42).reset_index(drop=True)\ntrain_df, val_df = train_test_split(balanced_df, test_size=0.2, random_state=42, stratify=balanced_df[\"do_not_touch\"])\n\ntrain_dataset = PlantDataset(train_df, transform=image_transforms)\nval_dataset   = PlantDataset(val_df,   transform=image_transforms)\ntrain_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True, num_workers=NUM_WORKERS)\nval_loader   = DataLoader(val_dataset,   batch_size=BATCH_SIZE, shuffle=False, num_workers=NUM_WORKERS)\n\n# test loader\ntest_dataset = PlantDataset(test_df, transform=image_transforms)\ntest_loader  = DataLoader(test_dataset, batch_size=BATCH_SIZE, shuffle=False, num_workers=NUM_WORKERS)\n\nprint(\"Train batches:\", len(train_loader))\nprint(\"Val batches:\", len(val_loader))\nprint(\"Realistic test size (samples):\", len(test_df))\nprint(\"Realistic test batches:\", len(test_loader))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-04T00:22:40.425663Z","iopub.execute_input":"2025-12-04T00:22:40.426039Z","iopub.status.idle":"2025-12-04T00:22:40.805268Z","shell.execute_reply.started":"2025-12-04T00:22:40.426016Z","shell.execute_reply":"2025-12-04T00:22:40.804564Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### dataloader for my balanced train / val","metadata":{}},{"cell_type":"code","source":"def evaluate_loader_timed(model, loader, device, use_amp=True, print_progress_every=500):\n    model.eval()\n    all_probs = []\n    all_labels = []\n    batch_times = []\n    start = time.time()\n    with torch.no_grad():\n        for i, (images, labels) in enumerate(loader, start=1):\n            t0 = time.time()\n            images = images.to(device)\n            labels = labels.to(device)\n            if use_amp and device.startswith(\"cuda\"):\n                with torch.amp.autocast(\"cuda\"):\n                    outputs = model(images)\n            else:\n                outputs = model(images)\n\n            probs = torch.sigmoid(outputs).squeeze(1).cpu().numpy()\n            all_probs.append(probs)\n            all_labels.append(labels.cpu().numpy())\n            batch_time = time.time() - t0\n            batch_times.append(batch_time)\n            if (i % print_progress_every) == 0:\n                elapsed = time.time() - start\n                avg_b = sum(batch_times) / len(batch_times)\n                processed = i * loader.batch_size\n                rate = processed / elapsed if elapsed > 0 else 0.0\n                print(f\"[eval] batches={i} avg_batch_s={avg_b:.4f} processed={processed} rate_samples/s={rate:.2f}\")\n    if len(all_probs) == 0:\n        return np.array([]), np.array([]), 0.0\n    all_probs = np.concatenate(all_probs, axis=0)\n    all_labels = np.concatenate(all_labels, axis=0)\n    avg_batch_time = sum(batch_times) / len(batch_times) if batch_times else 0.0\n    return all_labels, all_probs, avg_batch_time","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-04T00:22:40.805969Z","iopub.execute_input":"2025-12-04T00:22:40.8062Z","iopub.status.idle":"2025-12-04T00:22:41.116188Z","shell.execute_reply.started":"2025-12-04T00:22:40.806178Z","shell.execute_reply":"2025-12-04T00:22:41.115547Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### My sanity check before I set it to run and find an issue","metadata":{}},{"cell_type":"code","source":"EPOCHS_INITIAL = 0 ##no longer need :)\nfor epoch in range(EPOCHS_INITIAL):\n    model.train()\n    running_loss = 0.0\n    n = 0\n    for images, labels in train_loader:\n        images = images.to(device)\n        labels = labels.float().unsqueeze(1).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() * images.size(0)\n        n += images.size(0)\n    avg_train_loss = running_loss / n\n    val_labels, val_probs, _ = evaluate_loader_timed(model, val_loader, device, use_amp=True, print_progress_every=200)\n    val_metrics = compute_metrics(val_labels, val_probs, threshold=0.5)\n    print(f\"Initial Epoch {epoch+1}: Train Loss {avg_train_loss:.4f} | Val metrics:\", val_metrics)\n    ckpt = {\n        \"model_state_dict\": model.state_dict(),\n        \"optimizer_state_dict\": optimizer.state_dict(),\n        \"epoch\": epoch,\n        \"val_metrics\": val_metrics,\n    }\n    torch.save(ckpt, \"initial_checkpointv2.pth\")\n    print(\"Saved initial checkpoint: initial_checkpointv2.pth\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-04T00:22:41.116874Z","iopub.execute_input":"2025-12-04T00:22:41.117076Z","iopub.status.idle":"2025-12-04T00:22:41.122984Z","shell.execute_reply.started":"2025-12-04T00:22:41.117053Z","shell.execute_reply":"2025-12-04T00:22:41.122266Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"safe_shuffled = safe_df.sample(frac=1.0, random_state=123).reset_index(drop=True)\nrot_model_path = \"balanced_do_not_touch_model_rotating_safe_40rotations_On_20_rotations.pth\"\nbest_rot_val_f1 = 0.0\n\nfor rot in range(NUM_ROTATIONS):\n    print(\"=\"*60)\n    print(f\"=== ROTATION {rot+1}/{NUM_ROTATIONS} ===\")\n    start = (rot * SAFE_PER_ROTATION) % len(safe_shuffled)\n    end = start + SAFE_PER_ROTATION\n    if end <= len(safe_shuffled):\n        safe_chunk = safe_shuffled.iloc[start:end]\n    else:\n        safe_chunk = pd.concat([safe_shuffled.iloc[start:], safe_shuffled.iloc[:end - len(safe_shuffled)]]).reset_index(drop=True)\n\n    train_chunk_df = pd.concat([toxic_df, safe_chunk]).sample(frac=1.0, random_state=42 + rot).reset_index(drop=True)\n    train_dataset_rot = PlantDataset(train_chunk_df, transform=strong_transforms)\n    train_loader_rot = DataLoader(train_dataset_rot, batch_size=BATCH_SIZE, shuffle=True, num_workers=NUM_WORKERS)\n\n    # train one pass over this chunk\n    model.train()\n    train_loss = 0.0\n    train_total = 0\n    for images, labels in train_loader_rot:\n        images = images.to(device)\n        labels = labels.float().unsqueeze(1).to(device)\n        optimizer.zero_grad()\n        outputs = model(images)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        train_loss += loss.item() * images.size(0)\n        train_total += images.size(0)\n    avg_train_loss = train_loss / train_total\n\n    # evaluate on balanced val\n    val_labels, val_probs, _ = evaluate_loader_timed(model, val_loader, device, use_amp=True, print_progress_every=200)\n    val_metrics = compute_metrics(val_labels, val_probs, threshold=0.5)\n\n    print(f\"Rotation {rot+1}/{NUM_ROTATIONS}: Train Loss {avg_train_loss:.4f} | Val metrics: {val_metrics}\")\n\n    ckpt = {\n        \"model_state_dict\": model.state_dict(),\n        \"optimizer_state_dict\": optimizer.state_dict(),\n        \"rotation\": rot,\n        \"val_metrics\": val_metrics,\n    }\n    torch.save(ckpt, f\"rotation_{rot+1}_checkpoint.pth\")\n\n    current_val_f1 = val_metrics.get(\"f1\", 0.0)\n    scheduler.step(current_val_f1)\n\n    if current_val_f1 > best_val_f1 + 1e-6:\n        best_val_f1 = current_val_f1\n        patience_counter = 0\n        best_rot_val_f1 = current_val_f1\n        torch.save(ckpt, rot_model_path)\n        print(f\"  New BEST rotating model saved to: {rot_model_path} (val_f1={current_val_f1:.4f})\")\n    else:\n        patience_counter += 1\n        print(f\"  No improvement. patience_counter={patience_counter}/{PATIENCE}\")\n\n    if patience_counter >= PATIENCE:\n        print(\"Early stopping triggered (no val_f1 improvement). Breaking rotations.\")\n        break\n\nprint(\"Done rotating training. Best val_f1:\", best_rot_val_f1)","metadata":{"execution":{"iopub.status.busy":"2025-12-04T00:22:47.826278Z","iopub.execute_input":"2025-12-04T00:22:47.826607Z","execution_failed":"2025-12-04T00:33:19.453Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"\\nEvaluating final model on realistic (unbalanced) test set (timed with AMP)...\")\ntest_labels, test_probs, avg_batch_s_test = evaluate_loader_timed(model, test_loader, device, use_amp=True, print_progress_every=500)\ntest_metrics = compute_metrics(test_labels, test_probs, threshold=0.5)\nestimated_full_time_h = (len(test_loader) * avg_batch_s_test) / 3600.0 if avg_batch_s_test > 0 else None\nprint(f\"Avg batch seconds during full eval approx: {avg_batch_s_test:.4f}; estimated full-test hours: {estimated_full_time_h:.2f}\")\nprint(\"Realistic test metrics:\", test_metrics)\n\nfinal_ckpt = {\n    \"model_state_dict\": model.state_dict(),\n    \"optimizer_state_dict\": optimizer.state_dict(),\n    \"final_test_metrics\": test_metrics,\n    \"note\": \"trained with rotating safe strategy; threshold tuning recommended\",\n}\ntorch.save(final_ckpt, \"final_rich_checkpointv2baby.pth\")\nprint(\"Saved final rich checkpoint: final_rich_checkpointv2baby.pth\")","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}