{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":4829,"databundleVersionId":44847,"sourceType":"competition"}],"dockerImageVersionId":30587,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"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\nfor 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-11-26T15:25:26.676083Z","iopub.execute_input":"2023-11-26T15:25:26.677112Z","iopub.status.idle":"2023-11-26T15:25:27.126993Z","shell.execute_reply.started":"2023-11-26T15:25:26.677060Z","shell.execute_reply":"2023-11-26T15:25:27.125864Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%reload_ext autoreload\n%autoreload 2\n%matplotlib inline\n\nfrom fastai import *\nfrom fastai.vision import *\nfrom fastai.metrics import accuracy\nfrom fastai.vision.all import *\n\nfrom fastai.metrics import error_rate\nfrom IPython.display import Image\nfrom pathlib import Path","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:25:27.128939Z","iopub.execute_input":"2023-11-26T15:25:27.129384Z","iopub.status.idle":"2023-11-26T15:25:34.652276Z","shell.execute_reply.started":"2023-11-26T15:25:27.129324Z","shell.execute_reply":"2023-11-26T15:25:34.651073Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import tarfile\n# Specify the path to the tarball\ntarball1_path = '/kaggle/input/yelp-restaurant-photo-classification/train_photo_to_biz_ids.csv.tgz'\n\n# Extract the tarball\nwith tarfile.open(tarball1_path, 'r:gz') as tar:\n    # Extract all files to the current working directory\n    tar.extractall()\n\n# Now read the CSV file\ntrain_photos = pd.read_csv('train_photo_to_biz_ids.csv')\n#train_photos = pd.read_csv('/kaggle/input/yelp-restaurant-photo-classification/train_photo_to_biz_ids.csv.tgz')\n\n# Specify the path to the tarball\ntarball2_path = '/kaggle/input/yelp-restaurant-photo-classification/train.csv.tgz'\nwith tarfile.open(tarball2_path, 'r:gz') as tar:\n    # Extract all files to the current working directory\n    tar.extractall()\ntrain_attr = pd.read_csv('train.csv')\n\n\n# Specify the path to the tarball\ntarball3_path = '/kaggle/input/yelp-restaurant-photo-classification/train.csv.tgz'\nwith tarfile.open(tarball3_path, 'r:gz') as tar:\n    tar.extractall()\ntrain_id = pd.read_csv('train_photo_to_biz_ids.csv')\n\n\n# Specify the path to the tarball\ntarball4_path = '/kaggle/input/yelp-restaurant-photo-classification/test_photo_to_biz.csv.tgz'\nwith tarfile.open(tarball4_path, 'r:gz') as tar:\n    tar.extractall()\ntest_photos = pd.read_csv('test_photo_to_biz.csv')","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:25:34.653700Z","iopub.execute_input":"2023-11-26T15:25:34.654022Z","iopub.status.idle":"2023-11-26T15:25:35.517906Z","shell.execute_reply.started":"2023-11-26T15:25:34.653994Z","shell.execute_reply":"2023-11-26T15:25:35.516739Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Extract Files**","metadata":{}},{"cell_type":"code","source":"import time\n\nstart_time = time.time()\n\n!tar -xzf /kaggle/input/yelp-restaurant-photo-classification/sample_submission.csv.tgz\n!tar -xzf /kaggle/input/yelp-restaurant-photo-classification/test_photo_to_biz.csv.tgz\n!tar -xzf /kaggle/input/yelp-restaurant-photo-classification/test_photos.tgz\n!tar -xzf /kaggle/input/yelp-restaurant-photo-classification/train.csv.tgz\n!tar -xzf /kaggle/input/yelp-restaurant-photo-classification/train_photo_to_biz_ids.csv.tgz\n!tar -xzf /kaggle/input/yelp-restaurant-photo-classification/train_photos.tgz\n\nend_time = time.time()\nprint(\"Run time: {:.2f} seconds\".format(end_time - start_time))\nprint(\"Run time: {:.2f} minutes\".format((end_time - start_time)/60))","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:25:35.520586Z","iopub.execute_input":"2023-11-26T15:25:35.520885Z","iopub.status.idle":"2023-11-26T15:29:59.887938Z","shell.execute_reply.started":"2023-11-26T15:25:35.520859Z","shell.execute_reply":"2023-11-26T15:29:59.886618Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Read CSV File","metadata":{}},{"cell_type":"code","source":"train_photos = pd.read_csv(\"/kaggle/working/train_photo_to_biz_ids.csv\")\ntest_photos = pd.read_csv(\"/kaggle/working/test_photo_to_biz.csv\")\ntrain_attr = pd.read_csv(\"/kaggle/working/train.csv\")\ntrain_id = pd.read_csv('train_photo_to_biz_ids.csv')\nsub = pd.read_csv('sample_submission.csv')","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:29:59.889591Z","iopub.execute_input":"2023-11-26T15:29:59.889961Z","iopub.status.idle":"2023-11-26T15:30:00.559451Z","shell.execute_reply.started":"2023-11-26T15:29:59.889931Z","shell.execute_reply":"2023-11-26T15:30:00.558292Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Detail of Data & Prep DataFrame","metadata":{}},{"cell_type":"code","source":"test_photos","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:00.561478Z","iopub.execute_input":"2023-11-26T15:30:00.561941Z","iopub.status.idle":"2023-11-26T15:30:00.647083Z","shell.execute_reply.started":"2023-11-26T15:30:00.561901Z","shell.execute_reply":"2023-11-26T15:30:00.646081Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:00.648409Z","iopub.execute_input":"2023-11-26T15:30:00.648781Z","iopub.status.idle":"2023-11-26T15:30:00.709529Z","shell.execute_reply.started":"2023-11-26T15:30:00.648753Z","shell.execute_reply":"2023-11-26T15:30:00.708425Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_csv = pd.merge(test_photos, sub, how = \"inner\")\ntest_csv","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:00.710690Z","iopub.execute_input":"2023-11-26T15:30:00.710991Z","iopub.status.idle":"2023-11-26T15:30:01.058062Z","shell.execute_reply.started":"2023-11-26T15:30:00.710965Z","shell.execute_reply":"2023-11-26T15:30:01.056984Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_csv = test_csv.groupby(['business_id'], as_index=False).first()\ntest_csv = test_csv.groupby(['photo_id'], as_index=False).first()\ntest_csv","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:01.059430Z","iopub.execute_input":"2023-11-26T15:30:01.060091Z","iopub.status.idle":"2023-11-26T15:30:01.371092Z","shell.execute_reply.started":"2023-11-26T15:30:01.060054Z","shell.execute_reply":"2023-11-26T15:30:01.370101Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_photos","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:01.377558Z","iopub.execute_input":"2023-11-26T15:30:01.377848Z","iopub.status.idle":"2023-11-26T15:30:01.440080Z","shell.execute_reply.started":"2023-11-26T15:30:01.377823Z","shell.execute_reply":"2023-11-26T15:30:01.439197Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_attr = train_attr.dropna(axis=0)","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:01.441532Z","iopub.execute_input":"2023-11-26T15:30:01.441858Z","iopub.status.idle":"2023-11-26T15:30:01.498989Z","shell.execute_reply.started":"2023-11-26T15:30:01.441815Z","shell.execute_reply":"2023-11-26T15:30:01.498067Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_attr","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:01.501148Z","iopub.execute_input":"2023-11-26T15:30:01.501551Z","iopub.status.idle":"2023-11-26T15:30:01.562662Z","shell.execute_reply.started":"2023-11-26T15:30:01.501518Z","shell.execute_reply":"2023-11-26T15:30:01.561624Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_label = pd.merge(train_photos,train_attr,how = \"inner\")","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:01.564044Z","iopub.execute_input":"2023-11-26T15:30:01.564920Z","iopub.status.idle":"2023-11-26T15:30:01.638020Z","shell.execute_reply.started":"2023-11-26T15:30:01.564885Z","shell.execute_reply":"2023-11-26T15:30:01.637174Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_label.isnull().sum()","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:01.639105Z","iopub.execute_input":"2023-11-26T15:30:01.639430Z","iopub.status.idle":"2023-11-26T15:30:01.714385Z","shell.execute_reply.started":"2023-11-26T15:30:01.639404Z","shell.execute_reply":"2023-11-26T15:30:01.713532Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_label = train_label.sample(n=30000)\ntrain_label","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:01.715484Z","iopub.execute_input":"2023-11-26T15:30:01.715796Z","iopub.status.idle":"2023-11-26T15:30:01.785211Z","shell.execute_reply.started":"2023-11-26T15:30:01.715772Z","shell.execute_reply":"2023-11-26T15:30:01.784400Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# train_label = train_label.groupby(['business_id'], as_index=False).first()\n# train_label = train_label.groupby(['photo_id'], as_index=False).first()\n# train_label","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:01.786607Z","iopub.execute_input":"2023-11-26T15:30:01.787332Z","iopub.status.idle":"2023-11-26T15:30:01.841186Z","shell.execute_reply.started":"2023-11-26T15:30:01.787279Z","shell.execute_reply":"2023-11-26T15:30:01.840356Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_label.info()","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:01.842279Z","iopub.execute_input":"2023-11-26T15:30:01.842641Z","iopub.status.idle":"2023-11-26T15:30:01.912404Z","shell.execute_reply.started":"2023-11-26T15:30:01.842615Z","shell.execute_reply":"2023-11-26T15:30:01.911464Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import os\n\n# def check_file_exists(filename):\n#     if os.path.isfile(filename):\n#         print(f\"The file '{filename}' exists in the directory.\")\n#     else:\n#         print(f\"The file '{filename}' does not exist in the directory.\")\n\n# # Example usage\n# filename_to_check = input(\"Input\")\n# check_file_exists(filename_to_check)","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:01.913481Z","iopub.execute_input":"2023-11-26T15:30:01.913744Z","iopub.status.idle":"2023-11-26T15:30:01.965460Z","shell.execute_reply.started":"2023-11-26T15:30:01.913720Z","shell.execute_reply":"2023-11-26T15:30:01.964527Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Export DataFrame to CSV**","metadata":{}},{"cell_type":"code","source":"test_csv.to_csv('test_csv.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:01.966774Z","iopub.execute_input":"2023-11-26T15:30:01.967243Z","iopub.status.idle":"2023-11-26T15:30:02.048048Z","shell.execute_reply.started":"2023-11-26T15:30:01.967207Z","shell.execute_reply":"2023-11-26T15:30:02.047385Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_label.to_csv('train_label.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:02.049058Z","iopub.execute_input":"2023-11-26T15:30:02.049322Z","iopub.status.idle":"2023-11-26T15:30:02.181805Z","shell.execute_reply.started":"2023-11-26T15:30:02.049300Z","shell.execute_reply":"2023-11-26T15:30:02.180917Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"path_images = '/kaggle/working/train_photos'\nfilenames = get_image_files(path_images)\nfilenames[:15]","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:02.183218Z","iopub.execute_input":"2023-11-26T15:30:02.183500Z","iopub.status.idle":"2023-11-26T15:30:05.373044Z","shell.execute_reply.started":"2023-11-26T15:30:02.183476Z","shell.execute_reply":"2023-11-26T15:30:05.371969Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batchsize = 16\ntfms = aug_transforms(flip_vert=True, max_lighting=0.1, \n                      max_zoom=1.05, max_warp=0.)","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:05.374453Z","iopub.execute_input":"2023-11-26T15:30:05.374777Z","iopub.status.idle":"2023-11-26T15:30:05.438560Z","shell.execute_reply.started":"2023-11-26T15:30:05.374749Z","shell.execute_reply":"2023-11-26T15:30:05.437318Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# conda install -c fastai fastai","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:05.440117Z","iopub.execute_input":"2023-11-26T15:30:05.440612Z","iopub.status.idle":"2023-11-26T15:30:05.502840Z","shell.execute_reply.started":"2023-11-26T15:30:05.440563Z","shell.execute_reply":"2023-11-26T15:30:05.501805Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# np.random.seed(42)\n# datasource = (ImageList.from_csv(path, 'train_v2.csv', \n#                                  folder='train-jpg', suffix='.jpg')\n#        .split_by_rand_pct(0.2)\n#               \\\n#        .label_from_df(label_delim=' '))","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:05.504541Z","iopub.execute_input":"2023-11-26T15:30:05.505484Z","iopub.status.idle":"2023-11-26T15:30:05.568003Z","shell.execute_reply.started":"2023-11-26T15:30:05.505443Z","shell.execute_reply":"2023-11-26T15:30:05.566995Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport torch\nimport pandas as pd\nimport numpy as np\nfrom torch.utils.data import Dataset, random_split, DataLoader\nfrom PIL import Image\nimport torchvision.models as models\nimport matplotlib.pyplot as plt\nimport torchvision.transforms as transforms\nfrom sklearn.metrics import f1_score\nimport torch.nn.functional as F\nimport torch.nn as nn\nfrom torchvision.utils import make_grid\n%matplotlib inline","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:05.569710Z","iopub.execute_input":"2023-11-26T15:30:05.570042Z","iopub.status.idle":"2023-11-26T15:30:05.634376Z","shell.execute_reply.started":"2023-11-26T15:30:05.570014Z","shell.execute_reply":"2023-11-26T15:30:05.633188Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DATA_DIR = '/kaggle/working/'\nTRAIN_DIR = DATA_DIR + 'train_photos' \nTEST_DIR = DATA_DIR + 'test_photos'\nTRAIN_CSV = DATA_DIR + 'train_label.csv'\nTEST_CSV = DATA_DIR + 'test_csv.csv'\nTEST_TO_BUS = DATA_DIR + 'test_photo_to_biz.csv'","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:05.635722Z","iopub.execute_input":"2023-11-26T15:30:05.636066Z","iopub.status.idle":"2023-11-26T15:30:05.699276Z","shell.execute_reply.started":"2023-11-26T15:30:05.636038Z","shell.execute_reply":"2023-11-26T15:30:05.698000Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!head \"{TEST_TO_BUS}\"","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:05.701279Z","iopub.execute_input":"2023-11-26T15:30:05.701784Z","iopub.status.idle":"2023-11-26T15:30:06.869070Z","shell.execute_reply.started":"2023-11-26T15:30:05.701739Z","shell.execute_reply":"2023-11-26T15:30:06.867920Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!head \"{TRAIN_CSV}\"","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:06.879520Z","iopub.execute_input":"2023-11-26T15:30:06.880344Z","iopub.status.idle":"2023-11-26T15:30:07.951494Z","shell.execute_reply.started":"2023-11-26T15:30:06.880314Z","shell.execute_reply":"2023-11-26T15:30:07.950364Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!head \"{TEST_CSV}\"","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:07.953267Z","iopub.execute_input":"2023-11-26T15:30:07.954199Z","iopub.status.idle":"2023-11-26T15:30:09.123991Z","shell.execute_reply.started":"2023-11-26T15:30:07.954159Z","shell.execute_reply":"2023-11-26T15:30:09.122974Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!ls \"{TEST_DIR}\" | head","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:09.125480Z","iopub.execute_input":"2023-11-26T15:30:09.125810Z","iopub.status.idle":"2023-11-26T15:30:10.677393Z","shell.execute_reply.started":"2023-11-26T15:30:09.125782Z","shell.execute_reply":"2023-11-26T15:30:10.676057Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!ls \"{TRAIN_DIR}\" | head","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:10.679089Z","iopub.execute_input":"2023-11-26T15:30:10.679462Z","iopub.status.idle":"2023-11-26T15:30:12.305075Z","shell.execute_reply.started":"2023-11-26T15:30:10.679430Z","shell.execute_reply":"2023-11-26T15:30:12.303767Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = pd.read_csv(TRAIN_CSV)\ntrain_df","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:12.306625Z","iopub.execute_input":"2023-11-26T15:30:12.306974Z","iopub.status.idle":"2023-11-26T15:30:12.394761Z","shell.execute_reply.started":"2023-11-26T15:30:12.306945Z","shell.execute_reply":"2023-11-26T15:30:12.393788Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Prep Data","metadata":{}},{"cell_type":"code","source":"from fastai import *\nfrom fastai.vision import *\nfrom fastai.vision.all import *\nimport torch\npath_images = '/kaggle/working/train_photos'\nfilenames = get_image_files(path_images)\nfilenames[:10]","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:12.395936Z","iopub.execute_input":"2023-11-26T15:30:12.396251Z","iopub.status.idle":"2023-11-26T15:30:15.384026Z","shell.execute_reply.started":"2023-11-26T15:30:12.396224Z","shell.execute_reply":"2023-11-26T15:30:15.382997Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labels = {0: 'good_for_lunch', \n         1: 'good_for_dinner', \n         2: 'takes_reservations', \n         3: 'outdoor_seating',\n         4: 'restaurant_is_expensive', \n         5: 'has_alcohol', \n         6: 'has_table_service', \n         7: 'ambience_is_classy',\n         8: 'good_for_kids'}","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:15.385146Z","iopub.execute_input":"2023-11-26T15:30:15.385461Z","iopub.status.idle":"2023-11-26T15:30:15.441492Z","shell.execute_reply.started":"2023-11-26T15:30:15.385436Z","shell.execute_reply":"2023-11-26T15:30:15.440741Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def encode_label(label):\n    target = torch.zeros(9)\n    for l in str(label).split(' '):\n        target[int(l)] = 1.\n    return target","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:15.442788Z","iopub.execute_input":"2023-11-26T15:30:15.443524Z","iopub.status.idle":"2023-11-26T15:30:15.495805Z","shell.execute_reply.started":"2023-11-26T15:30:15.443491Z","shell.execute_reply":"2023-11-26T15:30:15.494969Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def decode_target(target, text_labels=False, threshold=0.5):\n    result = []\n    for i, x in enumerate(target):\n        if (x >= threshold):\n            if text_labels:\n                result.append(labels[i] + \"(\" + str(i) + \")\")\n            else:\n                result.append(str(i))\n    return ' '.join(result)","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:15.496945Z","iopub.execute_input":"2023-11-26T15:30:15.497718Z","iopub.status.idle":"2023-11-26T15:30:15.552612Z","shell.execute_reply.started":"2023-11-26T15:30:15.497679Z","shell.execute_reply":"2023-11-26T15:30:15.551685Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"encode_label('0 3 8')","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:15.553621Z","iopub.execute_input":"2023-11-26T15:30:15.553866Z","iopub.status.idle":"2023-11-26T15:30:15.716351Z","shell.execute_reply.started":"2023-11-26T15:30:15.553845Z","shell.execute_reply":"2023-11-26T15:30:15.715378Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"decode_target(torch.tensor([0., 1., 0., 1., 0., 0., 0., 0., 1.,]))","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:15.717588Z","iopub.execute_input":"2023-11-26T15:30:15.717901Z","iopub.status.idle":"2023-11-26T15:30:15.773102Z","shell.execute_reply.started":"2023-11-26T15:30:15.717875Z","shell.execute_reply":"2023-11-26T15:30:15.772144Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"decode_target(torch.tensor([0., 1., 0., 1., 0., 0., 0., 0., 1.]), text_labels=True)","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:15.774513Z","iopub.execute_input":"2023-11-26T15:30:15.775097Z","iopub.status.idle":"2023-11-26T15:30:15.831785Z","shell.execute_reply.started":"2023-11-26T15:30:15.775064Z","shell.execute_reply":"2023-11-26T15:30:15.830702Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class YelpResPhoto(Dataset):\n    def __init__(self, csv_file, root_dir, transform=None):\n        self.df = pd.read_csv(csv_file)\n        self.transform = transform\n        self.root_dir = root_dir\n        \n    def __len__(self):\n        return len(self.df)    \n    \n    def __getitem__(self, idx):\n        row = self.df.loc[idx]\n        img_id, img_label = row['photo_id'], row['labels']\n        #img_id, img_label = row['business_id'], row['labels']\n        img_fname = self.root_dir + \"/\" + str(img_id) + \".jpg\"\n        img = Image.open(img_fname)\n        \n        if self.transform:\n            img = self.transform(img)\n        return img, encode_label(img_label)\n        #return img, img_label","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:15.832880Z","iopub.execute_input":"2023-11-26T15:30:15.833169Z","iopub.status.idle":"2023-11-26T15:30:15.892491Z","shell.execute_reply.started":"2023-11-26T15:30:15.833143Z","shell.execute_reply":"2023-11-26T15:30:15.891500Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transform = transforms.Compose([transforms.ToTensor()])\ndataset = YelpResPhoto(TRAIN_CSV, TRAIN_DIR, transform=transform)","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:15.893583Z","iopub.execute_input":"2023-11-26T15:30:15.894104Z","iopub.status.idle":"2023-11-26T15:30:15.966025Z","shell.execute_reply.started":"2023-11-26T15:30:15.894076Z","shell.execute_reply":"2023-11-26T15:30:15.965267Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TRAIN_DIR","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:15.967222Z","iopub.execute_input":"2023-11-26T15:30:15.967824Z","iopub.status.idle":"2023-11-26T15:30:16.025888Z","shell.execute_reply.started":"2023-11-26T15:30:15.967787Z","shell.execute_reply":"2023-11-26T15:30:16.024847Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:16.027501Z","iopub.execute_input":"2023-11-26T15:30:16.028196Z","iopub.status.idle":"2023-11-26T15:30:16.090129Z","shell.execute_reply.started":"2023-11-26T15:30:16.028156Z","shell.execute_reply":"2023-11-26T15:30:16.089058Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(dataset)","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:16.091620Z","iopub.execute_input":"2023-11-26T15:30:16.092181Z","iopub.status.idle":"2023-11-26T15:30:16.156893Z","shell.execute_reply.started":"2023-11-26T15:30:16.092150Z","shell.execute_reply":"2023-11-26T15:30:16.155723Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def show_sample(img, target, invert=False):\n    if invert:\n        plt.imshow(1 - img.permute((1, 2, 0)))\n    else:\n        plt.imshow(img.permute(1, 2, 0))\n    print('Labels:', decode_target(target, text_labels=True))","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:16.158453Z","iopub.execute_input":"2023-11-26T15:30:16.158784Z","iopub.status.idle":"2023-11-26T15:30:16.221528Z","shell.execute_reply.started":"2023-11-26T15:30:16.158756Z","shell.execute_reply":"2023-11-26T15:30:16.220478Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"show_sample(*dataset[0])","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:16.223190Z","iopub.execute_input":"2023-11-26T15:30:16.223831Z","iopub.status.idle":"2023-11-26T15:30:16.756792Z","shell.execute_reply.started":"2023-11-26T15:30:16.223785Z","shell.execute_reply":"2023-11-26T15:30:16.755638Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"show_sample(*dataset[0],invert=True)","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:16.758381Z","iopub.execute_input":"2023-11-26T15:30:16.759207Z","iopub.status.idle":"2023-11-26T15:30:17.204538Z","shell.execute_reply.started":"2023-11-26T15:30:16.759167Z","shell.execute_reply":"2023-11-26T15:30:17.203406Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.manual_seed(10)","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:17.206161Z","iopub.execute_input":"2023-11-26T15:30:17.206519Z","iopub.status.idle":"2023-11-26T15:30:17.266471Z","shell.execute_reply.started":"2023-11-26T15:30:17.206492Z","shell.execute_reply":"2023-11-26T15:30:17.265552Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_pct = 0.1\nval_size = int(val_pct * len(dataset))\ntrain_size = len(dataset) - val_size","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:17.267738Z","iopub.execute_input":"2023-11-26T15:30:17.268090Z","iopub.status.idle":"2023-11-26T15:30:17.320227Z","shell.execute_reply.started":"2023-11-26T15:30:17.268040Z","shell.execute_reply":"2023-11-26T15:30:17.319459Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_size","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:17.321451Z","iopub.execute_input":"2023-11-26T15:30:17.321717Z","iopub.status.idle":"2023-11-26T15:30:17.376984Z","shell.execute_reply.started":"2023-11-26T15:30:17.321694Z","shell.execute_reply":"2023-11-26T15:30:17.376122Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_ds, val_ds = random_split(dataset, [train_size, val_size])\nlen(train_ds), len(val_ds)","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:17.378202Z","iopub.execute_input":"2023-11-26T15:30:17.378598Z","iopub.status.idle":"2023-11-26T15:30:17.439252Z","shell.execute_reply.started":"2023-11-26T15:30:17.378571Z","shell.execute_reply":"2023-11-26T15:30:17.438355Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch_size = 64","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:17.440483Z","iopub.execute_input":"2023-11-26T15:30:17.440815Z","iopub.status.idle":"2023-11-26T15:30:17.495452Z","shell.execute_reply.started":"2023-11-26T15:30:17.440784Z","shell.execute_reply":"2023-11-26T15:30:17.494389Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dl = DataLoader(train_ds, item_tfms=Resize(460),  # Resize all images to have the same dimensions\n                      batch_tfms=aug_transforms(size=224), shuffle=True, num_workers=2, pin_memory=True)\nval_dl = DataLoader(val_ds, item_tfms=Resize(460),  # Resize all images to have the same dimensions\n                    batch_tfms=aug_transforms(size=224), num_workers=2, pin_memory=True)","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:17.496621Z","iopub.execute_input":"2023-11-26T15:30:17.496933Z","iopub.status.idle":"2023-11-26T15:30:17.557096Z","shell.execute_reply.started":"2023-11-26T15:30:17.496907Z","shell.execute_reply":"2023-11-26T15:30:17.556308Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(train_dl)","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:17.558466Z","iopub.execute_input":"2023-11-26T15:30:17.558808Z","iopub.status.idle":"2023-11-26T15:30:17.616733Z","shell.execute_reply.started":"2023-11-26T15:30:17.558780Z","shell.execute_reply":"2023-11-26T15:30:17.615843Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def show_batch(dl): # ,invert = TRUE\n    for images, labels in dl:\n        fig, ax = plt.subplots(figsize=(9, 4))\n        ax.set_xticks([]); ax.set_yticks([])\n        #data = 1-images if invert else images\n        data = images\n        ax.imshow(make_grid(data, nrow=16).permute(1, 2, 0))\n        break","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:17.618037Z","iopub.execute_input":"2023-11-26T15:30:17.618379Z","iopub.status.idle":"2023-11-26T15:30:17.672925Z","shell.execute_reply.started":"2023-11-26T15:30:17.618323Z","shell.execute_reply":"2023-11-26T15:30:17.672038Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"show_batch(train_dl)","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:17.673994Z","iopub.execute_input":"2023-11-26T15:30:17.674577Z","iopub.status.idle":"2023-11-26T15:30:21.698895Z","shell.execute_reply.started":"2023-11-26T15:30:17.674552Z","shell.execute_reply":"2023-11-26T15:30:21.697608Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def F_score(output, label, threshold=0.5, beta=1):\n    prob = output > threshold\n    label = label > threshold\n\n    TP = (prob & label).sum(1).float()\n    TN = ((~prob) & (~label)).sum(1).float()\n    FP = (prob & (~label)).sum(1).float()\n    FN = ((~prob) & label).sum(1).float()\n\n    precision = torch.mean(TP / (TP + FP + 1e-12))\n    recall = torch.mean(TP / (TP + FN + 1e-12))\n    F2 = (1 + beta**2) * precision * recall / (beta**2 * precision + recall + 1e-12)\n    return F2.mean(0)","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:21.701278Z","iopub.execute_input":"2023-11-26T15:30:21.701804Z","iopub.status.idle":"2023-11-26T15:30:21.774986Z","shell.execute_reply.started":"2023-11-26T15:30:21.701738Z","shell.execute_reply":"2023-11-26T15:30:21.773736Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MultilabelImageClassificationBase(nn.Module):\n    def training_step(self, batch):\n        images, targets = batch \n        \n        images = images.unsqueeze(0)\n        targets = targets.unsqueeze(0)\n        \n        out = self(images)                      \n        loss = F.binary_cross_entropy(out, targets)      \n        return loss\n    \n    def validation_step(self, batch):\n        images, targets = batch \n        images = images.unsqueeze(0)\n        targets = targets.unsqueeze(0)\n        out = self(images)                           # Generate predictions\n        loss = F.binary_cross_entropy(out, targets)  # Calculate loss\n        score = F_score(out, targets)\n        return {'val_loss': loss.detach(), 'val_score': score.detach() }\n        \n    def validation_epoch_end(self, outputs):\n        batch_losses = [x['val_loss'] for x in outputs]\n        epoch_loss = torch.stack(batch_losses).mean()   # Combine losses\n        batch_scores = [x['val_score'] for x in outputs]\n        epoch_score = torch.stack(batch_scores).mean()      # Combine accuracies\n        return {'val_loss': epoch_loss.item(), 'val_score': epoch_score.item()}\n    \n    def epoch_end(self, epoch, result):\n        print(\"Epoch [{}], train_loss: {:.4f}, val_loss: {:.4f}, val_score: {:.4f}\".format(\n            epoch, result['train_loss'], result['val_loss'], result['val_score']))","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:21.776828Z","iopub.execute_input":"2023-11-26T15:30:21.777580Z","iopub.status.idle":"2023-11-26T15:30:21.845836Z","shell.execute_reply.started":"2023-11-26T15:30:21.777543Z","shell.execute_reply":"2023-11-26T15:30:21.844722Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class YelpResCnnModel(MultilabelImageClassificationBase):\n    def __init__(self, num_classes=9):\n        super().__init__()\n        self.network = nn.Sequential(\n            nn.Conv2d(3, 32, kernel_size=3, padding=1),\n            nn.ReLU(),\n            nn.MaxPool2d(2, 2),\n            nn.Conv2d(32, 64, kernel_size=3, stride=1, padding=1),\n            nn.ReLU(),\n            nn.MaxPool2d(2, 2),\n\n            nn.Conv2d(64, 64, kernel_size=3, stride=1, padding=1),\n            nn.ReLU(),\n            nn.MaxPool2d(2, 2),\n            nn.Conv2d(64, 128, kernel_size=3, stride=1, padding=1),\n            nn.ReLU(),\n            nn.MaxPool2d(2, 2),\n\n            nn.Conv2d(128, 256, kernel_size=3, stride=1, padding=1),\n            nn.ReLU(),\n            nn.AdaptiveAvgPool2d(1),\n\n            nn.Flatten(),\n            nn.Linear(256, num_classes),  # Update the output size to match num_classes\n            nn.Sigmoid(),\n        )\n        \n    def forward(self, xb):\n        return self.network(xb)","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:21.847456Z","iopub.execute_input":"2023-11-26T15:30:21.847812Z","iopub.status.idle":"2023-11-26T15:30:21.913399Z","shell.execute_reply.started":"2023-11-26T15:30:21.847778Z","shell.execute_reply":"2023-11-26T15:30:21.912397Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# model = YelpResCnnModel()\n# model","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:21.915007Z","iopub.execute_input":"2023-11-26T15:30:21.915417Z","iopub.status.idle":"2023-11-26T15:30:21.977162Z","shell.execute_reply.started":"2023-11-26T15:30:21.915383Z","shell.execute_reply":"2023-11-26T15:30:21.976105Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_default_device():\n    \"\"\"Pick GPU if available, else CPU\"\"\"\n    if torch.cuda.is_available():\n        return torch.device('cuda')\n    else:\n        return torch.device('cpu')\n    \ndef to_device(data, device):\n    \"\"\"Move tensor(s) to chosen device\"\"\"\n    if isinstance(data, (list,tuple)):\n        return [to_device(x, device) for x in data]\n    return data.to(device, non_blocking=True)\n\nclass DeviceDataLoader():\n    \"\"\"Wrap a dataloader to move data to a device\"\"\"\n    def __init__(self, dl, device):\n        self.dl = dl\n        self.device = device\n        \n    def __iter__(self):\n        \"\"\"Yield a batch of data after moving it to device\"\"\"\n        for b in self.dl: \n            yield to_device(b, self.device)\n\n    def __len__(self):\n        \"\"\"Number of batches\"\"\"\n        return len(self.dl)","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:21.978854Z","iopub.execute_input":"2023-11-26T15:30:21.979660Z","iopub.status.idle":"2023-11-26T15:30:22.045424Z","shell.execute_reply.started":"2023-11-26T15:30:21.979619Z","shell.execute_reply":"2023-11-26T15:30:22.044368Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = get_default_device()\ndevice","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:22.047003Z","iopub.execute_input":"2023-11-26T15:30:22.047469Z","iopub.status.idle":"2023-11-26T15:30:22.110772Z","shell.execute_reply.started":"2023-11-26T15:30:22.047427Z","shell.execute_reply":"2023-11-26T15:30:22.109683Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class YelpResCnnModel2(MultilabelImageClassificationBase):\n    def __init__(self):\n        super().__init__()\n        # Use a pretrained model\n        self.network = models.resnet34(pretrained=True)\n        # Replace last layer\n        num_ftrs = self.network.fc.in_features\n        self.network.fc = nn.Linear(num_ftrs, 9)\n    \n    def forward(self, xb):\n        return torch.sigmoid(self.network(xb))","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:22.113443Z","iopub.execute_input":"2023-11-26T15:30:22.114237Z","iopub.status.idle":"2023-11-26T15:30:22.175914Z","shell.execute_reply.started":"2023-11-26T15:30:22.114193Z","shell.execute_reply":"2023-11-26T15:30:22.174799Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = YelpResCnnModel2()\nmodel","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:22.177312Z","iopub.execute_input":"2023-11-26T15:30:22.177683Z","iopub.status.idle":"2023-11-26T15:30:23.228925Z","shell.execute_reply.started":"2023-11-26T15:30:22.177650Z","shell.execute_reply":"2023-11-26T15:30:23.227811Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dl = DeviceDataLoader(train_dl, device)\nval_dl = DeviceDataLoader(val_dl, device)\nto_device(model, device);","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:23.230077Z","iopub.execute_input":"2023-11-26T15:30:23.230386Z","iopub.status.idle":"2023-11-26T15:30:23.311403Z","shell.execute_reply.started":"2023-11-26T15:30:23.230361Z","shell.execute_reply":"2023-11-26T15:30:23.310523Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def try_batch(dl):\n    for images, labels in dl:\n        images = images.unsqueeze(0)\n        print('images.shape:', images.shape)\n        out = model(images)\n        print('out.shape:', out.shape)\n        print('out[0]:', out[0])\n        break\n\ntry_batch(train_dl)","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:23.313171Z","iopub.execute_input":"2023-11-26T15:30:23.313788Z","iopub.status.idle":"2023-11-26T15:30:30.170786Z","shell.execute_reply.started":"2023-11-26T15:30:23.313753Z","shell.execute_reply":"2023-11-26T15:30:30.169422Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tqdm.notebook import tqdm","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:30.172537Z","iopub.execute_input":"2023-11-26T15:30:30.172919Z","iopub.status.idle":"2023-11-26T15:30:30.242240Z","shell.execute_reply.started":"2023-11-26T15:30:30.172887Z","shell.execute_reply":"2023-11-26T15:30:30.241118Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@torch.no_grad()\ndef evaluate(model, val_loader):\n    model.eval()\n    outputs = [model.validation_step(batch) for batch in val_loader]\n    return model.validation_epoch_end(outputs)\n\ndef fit(epochs, lr, model, train_loader, val_loader, opt_func=torch.optim.SGD):\n    torch.cuda.empty_cache()\n    history = []\n    optimizer = opt_func(model.parameters(), lr)\n    for epoch in range(epochs):\n        # Training Phase \n        model.train()\n        train_losses = []\n        for batch in tqdm(train_loader):\n            loss = model.training_step(batch)\n            train_losses.append(loss)\n            loss.backward()\n            optimizer.step()\n            optimizer.zero_grad()\n        # Validation phase\n        result = evaluate(model, val_loader)\n        result['train_loss'] = torch.stack(train_losses).mean().item()\n        model.epoch_end(epoch, result)\n        history.append(result)\n    return history","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:30.243954Z","iopub.execute_input":"2023-11-26T15:30:30.244860Z","iopub.status.idle":"2023-11-26T15:30:30.310183Z","shell.execute_reply.started":"2023-11-26T15:30:30.244809Z","shell.execute_reply":"2023-11-26T15:30:30.309109Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Modeling","metadata":{}},{"cell_type":"code","source":"model = to_device(YelpResCnnModel2(), device)","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:30.311638Z","iopub.execute_input":"2023-11-26T15:30:30.312360Z","iopub.status.idle":"2023-11-26T15:30:30.862861Z","shell.execute_reply.started":"2023-11-26T15:30:30.312298Z","shell.execute_reply":"2023-11-26T15:30:30.862086Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"evaluate(model, val_dl)","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:30.864023Z","iopub.execute_input":"2023-11-26T15:30:30.864334Z","iopub.status.idle":"2023-11-26T15:30:57.801582Z","shell.execute_reply.started":"2023-11-26T15:30:30.864307Z","shell.execute_reply":"2023-11-26T15:30:57.800189Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"num_epochs = 2\nopt_func = torch.optim.Adam\nlr = 1e-2\n\nhistory = fit(num_epochs, lr, model, train_dl, val_dl, opt_func)","metadata":{"execution":{"iopub.status.busy":"2023-11-26T15:30:57.803058Z","iopub.execute_input":"2023-11-26T15:30:57.803810Z","iopub.status.idle":"2023-11-26T16:00:17.737253Z","shell.execute_reply.started":"2023-11-26T15:30:57.803782Z","shell.execute_reply":"2023-11-26T16:00:17.736194Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def predict_single(image):\n    xb = image.unsqueeze(0)\n    xb = to_device(xb, device)\n    preds = model(xb)\n    prediction = preds[0]\n    print(\"Prediction: \", prediction)\n    show_sample(image, prediction)","metadata":{"execution":{"iopub.status.busy":"2023-11-26T16:00:18.019668Z","iopub.execute_input":"2023-11-26T16:00:18.020152Z","iopub.status.idle":"2023-11-26T16:00:18.133861Z","shell.execute_reply.started":"2023-11-26T16:00:18.020118Z","shell.execute_reply":"2023-11-26T16:00:18.132713Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dataset = YelpResPhoto(TEST_CSV, TEST_DIR, transform=transform)","metadata":{"execution":{"iopub.status.busy":"2023-11-26T16:00:18.136501Z","iopub.execute_input":"2023-11-26T16:00:18.137232Z","iopub.status.idle":"2023-11-26T16:00:18.220209Z","shell.execute_reply.started":"2023-11-26T16:00:18.137186Z","shell.execute_reply":"2023-11-26T16:00:18.219175Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dataset","metadata":{"execution":{"iopub.status.busy":"2023-11-26T16:00:18.222287Z","iopub.execute_input":"2023-11-26T16:00:18.222767Z","iopub.status.idle":"2023-11-26T16:00:18.287822Z","shell.execute_reply.started":"2023-11-26T16:00:18.222724Z","shell.execute_reply":"2023-11-26T16:00:18.286687Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(test_dataset)","metadata":{"execution":{"iopub.status.busy":"2023-11-26T16:00:18.289923Z","iopub.execute_input":"2023-11-26T16:00:18.290643Z","iopub.status.idle":"2023-11-26T16:00:18.359657Z","shell.execute_reply.started":"2023-11-26T16:00:18.290597Z","shell.execute_reply":"2023-11-26T16:00:18.358212Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img, target = test_dataset[0]\nimg.shape","metadata":{"execution":{"iopub.status.busy":"2023-11-26T16:00:18.362038Z","iopub.execute_input":"2023-11-26T16:00:18.362474Z","iopub.status.idle":"2023-11-26T16:00:18.436962Z","shell.execute_reply.started":"2023-11-26T16:00:18.362439Z","shell.execute_reply":"2023-11-26T16:00:18.435822Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predict_single(test_dataset[100][0])","metadata":{"execution":{"iopub.status.busy":"2023-11-26T16:00:18.438724Z","iopub.execute_input":"2023-11-26T16:00:18.439516Z","iopub.status.idle":"2023-11-26T16:00:19.055748Z","shell.execute_reply.started":"2023-11-26T16:00:18.439471Z","shell.execute_reply":"2023-11-26T16:00:19.054652Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predict_single(test_dataset[63][0])","metadata":{"execution":{"iopub.status.busy":"2023-11-26T16:00:19.057563Z","iopub.execute_input":"2023-11-26T16:00:19.058163Z","iopub.status.idle":"2023-11-26T16:00:19.527431Z","shell.execute_reply.started":"2023-11-26T16:00:19.058113Z","shell.execute_reply":"2023-11-26T16:00:19.526351Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dl = DeviceDataLoader(DataLoader(test_dataset, item_tfms=Resize(460), batch_tfms=aug_transforms(size=224), num_workers=2, pin_memory=True), device)","metadata":{"execution":{"iopub.status.busy":"2023-11-26T16:00:19.528907Z","iopub.execute_input":"2023-11-26T16:00:19.529271Z","iopub.status.idle":"2023-11-26T16:00:19.590226Z","shell.execute_reply.started":"2023-11-26T16:00:19.529238Z","shell.execute_reply":"2023-11-26T16:00:19.589137Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#@torch.no_grad()\n# def predict_dl(dl, model):\n#     torch.cuda.empty_cache()\n#     batch_probs = []\n#     for xb, _ in tqdm(dl):\n        \n#         xb = xb.unsqueeze(0)\n        \n#         probs = model(xb)\n#         batch_probs.append(probs.cpu().detach())\n#     batch_probs = torch.cat(batch_probs)\n#     return [decode_target(x) for x in batch_probs]\n\ndef predict_dl(dl, model):\n    torch.cuda.empty_cache()\n    model.eval()\n    batch_probs = []\n\n    for xb, _ in tqdm(dl):\n        # Add an extra dimension to the input tensor\n        xb = xb.unsqueeze(0)\n        \n        with torch.no_grad():\n            probs = model(xb)\n        \n        batch_probs.append(probs.cpu().detach())\n\n    batch_probs = torch.cat(batch_probs)\n    return [decode_target(x) for x in batch_probs]","metadata":{"execution":{"iopub.status.busy":"2023-11-26T16:00:19.594410Z","iopub.execute_input":"2023-11-26T16:00:19.594739Z","iopub.status.idle":"2023-11-26T16:00:19.652486Z","shell.execute_reply.started":"2023-11-26T16:00:19.594702Z","shell.execute_reply":"2023-11-26T16:00:19.651462Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_preds = predict_dl(test_dl, model)","metadata":{"execution":{"iopub.status.busy":"2023-11-26T16:00:19.656221Z","iopub.execute_input":"2023-11-26T16:00:19.657689Z","iopub.status.idle":"2023-11-26T16:02:04.817138Z","shell.execute_reply.started":"2023-11-26T16:00:19.657654Z","shell.execute_reply":"2023-11-26T16:02:04.815862Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_preds","metadata":{"execution":{"iopub.status.busy":"2023-11-26T16:02:04.832724Z","iopub.execute_input":"2023-11-26T16:02:04.833155Z","iopub.status.idle":"2023-11-26T16:02:04.929747Z","shell.execute_reply.started":"2023-11-26T16:02:04.833118Z","shell.execute_reply":"2023-11-26T16:02:04.928208Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_df = pd.read_csv(TEST_CSV)\nsubmission_df['labels'] = test_preds\nsubmission_df = submission_df.drop('photo_id', axis=1)\nsubmission_df.sample(n=15)","metadata":{"execution":{"iopub.status.busy":"2023-11-26T16:02:04.931852Z","iopub.execute_input":"2023-11-26T16:02:04.932449Z","iopub.status.idle":"2023-11-26T16:02:05.028477Z","shell.execute_reply.started":"2023-11-26T16:02:04.932389Z","shell.execute_reply":"2023-11-26T16:02:05.027177Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_fname = 'submission.csv'\nsubmission_df.to_csv(sub_fname, index=False)","metadata":{"execution":{"iopub.status.busy":"2023-11-26T16:02:05.030605Z","iopub.execute_input":"2023-11-26T16:02:05.031064Z","iopub.status.idle":"2023-11-26T16:02:05.121193Z","shell.execute_reply.started":"2023-11-26T16:02:05.031021Z","shell.execute_reply":"2023-11-26T16:02:05.120092Z"},"trusted":true},"execution_count":null,"outputs":[]}]}