{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\n# for dirname, _, filenames in os.walk('/kaggle/input'):\n#     for filename in filenames:\n#         print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-03-11T14:40:58.530901Z","iopub.execute_input":"2023-03-11T14:40:58.531330Z","iopub.status.idle":"2023-03-11T14:40:58.548034Z","shell.execute_reply.started":"2023-03-11T14:40:58.531279Z","shell.execute_reply":"2023-03-11T14:40:58.546466Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Origional Ntb https://www.kaggle.com/code/lizhecheng/vision-transformer-vit-resnet-baseline","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport torch\nimport torch.nn as nn\nimport torchvision.transforms as transforms\nimport timm\nimport cv2\nimport torch.optim as optim\nimport torchvision.models as models\nimport torchvision\n\nfrom PIL import Image\nfrom tqdm.notebook import tqdm\nfrom sklearn import model_selection, metrics\nfrom torch.utils.data import DataLoader\nfrom tqdm import tqdm\nfrom transformers import ViTModel\nfrom torch.optim.lr_scheduler import StepLR\nfrom transformers import ViTForImageClassification, ViTConfig\nfrom statistics import mean\n\nplt.style.use('fivethirtyeight')","metadata":{"execution":{"iopub.status.busy":"2023-03-11T14:40:58.550145Z","iopub.execute_input":"2023-03-11T14:40:58.550541Z","iopub.status.idle":"2023-03-11T14:41:03.668163Z","shell.execute_reply.started":"2023-03-11T14:40:58.550499Z","shell.execute_reply":"2023-03-11T14:41:03.667091Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"use_cuda = torch.cuda.is_available()\ndevice = torch.device('cuda' if use_cuda else 'cpu')\nuse_cuda, device","metadata":{"execution":{"iopub.status.busy":"2023-03-11T14:41:03.669676Z","iopub.execute_input":"2023-03-11T14:41:03.671212Z","iopub.status.idle":"2023-03-11T14:41:03.763775Z","shell.execute_reply.started":"2023-03-11T14:41:03.671167Z","shell.execute_reply":"2023-03-11T14:41:03.762395Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BASE_PATH = '/kaggle/input/spr-x-ray-gender/kaggle/kaggle'","metadata":{"execution":{"iopub.status.busy":"2023-03-11T14:41:03.768115Z","iopub.execute_input":"2023-03-11T14:41:03.768435Z","iopub.status.idle":"2023-03-11T14:41:03.774825Z","shell.execute_reply.started":"2023-03-11T14:41:03.768399Z","shell.execute_reply":"2023-03-11T14:41:03.773568Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image = cv2.imread('/kaggle/input/spr-x-ray-gender/kaggle/kaggle/train/000000.png')\n# image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\nplt.imshow(image)\nplt.show()\nimage.shape","metadata":{"execution":{"iopub.status.busy":"2023-03-11T14:41:03.776761Z","iopub.execute_input":"2023-03-11T14:41:03.777582Z","iopub.status.idle":"2023-03-11T14:41:04.406833Z","shell.execute_reply.started":"2023-03-11T14:41:03.777537Z","shell.execute_reply":"2023-03-11T14:41:04.405784Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"IMG_SIZE = 224\nBATCH_SIZE = 16\nLR = 1e-4\nEPOCHS = 10\nnum_classes = 2\n\n# mean = [0.485, 0.456, 0.406]\n# std = [0.229, 0.224, 0.225]","metadata":{"execution":{"iopub.status.busy":"2023-03-11T14:41:04.407816Z","iopub.execute_input":"2023-03-11T14:41:04.408987Z","iopub.status.idle":"2023-03-11T14:41:04.414768Z","shell.execute_reply.started":"2023-03-11T14:41:04.408950Z","shell.execute_reply":"2023-03-11T14:41:04.413445Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = pd.read_csv('/kaggle/input/spr-x-ray-gender/train_gender.csv')\ntrain.head()","metadata":{"execution":{"iopub.status.busy":"2023-03-11T14:41:04.416615Z","iopub.execute_input":"2023-03-11T14:41:04.417440Z","iopub.status.idle":"2023-03-11T14:41:04.437156Z","shell.execute_reply.started":"2023-03-11T14:41:04.417401Z","shell.execute_reply":"2023-03-11T14:41:04.436324Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train.shape","metadata":{"execution":{"iopub.status.busy":"2023-03-11T14:41:04.438618Z","iopub.execute_input":"2023-03-11T14:41:04.439322Z","iopub.status.idle":"2023-03-11T14:41:04.445779Z","shell.execute_reply.started":"2023-03-11T14:41:04.439282Z","shell.execute_reply":"2023-03-11T14:41:04.444656Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(4, 3))\n\ntrain['gender'].value_counts().plot(\n    kind='bar',\n    color='#558364',\n    width=0.5\n)\n\nplt.xlabel('Gender', fontsize=12)\nplt.ylabel('Count', fontsize=12)\nplt.title('Distribution of Gender', fontsize=15)\nplt.xticks(rotation=360)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-03-11T14:41:04.447409Z","iopub.execute_input":"2023-03-11T14:41:04.448122Z","iopub.status.idle":"2023-03-11T14:41:04.623037Z","shell.execute_reply.started":"2023-03-11T14:41:04.448084Z","shell.execute_reply":"2023-03-11T14:41:04.622078Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df, val_df = model_selection.train_test_split(\n    train, test_size=0.1, random_state=42, stratify=train['gender'].values\n)","metadata":{"execution":{"iopub.status.busy":"2023-03-11T14:41:04.624425Z","iopub.execute_input":"2023-03-11T14:41:04.626192Z","iopub.status.idle":"2023-03-11T14:41:04.638446Z","shell.execute_reply.started":"2023-03-11T14:41:04.626161Z","shell.execute_reply":"2023-03-11T14:41:04.637207Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class LeafDataset(torch.utils.data.Dataset):\n    \n    def __init__(self, df, data_path=BASE_PATH, mode='train', transforms=None):\n        super().__init__()\n        self.df_data = df.values\n        self.data_path = data_path\n        self.transforms = transforms \n        self.mode = mode\n        self.data_dir = 'train' if mode == 'train' else 'test'\n    \n    def __len__(self):\n        return len(self.df_data)\n    \n    def __getitem__(self, index):\n        img_name, label = self.df_data[index]\n        img_name = int(img_name)  \n        \n        if img_name <= 9:\n            img_name = '00000' + str(img_name)\n        elif img_name <= 99:\n            img_name = '0000' + str(img_name)\n        elif img_name <= 999:\n            img_name = '000' + str(img_name)\n        elif img_name <= 9999:\n            img_name = '00' + str(img_name)\n        else:\n            img_name = '0' + str(img_name)\n            \n        relative_path = img_name + '.png'\n        img_path = os.path.join(self.data_path, self.data_dir, relative_path)\n        img = Image.open(img_path).convert('RGB')\n        \n        if self.transforms is not None:\n            img = self.transforms(img)\n            \n        return img, label","metadata":{"execution":{"iopub.status.busy":"2023-03-11T14:41:04.640043Z","iopub.execute_input":"2023-03-11T14:41:04.640456Z","iopub.status.idle":"2023-03-11T14:41:04.650517Z","shell.execute_reply.started":"2023-03-11T14:41:04.640417Z","shell.execute_reply":"2023-03-11T14:41:04.649460Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transforms_train = transforms.Compose([\n    transforms.Resize((IMG_SIZE, IMG_SIZE)),\n    transforms.RandomHorizontalFlip(p=0.2),\n#     transforms.RandomVerticalFlip(p=0.1),\n#     transforms.RandomResizedCrop(IMG_SIZE),\n    transforms.ToTensor(),\n#     transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225))\n])\n\ntransforms_val = transforms.Compose([\n    transforms.Resize((IMG_SIZE, IMG_SIZE)),\n    transforms.ToTensor(),\n#     transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225))\n])","metadata":{"execution":{"iopub.status.busy":"2023-03-11T14:41:04.651889Z","iopub.execute_input":"2023-03-11T14:41:04.652527Z","iopub.status.idle":"2023-03-11T14:41:04.666362Z","shell.execute_reply.started":"2023-03-11T14:41:04.652490Z","shell.execute_reply":"2023-03-11T14:41:04.665272Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = LeafDataset(df=train_df, data_path=BASE_PATH, mode='train', transforms=transforms_train)\nval_dataset = LeafDataset(df=val_df, data_path=BASE_PATH, mode='train', transforms=transforms_val)","metadata":{"execution":{"iopub.status.busy":"2023-03-11T14:41:04.671149Z","iopub.execute_input":"2023-03-11T14:41:04.671441Z","iopub.status.idle":"2023-03-11T14:41:04.678620Z","shell.execute_reply.started":"2023-03-11T14:41:04.671415Z","shell.execute_reply":"2023-03-11T14:41:04.676756Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(train_dataset), len(val_dataset)","metadata":{"execution":{"iopub.status.busy":"2023-03-11T14:41:04.679981Z","iopub.execute_input":"2023-03-11T14:41:04.680834Z","iopub.status.idle":"2023-03-11T14:41:04.689550Z","shell.execute_reply.started":"2023-03-11T14:41:04.680795Z","shell.execute_reply":"2023-03-11T14:41:04.688464Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = ViTForImageClassification.from_pretrained(\"google/vit-base-patch16-224\")\nmodel.classifier = nn.Linear(in_features=768, out_features=num_classes, bias=True)\nmodel.to(device)","metadata":{"execution":{"iopub.status.busy":"2023-03-11T14:41:04.692516Z","iopub.execute_input":"2023-03-11T14:41:04.692800Z","iopub.status.idle":"2023-03-11T14:41:09.813409Z","shell.execute_reply.started":"2023-03-11T14:41:04.692773Z","shell.execute_reply":"2023-03-11T14:41:09.812282Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_losses = []\ntrain_accs = []\nval_losses = []\nval_accs = []\n\ndef train_model(model, train_dataset, val_dataset, learning_rate, epochs):\n\n    train_dataloader = torch.utils.data.DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True)\n    val_dataloader = torch.utils.data.DataLoader(val_dataset, batch_size=BATCH_SIZE)\n    \n    criterion = nn.CrossEntropyLoss()\n    optimizer = optim.Adam(model.parameters(), lr=learning_rate)\n    scheduler = StepLR(optimizer, step_size=1, gamma=0.9)\n    \n    if use_cuda:\n        model = model.cuda()\n        criterion = criterion.cuda()\n\n    for epoch_num in range(epochs):\n\n            total_acc_train = 0\n            total_loss_train = 0\n\n            for train_images, train_labels in tqdm(train_dataloader):\n                \n                train_images = train_images.to(device)\n                train_labels = train_labels.to(device)\n                \n                optimizer.zero_grad()\n\n                output = model(train_images)\n                \n                batch_loss = criterion(output.logits, train_labels.long())\n                total_loss_train += batch_loss.item()\n            \n                _, predicted = torch.max(output.logits.data, 1)\n                acc = (predicted == train_labels).sum().item()\n                total_acc_train += acc\n\n                batch_loss.backward()\n                optimizer.step()\n                \n            scheduler.step()\n            \n            total_acc_val = 0\n            total_loss_val = 0\n\n            with torch.no_grad():\n\n                for val_images, val_labels in val_dataloader:\n                    \n                    val_images = val_images.to(device)\n                    val_labels = val_labels.to(device)\n                    \n                    output = model(val_images)\n\n                    batch_loss = criterion(output.logits, val_labels.long())\n                    total_loss_val += batch_loss.item()\n                    \n                    _, predicted = torch.max(output.logits.data, 1)\n                    acc = (predicted == val_labels).sum().item()\n                    total_acc_val += acc\n            \n            print(f'Epochs: {epoch_num + 1} | Train Loss: {total_loss_train / len(train_dataset): .3f} \\\n            | Train Accuracy: {total_acc_train / len(train_dataset): .3f} \\\n            | Val Loss: {total_loss_val / len(val_dataset): .3f} \\\n            | Val Accuracy: {total_acc_val / len(val_dataset): .3f}')\n            \n            train_losses.append(total_loss_train / len(train_dataset))\n            train_accs.append(total_acc_train / len(train_dataset))\n            val_losses.append(total_loss_val / len(val_dataset))\n            val_accs.append(total_acc_val / len(val_dataset))","metadata":{"execution":{"iopub.status.busy":"2023-03-11T14:41:09.815470Z","iopub.execute_input":"2023-03-11T14:41:09.815867Z","iopub.status.idle":"2023-03-11T14:41:09.829107Z","shell.execute_reply.started":"2023-03-11T14:41:09.815827Z","shell.execute_reply":"2023-03-11T14:41:09.827824Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_model(model, train_dataset, val_dataset, LR, 3)","metadata":{"execution":{"iopub.status.busy":"2023-03-11T14:41:09.830815Z","iopub.execute_input":"2023-03-11T14:41:09.831552Z","iopub.status.idle":"2023-03-11T14:41:58.045985Z","shell.execute_reply.started":"2023-03-11T14:41:09.831513Z","shell.execute_reply":"2023-03-11T14:41:58.044222Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(8, 5))\n\nplt.plot(\n    train_losses, \n    label='Train_Losses', \n    color='red', \n    linewidth=1.5\n)\nplt.plot(\n    val_losses, \n    label='Val_Losses', \n    color='blue', \n    linewidth=1.5\n)\n\nplt.plot(\n    train_accs, \n    label='Train_Accuracy', \n    color='green', \n    linewidth=1.5\n)\nplt.plot(\n    val_accs, \n    label='Val_Accuracy', \n    color='pink', \n    linewidth=1.5\n)\n\nplt.xlabel('Epoch')\nplt.ylabel('Loss / Accuracy')\nplt.title('Loss / Accuracy on train / validation')\nplt.legend()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-03-11T14:41:58.047275Z","iopub.status.idle":"2023-03-11T14:41:58.048654Z","shell.execute_reply.started":"2023-03-11T14:41:58.048380Z","shell.execute_reply":"2023-03-11T14:41:58.048409Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df = pd.read_csv('/kaggle/input/spr-x-ray-gender/sample_submission_gender.csv')\nsub_df.head()","metadata":{"execution":{"iopub.status.busy":"2023-03-11T14:41:58.050181Z","iopub.status.idle":"2023-03-11T14:41:58.051097Z","shell.execute_reply.started":"2023-03-11T14:41:58.050821Z","shell.execute_reply":"2023-03-11T14:41:58.050849Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df.shape","metadata":{"execution":{"iopub.status.busy":"2023-03-11T14:41:58.052695Z","iopub.status.idle":"2023-03-11T14:41:58.053336Z","shell.execute_reply.started":"2023-03-11T14:41:58.053049Z","shell.execute_reply":"2023-03-11T14:41:58.053076Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preds = []\n\ntest_dataset = LeafDataset(df=sub_df, data_path=BASE_PATH, mode='test', transforms=transforms_val)\n\ndef predict(model, test_dataset):\n    \n    test_dataloader = torch.utils.data.DataLoader(test_dataset, batch_size=BATCH_SIZE)\n    \n    for test_images, test_labels in tqdm(test_dataloader):\n        test_images = test_images.to(device)\n        test_labels = test_labels.to(device)\n\n        output = model(test_images)\n\n        _, predicted = torch.max(output.logits.data, 1)\n        preds.extend(predicted.cpu().data.numpy())\n        \n    print(len(preds))\n        \npredict(model, test_dataset)","metadata":{"execution":{"iopub.status.busy":"2023-03-11T14:41:58.055356Z","iopub.status.idle":"2023-03-11T14:41:58.055860Z","shell.execute_reply.started":"2023-03-11T14:41:58.055598Z","shell.execute_reply":"2023-03-11T14:41:58.055625Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df['gender'] = preds\nsub_df.to_csv('submission.csv', index=False)\nsub_df.head()","metadata":{"execution":{"iopub.status.busy":"2023-03-11T14:41:58.057727Z","iopub.status.idle":"2023-03-11T14:41:58.058299Z","shell.execute_reply.started":"2023-03-11T14:41:58.058000Z","shell.execute_reply":"2023-03-11T14:41:58.058027Z"},"trusted":true},"execution_count":null,"outputs":[]}]}