{"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\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":"2022-07-26T14:46:33.664358Z","iopub.execute_input":"2022-07-26T14:46:33.664709Z","iopub.status.idle":"2022-07-26T14:46:33.674544Z","shell.execute_reply.started":"2022-07-26T14:46:33.664679Z","shell.execute_reply":"2022-07-26T14:46:33.673503Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn.functional as F\nimport torch.nn as nn\nimport pandas as pd\nimport numpy as np\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nfrom PIL import Image\nimport os\nimport matplotlib.pyplot as plt\nimport torchvision.models as models\n# This is for the progress bar.\nfrom tqdm import tqdm\nimport seaborn as sns","metadata":{"execution":{"iopub.status.busy":"2022-07-26T14:46:33.676785Z","iopub.execute_input":"2022-07-26T14:46:33.677519Z","iopub.status.idle":"2022-07-26T14:46:33.685205Z","shell.execute_reply.started":"2022-07-26T14:46:33.677481Z","shell.execute_reply":"2022-07-26T14:46:33.684260Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import zipfile\nwith zipfile.ZipFile('/kaggle/input/noaa-right-whale-recognition/imgs.zip', 'r') as zip_ref:\n    zip_ref.extractall('/kaggle/temp/')","metadata":{"execution":{"iopub.status.busy":"2022-07-26T14:46:33.687224Z","iopub.execute_input":"2022-07-26T14:46:33.687889Z","iopub.status.idle":"2022-07-26T14:49:33.577550Z","shell.execute_reply.started":"2022-07-26T14:46:33.687852Z","shell.execute_reply":"2022-07-26T14:49:33.576513Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def res_model(num_classes, feature_extract = False, use_pretrained=True):\n\n    model_ft = models.resnet34(pretrained=use_pretrained)\n    set_parameter_requires_grad(model_ft, feature_extract)\n    num_ftrs = model_ft.fc.in_features\n    model_ft.fc = nn.Sequential(nn.Linear(num_ftrs, num_classes))\n\n    return model_ft","metadata":{"execution":{"iopub.status.busy":"2022-07-26T14:49:33.580289Z","iopub.execute_input":"2022-07-26T14:49:33.581391Z","iopub.status.idle":"2022-07-26T14:49:33.590462Z","shell.execute_reply.started":"2022-07-26T14:49:33.581349Z","shell.execute_reply":"2022-07-26T14:49:33.589601Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_device():\n    return 'cuda' if torch.cuda.is_available() else 'cpu'\n\ndevice = get_device()\nprint(device)","metadata":{"execution":{"iopub.status.busy":"2022-07-26T14:49:33.591833Z","iopub.execute_input":"2022-07-26T14:49:33.592257Z","iopub.status.idle":"2022-07-26T14:49:34.555321Z","shell.execute_reply.started":"2022-07-26T14:49:33.592219Z","shell.execute_reply":"2022-07-26T14:49:34.554287Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class WhaleData(Dataset):\n    def __init__(self, csv_path, file_path, mode='train', valid_ratio=0.2, resize_height=256, resize_width=256):\n        self.resize_height = resize_height\n        self.resize_width = resize_width\n\n        self.file_path = file_path\n        self.mode = mode\n        self.data_info = pd.read_csv(csv_path, header=None)  #header=None是去掉表头部分\n        # 计算 length\n        self.data_len = len(self.data_info.index) - 1\n        self.train_len = int(self.data_len * (1 - valid_ratio))\n        \n        if mode == 'train':\n            # 第一列包含图像文件的名称\n            self.train_image = np.asarray(self.data_info.iloc[1:self.train_len, 0])  #self.data_info.iloc[1:,0]表示读取第一列，从第二行开始到train_len\n            # 第二列是图像的 label\n            self.train_label = np.asarray(self.data_info.iloc[1:self.train_len, 1])\n            self.image_arr = self.train_image \n            self.label_arr = self.train_label\n        elif mode == 'valid':\n            self.valid_image = np.asarray(self.data_info.iloc[self.train_len:, 0])  \n            self.valid_label = np.asarray(self.data_info.iloc[self.train_len:, 1])\n            self.image_arr = self.valid_image\n            self.label_arr = self.valid_label\n        elif mode == 'test':\n            self.test_image = np.asarray(self.data_info.iloc[1:, 0])\n            self.image_arr = self.test_image\n            \n        self.real_len = len(self.image_arr)\n\n        print('Finished reading the {} set of Leaves Dataset ({} samples found)'\n              .format(mode, self.real_len))\n        \n    def __getitem__(self, index):\n        # 从 image_arr中得到索引对应的文件名\n        single_image_name = self.image_arr[index]\n\n        # 读取图像文件\n        img_as_img = Image.open(self.file_path + single_image_name)\n        if self.mode == 'train':\n            transform = transforms.Compose([\n                transforms.Resize((224, 224)),\n                transforms.RandomHorizontalFlip(p=0.5),   #随机水平翻转 选择一个概率\n                transforms.ToTensor()\n            ])\n        else:\n            # valid和test不做数据增强\n            transform = transforms.Compose([\n                transforms.Resize((224, 224)),\n                transforms.ToTensor()\n            ])\n        \n        img_as_img = transform(img_as_img)\n        \n        if self.mode == 'test':\n            return img_as_img\n        else:\n            # 得到图像的 string label\n            label = self.label_arr[index]\n            # number label\n            number_label = class_to_num[label]\n\n            return img_as_img, number_label  #返回每一个index对应的图片数据和对应的label\n\n    def __len__(self):\n        return self.real_len","metadata":{"execution":{"iopub.status.busy":"2022-07-26T14:49:34.558738Z","iopub.execute_input":"2022-07-26T14:49:34.559485Z","iopub.status.idle":"2022-07-26T14:49:34.576438Z","shell.execute_reply.started":"2022-07-26T14:49:34.559436Z","shell.execute_reply":"2022-07-26T14:49:34.575422Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_data = pd.read_csv('../input/noaa-right-whale-recognition/sample_submission.csv')\ntest = test_data['Image']\ntest.to_csv('/kaggle/working/test.csv', index = False)\ntest","metadata":{"execution":{"iopub.status.busy":"2022-07-26T14:49:34.577831Z","iopub.execute_input":"2022-07-26T14:49:34.578287Z","iopub.status.idle":"2022-07-26T14:49:34.823628Z","shell.execute_reply.started":"2022-07-26T14:49:34.578248Z","shell.execute_reply":"2022-07-26T14:49:34.822670Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_path = '/kaggle/temp/imgs/'\ntest_path = '/kaggle/working/test.csv'\ntest_dataset = WhaleData(test_path, img_path, mode='test')\nprint(test_dataset)","metadata":{"execution":{"iopub.status.busy":"2022-07-26T14:49:34.825196Z","iopub.execute_input":"2022-07-26T14:49:34.825541Z","iopub.status.idle":"2022-07-26T14:49:34.837979Z","shell.execute_reply.started":"2022-07-26T14:49:34.825507Z","shell.execute_reply":"2022-07-26T14:49:34.836827Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def set_parameter_requires_grad(model, feature_extracting):\n    if feature_extracting:\n        model = model\n        for param in model.parameters():\n            param.requires_grad = False","metadata":{"execution":{"iopub.status.busy":"2022-07-26T14:49:34.839266Z","iopub.execute_input":"2022-07-26T14:49:34.840076Z","iopub.status.idle":"2022-07-26T14:49:34.845321Z","shell.execute_reply.started":"2022-07-26T14:49:34.840040Z","shell.execute_reply":"2022-07-26T14:49:34.844391Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_loader = torch.utils.data.DataLoader(\n        dataset=test_dataset,\n        batch_size=8, \n        shuffle=False,\n        num_workers=0\n    )","metadata":{"execution":{"iopub.status.busy":"2022-07-26T14:49:34.846446Z","iopub.execute_input":"2022-07-26T14:49:34.847486Z","iopub.status.idle":"2022-07-26T14:49:34.854552Z","shell.execute_reply.started":"2022-07-26T14:49:34.847439Z","shell.execute_reply":"2022-07-26T14:49:34.853674Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = res_model(447)\nprint(model)","metadata":{"execution":{"iopub.status.busy":"2022-07-26T14:49:34.855803Z","iopub.execute_input":"2022-07-26T14:49:34.856222Z","iopub.status.idle":"2022-07-26T14:49:35.729545Z","shell.execute_reply.started":"2022-07-26T14:49:34.856185Z","shell.execute_reply":"2022-07-26T14:49:35.728566Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = res_model(447)\n\n# create model and load weights from checkpoint\nmodel = model.to(device)\nmodel_path = '../input/kaggle/pre_res_model.ckpt'\nmodel.load_state_dict(torch.load(model_path))\n\n# Make sure the model is in eval mode.\n# Some modules like Dropout or BatchNorm affect if the model is in training mode.\nmodel.eval()\n\n# Initialize a list to store the predictions.\npredictions = []\n\nfor batch in tqdm(test_loader):\n    \n    imgs = batch\n    with torch.no_grad():\n        logits = model(imgs.to(device))\n    \n    # Take the class with greatest logit as prediction and record it.\n    predictions.extend((F.softmax(logits)).cpu().numpy().tolist())","metadata":{"execution":{"iopub.status.busy":"2022-07-26T15:20:56.779694Z","iopub.execute_input":"2022-07-26T15:20:56.780407Z","iopub.status.idle":"2022-07-26T15:39:32.933351Z","shell.execute_reply.started":"2022-07-26T15:20:56.780370Z","shell.execute_reply":"2022-07-26T15:39:32.932376Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(predictions)","metadata":{"execution":{"iopub.status.busy":"2022-07-26T15:41:06.486874Z","iopub.execute_input":"2022-07-26T15:41:06.487862Z","iopub.status.idle":"2022-07-26T15:41:06.494011Z","shell.execute_reply.started":"2022-07-26T15:41:06.487814Z","shell.execute_reply":"2022-07-26T15:41:06.493059Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(predictions[0])","metadata":{"execution":{"iopub.status.busy":"2022-07-26T15:41:03.303695Z","iopub.execute_input":"2022-07-26T15:41:03.304690Z","iopub.status.idle":"2022-07-26T15:41:03.312227Z","shell.execute_reply.started":"2022-07-26T15:41:03.304650Z","shell.execute_reply":"2022-07-26T15:41:03.311247Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_data","metadata":{"execution":{"iopub.status.busy":"2022-07-26T15:41:10.681485Z","iopub.execute_input":"2022-07-26T15:41:10.681854Z","iopub.status.idle":"2022-07-26T15:41:10.721903Z","shell.execute_reply.started":"2022-07-26T15:41:10.681822Z","shell.execute_reply":"2022-07-26T15:41:10.720818Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predictions = np.array(predictions)\npredictions.shape","metadata":{"execution":{"iopub.status.busy":"2022-07-26T15:41:15.828013Z","iopub.execute_input":"2022-07-26T15:41:15.828404Z","iopub.status.idle":"2022-07-26T15:41:16.177478Z","shell.execute_reply.started":"2022-07-26T15:41:15.828370Z","shell.execute_reply":"2022-07-26T15:41:16.176460Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#predictions = torch.from_numpy(predictions)\n#predictions = F.softmax(predictions)","metadata":{"execution":{"iopub.status.busy":"2022-07-26T15:20:30.231156Z","iopub.execute_input":"2022-07-26T15:20:30.231519Z","iopub.status.idle":"2022-07-26T15:20:30.237470Z","shell.execute_reply.started":"2022-07-26T15:20:30.231488Z","shell.execute_reply":"2022-07-26T15:20:30.236077Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predictions","metadata":{"execution":{"iopub.status.busy":"2022-07-26T15:41:19.394927Z","iopub.execute_input":"2022-07-26T15:41:19.395892Z","iopub.status.idle":"2022-07-26T15:41:19.403950Z","shell.execute_reply.started":"2022-07-26T15:41:19.395843Z","shell.execute_reply":"2022-07-26T15:41:19.402833Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(predictions[:,1])","metadata":{"execution":{"iopub.status.busy":"2022-07-26T15:41:24.782846Z","iopub.execute_input":"2022-07-26T15:41:24.783547Z","iopub.status.idle":"2022-07-26T15:41:24.789650Z","shell.execute_reply.started":"2022-07-26T15:41:24.783508Z","shell.execute_reply":"2022-07-26T15:41:24.788588Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"label = list(test_data)","metadata":{"execution":{"iopub.status.busy":"2022-07-26T15:41:26.957588Z","iopub.execute_input":"2022-07-26T15:41:26.959063Z","iopub.status.idle":"2022-07-26T15:41:26.969297Z","shell.execute_reply.started":"2022-07-26T15:41:26.959014Z","shell.execute_reply":"2022-07-26T15:41:26.968346Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for i in range(447):\n    test_data[label[i+1]] = pd.Series(predictions[:,i])","metadata":{"execution":{"iopub.status.busy":"2022-07-26T15:41:29.817630Z","iopub.execute_input":"2022-07-26T15:41:29.818686Z","iopub.status.idle":"2022-07-26T15:41:29.938569Z","shell.execute_reply.started":"2022-07-26T15:41:29.818641Z","shell.execute_reply":"2022-07-26T15:41:29.937608Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_data","metadata":{"execution":{"iopub.status.busy":"2022-07-26T15:41:31.820526Z","iopub.execute_input":"2022-07-26T15:41:31.820898Z","iopub.status.idle":"2022-07-26T15:41:31.864380Z","shell.execute_reply.started":"2022-07-26T15:41:31.820865Z","shell.execute_reply":"2022-07-26T15:41:31.863365Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"saveFileName = '/kaggle/working/submission.csv'\nsubmission = test_data\nsubmission.to_csv(saveFileName, index=False)","metadata":{"execution":{"iopub.status.busy":"2022-07-26T15:41:36.890665Z","iopub.execute_input":"2022-07-26T15:41:36.891246Z","iopub.status.idle":"2022-07-26T15:41:44.131121Z","shell.execute_reply.started":"2022-07-26T15:41:36.891209Z","shell.execute_reply":"2022-07-26T15:41:44.129859Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission","metadata":{"execution":{"iopub.status.busy":"2022-07-26T15:41:46.583840Z","iopub.execute_input":"2022-07-26T15:41:46.584202Z","iopub.status.idle":"2022-07-26T15:41:46.628727Z","shell.execute_reply.started":"2022-07-26T15:41:46.584170Z","shell.execute_reply":"2022-07-26T15:41:46.627740Z"},"trusted":true},"execution_count":null,"outputs":[]}]}