{ "nbformat": 4, "nbformat_minor": 0, "metadata": { "colab": { "provenance": [], "gpuType": "T4" }, "kernelspec": { "name": "python3", "display_name": "Python 3" }, "language_info": { "name": "python" }, "accelerator": "GPU" }, "cells": [ { "cell_type": "code", "execution_count": null, "metadata": { "id": "c5OJGv8KJ7tO" }, "outputs": [], "source": [ "!pip install -q transformers peft trl datasets accelerate bitsandbytes huggingface_hub pillow" ] }, { "cell_type": "code", "source": [ "from google.colab import drive\n", "drive.mount('/content/drive')\n", "\n", "from huggingface_hub import login\n", "login()" ], "metadata": { "id": "8Hdo98OWKKhD" }, "execution_count": null, "outputs": [] }, { "cell_type": "code", "source": [ "import pandas as pd\n", "import os\n", "\n", "DATA_DIR = \"/content/drive/MyDrive/synthetic_cards_data\"\n", "CSV_PATH = \"/content/drive/MyDrive/synthetic_cards_progress.csv\"\n", "\n", "df = pd.read_csv(CSV_PATH)\n", "\n", "print(\"Total Records:\", len(df))\n", "display(df.head())" ], "metadata": { "id": "pm8maH_QLOND" }, "execution_count": null, "outputs": [] }, { "cell_type": "code", "source": [ "from datasets import Dataset, Image\n", "\n", "df = df[\n", " [\n", " \"image_path\",\n", " \"doc_type\",\n", " \"side\",\n", " \"state\",\n", " ]\n", "].copy()\n", "\n", "missing = df[\"image_path\"].apply(lambda x: not os.path.exists(x)).sum()\n", "\n", "print(f\"Missing Images : {missing}\")\n", "print(f\"Available Images : {len(df) - missing}\")\n", "\n", "assert missing == 0, \"Some image paths are invalid!\"\n", "\n", "hf_dataset = Dataset.from_pandas(\n", " df,\n", " preserve_index=False,\n", ")\n", "\n", "hf_dataset = hf_dataset.cast_column(\n", " \"image_path\",\n", " Image()\n", ")\n", "\n", "print(hf_dataset)" ], "metadata": { "id": "sxhgrQz5Lrwk" }, "execution_count": null, "outputs": [] }, { "cell_type": "code", "source": [ "from datasets import DatasetDict\n", "\n", "dataset = hf_dataset.train_test_split(\n", " test_size=0.20,\n", " seed=42,\n", " shuffle=True,\n", ")\n", "\n", "train_dataset = dataset[\"train\"]\n", "test_dataset = dataset[\"test\"]\n", "\n", "print(f\"Training Images : {len(train_dataset)}\")\n", "print(f\"Testing Images : {len(test_dataset)}\")" ], "metadata": { "id": "Lk9MHUYWLzWQ" }, "execution_count": null, "outputs": [] }, { "cell_type": "code", "source": [ "doc_type_to_id = {\n", " label: idx\n", " for idx, label in enumerate(\n", " sorted(train_dataset.unique(\"doc_type\"))\n", " )\n", "}\n", "\n", "side_to_id = {\n", " label: idx\n", " for idx, label in enumerate(\n", " sorted(train_dataset.unique(\"side\"))\n", " )\n", "}\n", "\n", "state_to_id = {\n", " label: idx\n", " for idx, label in enumerate(\n", " sorted(train_dataset.unique(\"state\"))\n", " )\n", "}\n", "\n", "\n", "id_to_doc_type = {\n", " idx: label\n", " for label, idx in doc_type_to_id.items()\n", "}\n", "\n", "id_to_side = {\n", " idx: label\n", " for label, idx in side_to_id.items()\n", "}\n", "\n", "id_to_state = {\n", " idx: label\n", " for label, idx in state_to_id.items()\n", "}\n", "\n", "\n", "NUM_DOC_TYPES = len(doc_type_to_id)\n", "NUM_SIDES = len(side_to_id)\n", "NUM_STATES = len(state_to_id)\n", "\n", "\n", "print(\"Document Types :\", NUM_DOC_TYPES)\n", "print(\"Sides :\", NUM_SIDES)\n", "print(\"States :\", NUM_STATES)" ], "metadata": { "id": "-BUT3YwMMCwR" }, "execution_count": null, "outputs": [] }, { "cell_type": "code", "source": [ "from transformers import AutoImageProcessor\n", "\n", "MODEL_NAME = \"google/vit-base-patch16-224\"\n", "\n", "processor = AutoImageProcessor.from_pretrained(MODEL_NAME)\n", "\n", "print(\"Processor Loaded Successfully\")" ], "metadata": { "id": "0QllKW_iMGRi" }, "execution_count": null, "outputs": [] }, { "cell_type": "code", "source": [ "import torch\n", "\n", "def preprocess_batch(batch):\n", " images = batch[\"image_path\"]\n", "\n", " processed = processor(\n", " images=images,\n", " return_tensors=\"pt\"\n", " )\n", "\n", " return {\n", " \"pixel_values\": processed[\"pixel_values\"],\n", " \"doc_type_labels\": [\n", " doc_type_to_id[x] for x in batch[\"doc_type\"]\n", " ],\n", " \"side_labels\": [\n", " side_to_id[x] for x in batch[\"side\"]\n", " ],\n", " \"state_labels\": [\n", " state_to_id[x] for x in batch[\"state\"]\n", " ],\n", " }" ], "metadata": { "id": "C_4AtLtnMLKc" }, "execution_count": null, "outputs": [] }, { "cell_type": "code", "source": [ "from datasets import load_from_disk\n", "import os\n", "\n", "CACHE_PATH = \"/content/drive/MyDrive/vit_processed_cards\"\n", "\n", "if os.path.exists(CACHE_PATH):\n", " print(\"Loading cached dataset...\")\n", " dataset = load_from_disk(CACHE_PATH)\n", "\n", "else:\n", " print(\"Preprocessing images...\")\n", "\n", " dataset = DatasetDict({\n", " \"train\": train_dataset,\n", " \"test\": test_dataset,\n", " })\n", "\n", " dataset = dataset.map(\n", " preprocess_batch,\n", " batched=True,\n", " batch_size=32,\n", " remove_columns=dataset[\"train\"].column_names,\n", " desc=\"Preprocessing Images\",\n", " )\n", "\n", " dataset.save_to_disk(CACHE_PATH)\n", " print(\"Dataset cached successfully!\")" ], "metadata": { "id": "7KtD46OzMWha" }, "execution_count": null, "outputs": [] }, { "cell_type": "code", "source": [ "dataset.set_format(\n", " type=\"torch\",\n", " columns=[\n", " \"pixel_values\",\n", " \"doc_type_labels\",\n", " \"side_labels\",\n", " \"state_labels\",\n", " ]\n", ")\n", "\n", "train_dataset = dataset[\"train\"]\n", "test_dataset = dataset[\"test\"]\n", "\n", "print(train_dataset[0].keys())\n", "print(train_dataset[0][\"pixel_values\"].shape)" ], "metadata": { "id": "L2I4Xlt8WOLo" }, "execution_count": null, "outputs": [] }, { "cell_type": "code", "source": [ "from torch.utils.data import DataLoader\n", "\n", "BATCH_SIZE = 32\n", "\n", "train_loader = DataLoader(\n", " train_dataset,\n", " batch_size=BATCH_SIZE,\n", " shuffle=True,\n", " num_workers=2,\n", " pin_memory=True,\n", " persistent_workers=True,\n", ")\n", "\n", "test_loader = DataLoader(\n", " test_dataset,\n", " batch_size=BATCH_SIZE,\n", " shuffle=False,\n", " num_workers=2,\n", " pin_memory=True,\n", " persistent_workers=True,\n", ")\n", "\n", "print(len(train_loader), len(test_loader))" ], "metadata": { "id": "SdHzd6NoMZFR" }, "execution_count": null, "outputs": [] }, { "cell_type": "code", "source": [ "import torch\n", "import torch.nn as nn\n", "\n", "from transformers import ViTModel\n", "\n", "\n", "class MultiHeadViT(nn.Module):\n", "\n", " def __init__(self):\n", "\n", " super().__init__()\n", "\n", " self.backbone = ViTModel.from_pretrained(\n", " MODEL_NAME\n", " )\n", "\n", " hidden = self.backbone.config.hidden_size\n", "\n", " self.dropout = nn.Dropout(0.1)\n", "\n", " self.doc_classifier = nn.Linear(\n", " hidden,\n", " NUM_DOC_TYPES,\n", " )\n", "\n", " self.side_classifier = nn.Linear(\n", " hidden,\n", " NUM_SIDES,\n", " )\n", "\n", " self.state_classifier = nn.Linear(\n", " hidden,\n", " NUM_STATES,\n", " )\n", "\n", " def forward(self, pixel_values):\n", "\n", " outputs = self.backbone(\n", " pixel_values=pixel_values\n", " )\n", "\n", " cls = outputs.last_hidden_state[:, 0]\n", "\n", " cls = self.dropout(cls)\n", "\n", " return {\n", " \"doc\": self.doc_classifier(cls),\n", " \"side\": self.side_classifier(cls),\n", " \"state\": self.state_classifier(cls),\n", " }" ], "metadata": { "id": "IwbmY-6sazmk" }, "execution_count": null, "outputs": [] }, { "cell_type": "code", "source": [ "device = torch.device(\n", " \"cuda\" if torch.cuda.is_available() else \"cpu\"\n", ")\n", "\n", "model = MultiHeadViT().to(device)\n", "\n", "print(device)" ], "metadata": { "id": "vhAbEPToa21Z" }, "execution_count": null, "outputs": [] }, { "cell_type": "code", "source": [ "from transformers import get_cosine_schedule_with_warmup\n", "\n", "criterion = nn.CrossEntropyLoss()\n", "\n", "optimizer = torch.optim.AdamW(\n", " model.parameters(),\n", " lr=2e-5,\n", " weight_decay=0.01,\n", ")\n", "\n", "EPOCHS = 10\n", "\n", "total_steps = len(train_loader) * EPOCHS\n", "\n", "scheduler = get_cosine_schedule_with_warmup(\n", " optimizer,\n", " num_warmup_steps=int(0.1 * total_steps),\n", " num_training_steps=total_steps,\n", ")" ], "metadata": { "id": "_3QiYpSVa5MT" }, "execution_count": null, "outputs": [] }, { "cell_type": "code", "source": [ "from torch.amp import GradScaler, autocast\n", "\n", "scaler = GradScaler(\"cuda\")" ], "metadata": { "id": "wPV4Ipc2a85E" }, "execution_count": null, "outputs": [] }, { "cell_type": "code", "source": [ "import torch\n", "\n", "def evaluate(model, dataloader, criterion):\n", "\n", " model.eval()\n", "\n", " total_loss = 0\n", "\n", " doc_correct = 0\n", " side_correct = 0\n", " state_correct = 0\n", " combined_correct = 0\n", "\n", " total = 0\n", "\n", " with torch.no_grad():\n", "\n", " for batch in dataloader:\n", "\n", " pixel_values = batch[\"pixel_values\"].to(device)\n", "\n", " doc_labels = batch[\"doc_type_labels\"].to(device)\n", " side_labels = batch[\"side_labels\"].to(device)\n", " state_labels = batch[\"state_labels\"].to(device)\n", "\n", " outputs = model(pixel_values)\n", "\n", " doc_loss = criterion(outputs[\"doc\"], doc_labels)\n", " side_loss = criterion(outputs[\"side\"], side_labels)\n", " state_loss = criterion(outputs[\"state\"], state_labels)\n", "\n", " loss = doc_loss + side_loss + state_loss\n", "\n", " total_loss += loss.item()\n", "\n", " doc_pred = outputs[\"doc\"].argmax(1)\n", " side_pred = outputs[\"side\"].argmax(1)\n", " state_pred = outputs[\"state\"].argmax(1)\n", "\n", " doc_correct += (doc_pred == doc_labels).sum().item()\n", " side_correct += (side_pred == side_labels).sum().item()\n", " state_correct += (state_pred == state_labels).sum().item()\n", "\n", " combined_correct += (\n", " (doc_pred == doc_labels) &\n", " (side_pred == side_labels) &\n", " (state_pred == state_labels)\n", " ).sum().item()\n", "\n", " total += len(doc_labels)\n", "\n", " return {\n", " \"loss\": total_loss / len(dataloader),\n", " \"doc_acc\": 100 * doc_correct / total,\n", " \"side_acc\": 100 * side_correct / total,\n", " \"state_acc\": 100 * state_correct / total,\n", " \"combined_acc\": 100 * combined_correct / total,\n", " }" ], "metadata": { "id": "xsNBD9I7bKdO" }, "execution_count": null, "outputs": [] }, { "cell_type": "code", "source": [ "from tqdm.auto import tqdm\n", "\n", "def train_one_epoch(model, dataloader):\n", "\n", " model.train()\n", "\n", " running_loss = 0\n", "\n", " progress = tqdm(\n", " dataloader,\n", " leave=False,\n", " desc=\"Training\"\n", " )\n", "\n", " for batch in progress:\n", "\n", " pixel_values = batch[\"pixel_values\"].to(device)\n", "\n", " doc_labels = batch[\"doc_type_labels\"].to(device)\n", " side_labels = batch[\"side_labels\"].to(device)\n", " state_labels = batch[\"state_labels\"].to(device)\n", "\n", " optimizer.zero_grad(set_to_none=True)\n", "\n", " with autocast(\"cuda\"):\n", "\n", " outputs = model(pixel_values)\n", "\n", " doc_loss = criterion(outputs[\"doc\"], doc_labels)\n", " side_loss = criterion(outputs[\"side\"], side_labels)\n", " state_loss = criterion(outputs[\"state\"], state_labels)\n", "\n", " loss = doc_loss + side_loss + state_loss\n", "\n", " scaler.scale(loss).backward()\n", "\n", " scaler.unscale_(optimizer)\n", "\n", " torch.nn.utils.clip_grad_norm_(\n", " model.parameters(),\n", " 1.0\n", " )\n", "\n", " scaler.step(optimizer)\n", " scaler.update()\n", "\n", " scheduler.step()\n", "\n", " running_loss += loss.item()\n", "\n", " progress.set_postfix(\n", " loss=f\"{loss.item():.4f}\"\n", " )\n", "\n", " return running_loss / len(dataloader)" ], "metadata": { "id": "HjZs5pjFbp4J" }, "execution_count": null, "outputs": [] }, { "cell_type": "code", "source": [ "import os\n", "\n", "CHECKPOINT = \"/content/drive/MyDrive/multihead_vit_checkpoint.pth\"\n", "BEST_MODEL = \"/content/drive/MyDrive/multihead_vit_best.pth\"\n", "\n", "start_epoch = 0\n", "best_accuracy = 0\n", "\n", "if os.path.exists(CHECKPOINT):\n", "\n", " checkpoint = torch.load(\n", " CHECKPOINT,\n", " map_location=device\n", " )\n", "\n", " model.load_state_dict(\n", " checkpoint[\"model\"]\n", " )\n", "\n", " optimizer.load_state_dict(\n", " checkpoint[\"optimizer\"]\n", " )\n", "\n", " scheduler.load_state_dict(\n", " checkpoint[\"scheduler\"]\n", " )\n", "\n", " scaler.load_state_dict(\n", " checkpoint[\"scaler\"]\n", " )\n", "\n", " start_epoch = checkpoint[\"epoch\"] + 1\n", " best_accuracy = checkpoint[\"best_accuracy\"]\n", "\n", " print(f\"Resuming from epoch {start_epoch}\")" ], "metadata": { "id": "nwMLhEElbryn" }, "execution_count": null, "outputs": [] }, { "cell_type": "code", "source": [ "EARLY_STOPPING = 5\n", "\n", "patience = 0\n", "\n", "for epoch in range(start_epoch, EPOCHS):\n", "\n", " print(\"=\" * 60)\n", " print(f\"Epoch {epoch+1}/{EPOCHS}\")\n", "\n", " train_loss = train_one_epoch(\n", " model,\n", " train_loader,\n", " )\n", "\n", " metrics = evaluate(\n", " model,\n", " test_loader,\n", " criterion,\n", " )\n", "\n", " print(f\"Train Loss : {train_loss:.4f}\")\n", " print(f\"Validation Loss : {metrics['loss']:.4f}\")\n", "\n", " print(f\"Doc Accuracy : {metrics['doc_acc']:.2f}%\")\n", " print(f\"Side Accuracy : {metrics['side_acc']:.2f}%\")\n", " print(f\"State Accuracy : {metrics['state_acc']:.2f}%\")\n", " print(f\"Combined Acc : {metrics['combined_acc']:.2f}%\")\n", "\n", " torch.save(\n", " {\n", " \"epoch\": epoch,\n", " \"model\": model.state_dict(),\n", " \"optimizer\": optimizer.state_dict(),\n", " \"scheduler\": scheduler.state_dict(),\n", " \"scaler\": scaler.state_dict(),\n", " \"best_accuracy\": best_accuracy,\n", " },\n", " CHECKPOINT,\n", " )\n", "\n", " if metrics[\"combined_acc\"] > best_accuracy:\n", "\n", " best_accuracy = metrics[\"combined_acc\"]\n", "\n", " torch.save(\n", " model.state_dict(),\n", " BEST_MODEL,\n", " )\n", "\n", " print(\"✅ Best model updated.\")\n", "\n", " patience = 0\n", "\n", " else:\n", "\n", " patience += 1\n", "\n", " if patience >= EARLY_STOPPING:\n", "\n", " print(\"Early stopping.\")\n", "\n", " break" ], "metadata": { "id": "_H2izh0kbtn9" }, "execution_count": null, "outputs": [] }, { "cell_type": "code", "source": [ "model.load_state_dict(\n", " torch.load(BEST_MODEL, map_location=device)\n", ")\n", "\n", "model.eval()\n", "\n", "print(\"Best model loaded.\")" ], "metadata": { "id": "-QfPTZM9bvUU" }, "execution_count": null, "outputs": [] }, { "cell_type": "code", "source": [ "import json\n", "\n", "label_mappings = {\n", " \"doc_type_to_id\": doc_type_to_id,\n", " \"side_to_id\": side_to_id,\n", " \"state_to_id\": state_to_id,\n", "}\n", "\n", "with open(\"/content/label_mappings.json\", \"w\") as f:\n", " json.dump(label_mappings, f, indent=4)\n", "\n", "print(\"Label mappings saved.\")" ], "metadata": { "id": "03CVtVuZfkYQ" }, "execution_count": null, "outputs": [] }, { "cell_type": "code", "source": [ "from PIL import Image\n", "import torch\n", "\n", "# Reverse mappings\n", "id_to_doc_type = {v: k for k, v in doc_type_to_id.items()}\n", "id_to_side = {v: k for k, v in side_to_id.items()}\n", "id_to_state = {v: k for k, v in state_to_id.items()}\n", "\n", "\n", "def predict(image_path):\n", "\n", " image = Image.open(image_path).convert(\"RGB\")\n", "\n", " inputs = processor(\n", " images=image,\n", " return_tensors=\"pt\"\n", " )\n", "\n", " pixel_values = inputs[\"pixel_values\"].to(device)\n", "\n", " with torch.no_grad():\n", " outputs = model(pixel_values)\n", "\n", " doc = outputs[\"doc\"].argmax(1).item()\n", " side = outputs[\"side\"].argmax(1).item()\n", " state = outputs[\"state\"].argmax(1).item()\n", "\n", " print(\"=\" * 40)\n", " print(\"Prediction\")\n", " print(\"=\" * 40)\n", " print(\"Document Type :\", id_to_doc_type[doc])\n", " print(\"Side :\", id_to_side[side])\n", " print(\"State :\", id_to_state[state])\n", "\n", " return {\n", " \"doc_type\": id_to_doc_type[doc],\n", " \"side\": id_to_side[side],\n", " \"state\": id_to_state[state]\n", " }" ], "metadata": { "id": "Fs7DtXJqflc9" }, "execution_count": null, "outputs": [] }, { "cell_type": "code", "source": [ "import matplotlib.pyplot as plt\n", "\n", "def predict_and_show(image_path):\n", "\n", " image = Image.open(image_path).convert(\"RGB\")\n", "\n", " plt.figure(figsize=(6,6))\n", " plt.imshow(image)\n", " plt.axis(\"off\")\n", "\n", " prediction = predict(image_path)\n", "\n", " plt.title(\n", " f\"{prediction['doc_type']} | \"\n", " f\"{prediction['side']} | \"\n", " f\"{prediction['state']}\"\n", " )\n", "\n", " plt.show()" ], "metadata": { "id": "pAtqRT1dfnms" }, "execution_count": null, "outputs": [] }, { "cell_type": "code", "source": [ "predict_and_show(df.iloc[5][\"image_path\"])" ], "metadata": { "collapsed": true, "id": "pePG_Ummf2DE" }, "execution_count": null, "outputs": [] }, { "cell_type": "code", "source": [], "metadata": { "id": "7Y6gqVhjhWSa" }, "execution_count": null, "outputs": [] } ] }