{"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":45867,"databundleVersionId":6924515,"sourceType":"competition"},{"sourceId":7261880,"sourceType":"datasetVersion","datasetId":4208723}],"dockerImageVersionId":30627,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# Data handling\nimport pandas as pd\nimport numpy as np\n\n# Data visualization\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport cv2\nfrom PIL import Image\nImage.MAX_IMAGE_PIXELS = None\n\n# Preprocessing\nfrom sklearn.model_selection import train_test_split as tts\nfrom sklearn.utils.class_weight import compute_class_weight\n\n# Torch\nimport torch\nfrom torch import nn, optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision.models import vit_b_16, ViT_B_16_Weights\nfrom torchinfo import summary\n\n# Metrics\nfrom sklearn.metrics import balanced_accuracy_score\nfrom sklearn.metrics import confusion_matrix\n\n# os\nimport os\n\n# Path\nfrom pathlib import Path\n\n# random\nimport random\n\n# OrderedDict\nfrom collections import OrderedDict\n\n# tqdm\nfrom tqdm.auto import tqdm\n\n# warnings\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\nimport torchvision.transforms as transforms","metadata":{"execution":{"iopub.status.busy":"2023-12-23T12:14:22.655523Z","iopub.execute_input":"2023-12-23T12:14:22.65597Z","iopub.status.idle":"2023-12-23T12:14:28.427237Z","shell.execute_reply.started":"2023-12-23T12:14:22.655943Z","shell.execute_reply":"2023-12-23T12:14:28.426416Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nBATCH_SIZE = 1\nNUM_WORKERS = 4","metadata":{"execution":{"iopub.status.busy":"2023-12-23T12:14:28.428935Z","iopub.execute_input":"2023-12-23T12:14:28.429446Z","iopub.status.idle":"2023-12-23T12:14:28.501686Z","shell.execute_reply.started":"2023-12-23T12:14:28.42941Z","shell.execute_reply":"2023-12-23T12:14:28.499597Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"loaded_model = vit_b_16()\nloaded_model.heads = nn.Sequential(OrderedDict([('head', nn.Linear(in_features = 768, \n                                                                   out_features = 5))]))\nweights = '/kaggle/input/best-model-pth/best_model.pth'\ncheckpoint = torch.load(weights)\n   \nloaded_model.load_state_dict(checkpoint)\n    \nloaded_model.to(device)\n    \nloaded_model.eval()\n    \nlabel_map  = {'CC': 0, 'EC': 1, 'HGSC': 2, 'LGSC': 3, 'MC': 4}","metadata":{"execution":{"iopub.status.busy":"2023-12-23T12:14:28.503035Z","iopub.execute_input":"2023-12-23T12:14:28.503524Z","iopub.status.idle":"2023-12-23T12:14:36.015406Z","shell.execute_reply.started":"2023-12-23T12:14:28.503489Z","shell.execute_reply":"2023-12-23T12:14:36.014568Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"auto_transforms = transforms.Compose([\n    transforms.Resize([256, 256]),\n    transforms.CenterCrop(224),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n])","metadata":{"execution":{"iopub.status.busy":"2023-12-23T12:14:36.017409Z","iopub.execute_input":"2023-12-23T12:14:36.017738Z","iopub.status.idle":"2023-12-23T12:14:36.022951Z","shell.execute_reply.started":"2023-12-23T12:14:36.017712Z","shell.execute_reply":"2023-12-23T12:14:36.022069Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CustomImageDataset(Dataset):\n    def __init__(self, df:pd.DataFrame, image_transforms):\n        self.df = df\n        self.image_transforms = image_transforms\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        image_path = self.df.iloc[idx, 0]\n        image = Image.open(image_path).convert(\"RGB\")\n        image = self.image_transforms(image)\n        \n        return image, os.path.basename(image_path)","metadata":{"execution":{"iopub.status.busy":"2023-12-23T12:14:36.024197Z","iopub.execute_input":"2023-12-23T12:14:36.024547Z","iopub.status.idle":"2023-12-23T12:14:36.045771Z","shell.execute_reply.started":"2023-12-23T12:14:36.024514Z","shell.execute_reply":"2023-12-23T12:14:36.044891Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import multiprocessing\nmultiprocessing_context = multiprocessing.get_context('fork')\n\n# 从测试集文件夹中获取图像路径\ntest_image_folder = f\"/kaggle/input/UBC-OCEAN/test_images\"\ntest_image_paths = [os.path.join(test_image_folder, img) for img in os.listdir(test_image_folder)]\n\n# 创建一个包含图像路径的 DataFrame\ntest_df = pd.DataFrame(data=test_image_paths, columns=['image_path'])\n\n# 使用 CustomImageDataset 加载测试集数据\ntest_dataset = CustomImageDataset(test_df, auto_transforms)\n\ntest_dataloader = DataLoader(dataset = test_dataset, \n                             batch_size = BATCH_SIZE, \n                             shuffle = False, \n                             num_workers = NUM_WORKERS,\n                            multiprocessing_context=multiprocessing_context)","metadata":{"execution":{"iopub.status.busy":"2023-12-23T12:14:36.046987Z","iopub.execute_input":"2023-12-23T12:14:36.047245Z","iopub.status.idle":"2023-12-23T12:14:36.066163Z","shell.execute_reply.started":"2023-12-23T12:14:36.047223Z","shell.execute_reply":"2023-12-23T12:14:36.065362Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd  \nimport os  \n  \nresults = []  \nimage_names = []  \n  \nfor batch in test_dataloader:  \n    images = batch[0]  # 获取图像张量  \n    image_name = batch[1][0]  # 获取图像路径  \n    image_name = os.path.splitext(image_name)[0]  \n    image_names.append(image_name)  \n  \n    images = images.to(device, dtype=torch.float32)  \n  \n    # 模型推理  \n    y_pred_logit = loaded_model(images)  \n    y_pred_prob = y_pred_logit.softmax(dim=1)  \n    y_pred_class = y_pred_prob.argmax(dim=1)  \n    y_pred_class = int(y_pred_class[0])\n#     y_pred_class = int(y_pred_class)  \n#     predictions = label_map[y_pred_class[0].item()]\n    predictions = []  \n    for key, val in label_map.items():  \n        if val == y_pred_class:  \n            predictions.append(key)  \n            break  \n  \n    batch_results = {'image_id': image_names[0], 'label': predictions[0]}  # 从列表中提取单个值  \n    results.append(batch_results)  \n  \n# 将字典列表转换为 DataFrame  \ndf = pd.DataFrame(results)  \n  \n# 保存为 CSV 文件  \ndf.to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2023-12-23T12:14:36.067283Z","iopub.execute_input":"2023-12-23T12:14:36.06762Z","iopub.status.idle":"2023-12-23T12:15:03.93381Z","shell.execute_reply.started":"2023-12-23T12:14:36.067569Z","shell.execute_reply":"2023-12-23T12:15:03.932646Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\n\n# 从CSV文件读取数据\ndata = pd.read_csv('submission.csv')\n\n# 查看DataFrame对象的前几行\nprint(data.head())","metadata":{"execution":{"iopub.status.busy":"2023-12-23T12:15:03.935258Z","iopub.execute_input":"2023-12-23T12:15:03.935554Z","iopub.status.idle":"2023-12-23T12:15:03.952592Z","shell.execute_reply.started":"2023-12-23T12:15:03.935525Z","shell.execute_reply":"2023-12-23T12:15:03.951669Z"},"trusted":true},"execution_count":null,"outputs":[]}]}