{"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":"none","dataSources":[{"sourceId":45867,"databundleVersionId":6924515,"sourceType":"competition"},{"sourceId":6774400,"sourceType":"datasetVersion","datasetId":3895136},{"sourceId":6774553,"sourceType":"datasetVersion","datasetId":3898019},{"sourceId":6984590,"sourceType":"datasetVersion","datasetId":4014175},{"sourceId":7325443,"sourceType":"datasetVersion","datasetId":4042516}],"dockerImageVersionId":30588,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# ==============================\n# setup pyvips for GPU\n# ==============================\n\n!ls /kaggle/input/pyvips-python-and-deb-package-gpu\n# intall the deb packages\n!yes | dpkg -i --force-depends /kaggle/input/pyvips-python-and-deb-package-gpu/linux_packages/archives/*.deb\n# install the python wrapper\n!pip install pyvips -f /kaggle/input/pyvips-python-and-deb-package-gpu/python_packages/ --no-index\n\n\n\n# # ==============================\n# # setup pyvips for CPU\n# # ==============================\n\n# !ls /kaggle/input/pyvips-python-and-deb-package\n# # intall the deb packages\n# !dpkg -i --force-depends /kaggle/input/pyvips-python-and-deb-package/linux_packages/archives/*.deb\n# # install the python wrapper\n# !pip install pyvips -f /kaggle/input/pyvips-python-and-deb-package/python_packages/ --no-index\n# !pip list | grep pyvips\n\n\n\n# # ==============================\n# # setup pyvips without using external dataset\n# # ==============================\n\n# # setup pyvips\n\n# !sudo apt-get update\n# !sudo apt-get install libvips-dev -y --no-install-recommends --download-only -o dir::cache='./'\n\n# !mkdir ./libvips\n# !mv ./archives/* ./libvips\n# !rm -rf ./archives\n# !ls ./libvips\n\n# !yes | sudo dpkg -i ./libvips/*.deb\n\n# !pip install pyvips\n# !pip wheel pyvips\n# !mkdir pyvips\n# !mv *.whl ./pyvips","metadata":{"_kg_hide-output":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2024-01-03T09:27:25.621863Z","iopub.execute_input":"2024-01-03T09:27:25.622143Z","iopub.status.idle":"2024-01-03T09:28:27.279949Z","shell.execute_reply.started":"2024-01-03T09:27:25.622117Z","shell.execute_reply":"2024-01-03T09:28:27.278811Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport pandas as pd\nimport numpy as np\n\nimport statistics\n\nfrom PIL import Image\n\nimport os\nimport gc\nimport time\nfrom IPython import display\nimport glob\nimport random\nimport shutil\n\nfrom collections import defaultdict\nimport copy\nimport cv2\nimport seaborn as sns\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nimport timm\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import models, transforms, datasets\n\nfrom joblib import Parallel, delayed\nfrom joblib.externals.loky.backend.context import get_context\nimport multiprocessing as mproc\nfrom tqdm.auto import tqdm\n\n# if not os.path.exists('/kaggle/working/tmp'):\n#     os.mkdir('/kaggle/working/tmp')\n# os.environ['TMPDIR'] = '/kaggle/working/tmp'\n# !export TMPDIR='/kaggle/working/tmp'\n\nos.environ['VIPS_CONCURRENCY'] = '4'\nos.environ['VIPS_DISC_THRESHOLD'] = '14gb' #use disk caching instead of memory when the image exceeds threshold\nimport pyvips\n\n\n# ===========================================================================================================================\n# explore devices\n# ===========================================================================================================================\n\ndef try_gpu(i=0):\n    if torch.cuda.device_count()>=i+1:\n        return torch.device(f'cuda:{i}')\n    return torch.device('cpu')\n\ndef try_all_gpus():\n    devices=[torch.device(f'cuda:{i}') for i in range(torch.cuda.device_count())]\n    return devices if devices else [torch.device('cpu')]\n\nprint(try_all_gpus())","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2024-01-03T09:33:46.133176Z","iopub.execute_input":"2024-01-03T09:33:46.134379Z","iopub.status.idle":"2024-01-03T09:33:46.146605Z","shell.execute_reply.started":"2024-01-03T09:33:46.134335Z","shell.execute_reply":"2024-01-03T09:33:46.145331Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"### configs\n\n# ===========================================================================================================================\n# ===========================================================================================================================\n\nfc_out=[6,6] # wsi, tma\n\nnet_params=\"/kaggle/input/params-densenet201-224/20231230_uda_extraTma/bestVal9019.params\"\n\n# att_params_1=\"/kaggle/input/params-densenet201-224/20231230_uda_extraTma/bestVal9019_ATT-fc5_prefil_bestVal_0.8818.params\" # wsi\n# att_params_2=\"/kaggle/input/params-densenet201-224/20231230_uda_extraTma/bestVal9019_ATT-fc5_prefil_bestTma_0.9200.params\" # tma\n\natt_params_1=\"/kaggle/input/params-densenet201-224/20231230_uda_extraTma/bestVal9019_ATT-fc6_prefil_bestVal_0.8652.params\" # wsi\natt_params_2=\"/kaggle/input/params-densenet201-224/20231230_uda_extraTma/bestVal9019_ATT-fc6_prefil_bestTma_0.9200.params\" # tma\n\n# ===========================================================================================================================\n# ===========================================================================================================================\n\nTEST_IMG_DIR = \"/kaggle/input/UBC-OCEAN/test_images\"\nTEST_TBNLS_DIR = \"/kaggle/input/UBC-OCEAN/test_thumbnails\"\n\nCONFIG={\n    \"seed\": 42,\n    \"input_size\": 224,\n    \"tma_size_thr\": 5000,\n    \"hda\": True,\n    \"back_bone\": models.densenet201(),\n    \"num_features\": 1920,\n    \"net_params\": net_params,\n    \"att_params\": (att_params_1, att_params_2), # wsi, tma\n    \"self_att\":True,\n    \"fc_out\": fc_out,\n    \"num_labels\": 6,\n    \"label_dict\":{'CC':0, 'EC':1, 'HGSC':2, 'LGSC':3, 'MC':4, 'Other':5},\n    \"label_dict_reverse\":{0: 'CC', 1: 'EC', 2: 'HGSC', 3: 'LGSC', 4: 'MC', 5: 'Other'},\n    \"wsi_scale\":2,\n    \"tma_scale\":4,\n    \"wsi_sample_grid\":224,\n    \"tma_sample_grid\":120,\n    \"sampling_drop_thr\": 0.7,\n    \"sampling_white_thr\": 245,\n    \"max_samples\": 512,\n    \"att_prefilter\": True,\n    \"wsi_bag_conf_thr\": 0.4, # 越大越易变other\n    \"tma_bag_conf_thr\": 0.55, # 越大越易变other\n    \"num_workers\": 2\n}\n\n# ===========================================================================================================================\n# ===========================================================================================================================\n\ndef set_seed(seed=42):\n    '''Sets the seed of the entire notebook so results are the same every time we run.\n    This is for REPRODUCIBILITY.'''\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    # When running on the CuDNN backend, two further options must be set\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n    # Set a fixed value for the hash seed\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    \nset_seed(CONFIG['seed'])","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2024-01-03T09:33:50.743866Z","iopub.execute_input":"2024-01-03T09:33:50.744731Z","iopub.status.idle":"2024-01-03T09:33:51.103497Z","shell.execute_reply.started":"2024-01-03T09:33:50.744679Z","shell.execute_reply":"2024-01-03T09:33:51.102452Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data preparation","metadata":{}},{"cell_type":"code","source":"### dataframe\n\n# ==============================\n# daraframe\n# ==============================\n\ntest_df=pd.read_csv(\"/kaggle/input/UBC-OCEAN/test.csv\")\n\ndef is_tma(row):\n    return (row[\"image_width\"] <= CONFIG[\"tma_size_thr\"]) and (row[\"image_height\"] <= CONFIG[\"tma_size_thr\"])\n\n# 应用函数并创建两个不同的 DataFrame\ntma_df = test_df[test_df.apply(is_tma, axis=1)]\nwsi_df = test_df[~test_df.apply(is_tma, axis=1)]\nprint(len(wsi_df))\nprint(len(tma_df))","metadata":{"execution":{"iopub.status.busy":"2024-01-03T09:34:00.405551Z","iopub.execute_input":"2024-01-03T09:34:00.405946Z","iopub.status.idle":"2024-01-03T09:34:00.417508Z","shell.execute_reply.started":"2024-01-03T09:34:00.405912Z","shell.execute_reply":"2024-01-03T09:34:00.416672Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"### data pipelines: extract tiles from an image\n\n# ==============================\n# transform: color normalization\n# ==============================\n\n# ImageNet\nimg_color_mean=[0.485, 0.456, 0.406]\nimg_color_std=[0.229, 0.224, 0.225]\n\n# use albumentations for augmentation\nval_transform = A.Compose([\n        A.Normalize(mean=img_color_mean, std=img_color_std),\n        ToTensorV2()])\n\n\n# ============================================================\n# extract tiles from given image path\n# ============================================================\n\ndef extract_tiles(img_path, drop_thr=CONFIG[\"sampling_drop_thr\"], \n                  white_thr=CONFIG[\"sampling_white_thr\"], max_samples=CONFIG[\"max_samples\"]):\n\n    # print(f\"processing: {img_path}\")\n    im = pyvips.Image.new_from_file(img_path) #load image\n    w, h = im.width, im.height\n    \n    if w <= CONFIG[\"tma_size_thr\"] and h <= CONFIG[\"tma_size_thr\"]:\n        scale=CONFIG[\"tma_scale\"]\n        grid=CONFIG[\"tma_sample_grid\"]*scale\n    else:\n        scale=CONFIG[\"wsi_scale\"]\n        grid=CONFIG[\"wsi_sample_grid\"]*scale\n        \n    size=CONFIG[\"input_size\"]\n    raw_size=size*scale\n    \n    # scale tma up if needed\n    scale_factor = max(2 if w < 1600 or h < 1600 else 1, raw_size / w, raw_size / h)\n    if scale_factor > 1:\n        im = im.resize(scale_factor)\n        w, h = im.width, im.height\n        \n    idxs = [(y,x) for y in range(0, h - raw_size + 1, grid) for x in range(0, w - raw_size + 1, grid)]\n    \n    if len(idxs)>max_samples:\n        idxs=random.sample(idxs, max_samples)\n    \n    tiles = []\n    for (y,x) in idxs:\n        \n        tile = im.crop(x, y, raw_size, raw_size).write_to_memory()\n        tile = np.frombuffer(tile, dtype=np.uint8).reshape(raw_size, raw_size, -1)[..., :3]\n        tile = np.array(Image.fromarray(tile).resize((size,size), Image.LANCZOS))\n        \n        # emptry ratio detection\n        black_bg = np.sum(tile, axis=2) == 0\n        tile[black_bg, :] = 255\n        mask_bg = np.mean(tile, axis=2) > white_thr\n        if np.sum(mask_bg) >= (np.prod(mask_bg.shape) * drop_thr):\n            continue\n        tiles.append(tile)\n    \n    del im\n    gc.collect()\n    \n    return tiles","metadata":{"execution":{"iopub.status.busy":"2024-01-03T09:34:03.587844Z","iopub.execute_input":"2024-01-03T09:34:03.588485Z","iopub.status.idle":"2024-01-03T09:34:03.602277Z","shell.execute_reply.started":"2024-01-03T09:34:03.588454Z","shell.execute_reply":"2024-01-03T09:34:03.601306Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Models","metadata":{}},{"cell_type":"code","source":"### Model definitions\n\n# ==============================\n# feature extractor\n# ==============================\n\nclass Features(nn.Module):\n    def __init__(self, back_bone=CONFIG[\"back_bone\"]):\n        super(Features, self).__init__()\n        self.features = back_bone.features\n\n    def forward(self, x):\n        features = self.features(x)\n        features = F.relu(features, inplace=True)\n        features = F.adaptive_avg_pool2d(features, (1, 1))\n        features = torch.flatten(features, 1)\n        return features\n\n\n# ==============================\n# base model\n# ==============================\n\nclass BaseModel(nn.Module):\n    def __init__(self,\n                 back_bone=CONFIG[\"back_bone\"],\n                 num_features=CONFIG[\"num_features\"],\n                 num_classes=CONFIG[\"num_labels\"],\n                 hda: bool = False,  # whether use hda head\n                 toalign: bool = False,  # whether use toalign\n                 **kwargs\n                 ):\n        super().__init__()\n        self.num_classes = num_classes\n        self.fdim = num_features\n        self.extractor = Features(back_bone=back_bone)\n        self.fc = nn.Linear(self.fdim, self.num_classes)\n        nn.init.kaiming_normal_(self.fc.weight)\n        if self.fc.bias is not None:\n            nn.init.zeros_(self.fc.bias)\n        \n        # HDA\n        self.hda=hda\n        if self.hda:\n            self.fc0 = nn.Linear(self.fdim, self.num_classes)\n            self.fc1 = nn.Linear(self.fdim, self.num_classes)\n            self.fc2 = nn.Linear(self.fdim, self.num_classes)\n\n        # toalign\n        self.toalign = toalign\n\n    def forward_backbone(self, x):\n        return self.extractor(x)\n\n    def _get_toalign_weight(self, features, labels=None):\n        assert labels is not None, f'labels should be asigned'\n        w = self.fc.weight[labels].detach()  # [B, C]\n        if self.hda:\n            w0 = self.fc0.weight[labels].detach()\n            w1 = self.fc1.weight[labels].detach()\n            w2 = self.fc2.weight[labels].detach()\n            w = w - (w0 + w1 + w2)\n        eng_org = (features**2).sum(dim=1, keepdim=True)  # [B, 1]\n        eng_aft = ((features*w)**2).sum(dim=1, keepdim=True)  # [B, 1]\n        scalar = (eng_org / eng_aft).sqrt()\n        w_pos = w * scalar\n\n        return w_pos\n\n    def forward(self, x, toalign=False, labels=None) -> tuple:\n        \"\"\"\n        return: [f, y, ...]\n        \"\"\"\n        features = self.forward_backbone(x)  # output feature [B, C]\n\n        if toalign:\n            w_pos = self._get_toalign_weight(features, labels=labels)\n            features_pos = features * w_pos\n            y_pos = self.fc(features_pos)\n            if self.hda:\n                z_pos0 = self.fc0(features_pos)\n                z_pos1 = self.fc1(features_pos)\n                z_pos2 = self.fc2(features_pos)\n                z_pos = z_pos0 + z_pos1 + z_pos2\n                return features_pos, y_pos - z_pos, z_pos\n            else:\n                return features_pos, y_pos\n        else:\n            y_hat = self.fc(features)\n            if self.hda:\n                z0 = self.fc0(features)\n                z1 = self.fc1(features)\n                z2 = self.fc2(features)\n                z = z0 + z1 + z2\n                return features, y_hat - z, z\n            else:\n                return features, y_hat\n\n\n# ==============================\n# attention\n# ==============================\n\nclass Attention(nn.Module):\n    def __init__(self, num_features, fc_out=CONFIG[\"fc_out\"], self_att=CONFIG[\"self_att\"]):\n        super(Attention, self).__init__()\n        self.L = num_features\n        self.D = 256\n        self.fc_out=fc_out\n        self.self_att = self_att\n        \n        if self.self_att:\n            self.self_att = SelfAttention(self.L)\n\n        self.attention = nn.Sequential(\n            nn.Linear(self.L, self.D),\n            nn.Tanh(),\n            nn.Linear(self.D, 1)\n        )\n        \n# @@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@     \n        self.fc = nn.Linear(num_features, self.fc_out)\n# @@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@ \n\n    def forward(self, x):\n\n        if self.self_att:\n            x = self.self_att(x)  # BxNxL >> BxNxL\n\n        A = self.attention(x)  # BxNx1\n        A = A.transpose(1, 2)  # Bx1xN\n        A = F.softmax(A, dim=2)  # softmax over N\n        M = torch.bmm(A, x)  # Bx1xL\n        M = M.squeeze(1)# 移除1维，得到BxL\n    \n# @@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@   \n        M = self.fc(M)\n# @@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@ \n\n        return M     \n    \n\nclass SelfAttention(nn.Module):\n    def __init__(self, in_dim):\n        super(SelfAttention, self).__init__()\n        self.query_conv = nn.Conv1d(in_channels=in_dim, out_channels=in_dim // 8, kernel_size=1)\n        self.key_conv = nn.Conv1d(in_channels=in_dim, out_channels=in_dim // 8, kernel_size=1)\n        self.value_conv = nn.Conv1d(in_channels=in_dim, out_channels=in_dim, kernel_size=1)\n        self.gamma = nn.Parameter((torch.zeros(1)).to(try_gpu()))\n        self.softmax = nn.Softmax(dim=-1)\n        self.gamma_att = nn.Parameter((torch.ones(1)).to(try_gpu()))\n    \n    def forward(self, x):\n        # 输入形状：BxNxL\n        B, N, L = x.shape\n        x = x.permute(0, 2, 1)  # BxLxN\n\n        proj_query = self.query_conv(x).view(B, -1, N).permute(0, 2, 1)  # BxNx(L/8)\n        proj_key = self.key_conv(x).view(B, -1, N)  # Bx(L/8)xN\n\n        energy = torch.bmm(proj_query, proj_key)  # BxNxN\n\n        attention = self.softmax(energy)  # BxNxN\n\n        proj_value = self.value_conv(x).view(B, -1, N)  # BxLxN\n        \n        \n        out = torch.bmm(proj_value, attention.permute(0, 2, 1))  # BxLxN\n        out = out.view(B, L, N)\n        out = self.gamma * out + x\n        \n        return out.permute(0, 2, 1)  # BxNxL","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2024-01-03T09:34:08.048929Z","iopub.execute_input":"2024-01-03T09:34:08.049308Z","iopub.status.idle":"2024-01-03T09:34:08.243282Z","shell.execute_reply.started":"2024-01-03T09:34:08.049261Z","shell.execute_reply":"2024-01-03T09:34:08.242226Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"### models instance\n\ndevices = try_all_gpus()\ndevice = devices[0]\n\n# ==============================\n# remove module prefix from params\n# ==============================\n\ndef remove_module_prefix(state_dict):\n    new_state_dict = {}\n    for k, v in state_dict.items():\n        if k.startswith(\"module.\"):\n            # 移除'module.'前缀\n            new_key = k[len(\"module.\"):]\n        else:\n            new_key = k\n        new_state_dict[new_key] = v\n    return new_state_dict\n\n# ==============================\n# nets\n# ==============================\nnet=BaseModel(CONFIG[\"back_bone\"], CONFIG[\"num_features\"], CONFIG[\"num_labels\"], hda=CONFIG[\"hda\"], toalign=True).to(device)\nstate_dict = torch.load(CONFIG[\"net_params\"], map_location=device)\nstate_dict = remove_module_prefix(state_dict)\nnet.load_state_dict(state_dict)\nnet.eval()\n\natt_1 = Attention(net.fdim, fc_out=CONFIG[\"fc_out\"][0], self_att=CONFIG[\"self_att\"]).to(device)\nstate_dict = torch.load(CONFIG[\"att_params\"][0], map_location=device)\nstate_dict = remove_module_prefix(state_dict)\natt_1.load_state_dict(state_dict)\natt_1.eval()\n\natt_2 = Attention(net.fdim, fc_out=CONFIG[\"fc_out\"][1], self_att=CONFIG[\"self_att\"]).to(device)\nstate_dict = torch.load(CONFIG[\"att_params\"][1], map_location=device)\nstate_dict = remove_module_prefix(state_dict)\natt_2.load_state_dict(state_dict)\natt_2.eval()\n\n# ==============================\n# ==============================\n# move all model to right device\n# ==============================\nif len(devices)>1:\n    net = nn.DataParallel(net, device_ids=devices)\n    att_1 = nn.DataParallel(att_1, device_ids=devices)\n    att_2 = nn.DataParallel(att_2, device_ids=devices)","metadata":{"execution":{"iopub.status.busy":"2024-01-03T09:34:14.843581Z","iopub.execute_input":"2024-01-03T09:34:14.844241Z","iopub.status.idle":"2024-01-03T09:34:15.506265Z","shell.execute_reply.started":"2024-01-03T09:34:14.844204Z","shell.execute_reply":"2024-01-03T09:34:15.505422Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# infer helper\n\n# ==============================\n# attention\n# ==============================\ndef infer_row_attention(idx_row, net, attention, device=try_gpu()):\n    \n    row = dict(idx_row[1])\n    \n    # prepare data - cut and load tiles\n    img_path = os.path.join(TEST_IMG_DIR, f\"{str(row['image_id'])}.png\")\n    tiles = extract_tiles(img_path)\n\n    if len(tiles)==0:\n        row['label'] = \"Other\"\n        return row\n    \n    bag_conf_thr = CONFIG[\"tma_bag_conf_thr\"] if is_tma(row) else CONFIG[\"wsi_bag_conf_thr\"]\n    \n    with torch.no_grad():\n        \n        X = torch.stack([val_transform(image=tile)['image'] for tile in tiles]).to(device) # [N,224,224,3]\n        \n# @@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@\n        outputs = net(X)\n        features = outputs[0] # [N,3,224,224]>>[N,1920]\n    \n        if CONFIG[\"att_prefilter\"]:\n            y_hat =outputs[1]\n            max_indices = torch.argmax(y_hat, dim=1)\n            mask = max_indices != y_hat.size(1) - 1\n            features = features[mask]\n            if features.shape[0]==0:\n                row['label'] = \"Other\"\n                del tiles, features\n                gc.collect()\n                return row\n        \n        features = features.unsqueeze(0) # [N, 1920]>>[1, N, 1920]\n        pred = attention(features)\n        pred = F.softmax(pred, dim=1).cpu() # [1, fc_out]\n# @@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@\n\n    max_score = pred.max()\n    if max_score < bag_conf_thr:\n        lb = 5\n    else:\n        lb = pred.argmax(dim=1).item() # item()方法会自动将结果从GPU搬到CPU\n    row['label'] = CONFIG[\"label_dict_reverse\"][lb]\n\n    del tiles, features\n    gc.collect()\n    return row\n\n\n# ==============================\n# infer wrong\n# ==============================\ndef infer_single_image_numb(idx_row):\n    row = dict(idx_row[1])\n    row['label'] = \"AAA\"\n    return row\n\n# ==============================\n# infer other\n# ==============================\ndef infer_single_image_other(idx_row):\n    row = dict(idx_row[1])\n    row['label'] = \"Other\"\n    return row","metadata":{"execution":{"iopub.status.busy":"2024-01-03T09:34:20.180168Z","iopub.execute_input":"2024-01-03T09:34:20.180882Z","iopub.status.idle":"2024-01-03T09:34:20.193108Z","shell.execute_reply.started":"2024-01-03T09:34:20.180848Z","shell.execute_reply":"2024-01-03T09:34:20.192098Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Evaluating","metadata":{}},{"cell_type":"code","source":"# ### fake data for debugging\n\n# # ==============================\n# # ==============================\n\n# TEST_IMG_DIR = \"/kaggle/input/UBC-OCEAN/train_images\"\n# test_df=pd.read_csv(\"/kaggle/input/UBC-OCEAN/train.csv\")\n\n# test_df['area'] = test_df['image_width'] * test_df['image_height']\n\n# # 应用函数并创建两个不同的 DataFrame\n# wsi_df = test_df[~test_df.apply(is_tma, axis=1)].sort_values(by='area', ascending=False)[:2]\n# tma_df = test_df[test_df.apply(is_tma, axis=1)][:5]\n\n# test_df= pd.concat([tma_df, wsi_df], ignore_index=True)\n# # test_df= test_df[[\"image_id\", \"image_width\",\"image_height\"]]\n\n\n# print(wsi_df)\n# print(tma_df)","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2024-01-02T17:31:24.869319Z","iopub.execute_input":"2024-01-02T17:31:24.869854Z","iopub.status.idle":"2024-01-02T17:31:24.939691Z","shell.execute_reply.started":"2024-01-02T17:31:24.86981Z","shell.execute_reply":"2024-01-02T17:31:24.937813Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # ### speed and accuracy testing\n\n# start_time = time.time()\n\n# # ==============================\n\n# preds1 = [infer_row_attention(idx_row, net, att_1, device=try_gpu())\n#     for idx_row in tqdm(wsi_df.iterrows(), total=len(wsi_df))]\n\n# # preds1 = [infer_single_image_numb(idx_row) for idx_row in tqdm(wsi_df.iterrows(), total=len(wsi_df))]\n\n\n# preds2 = Parallel(n_jobs=2, backend='loky')(\n#     delayed(infer_row_attention)\n#     (idx_row, net, att_2, device=try_gpu())\n#     for idx_row in tqdm(tma_df.iterrows(), total=len(tma_df))\n# )\n\n# # preds2 = [infer_single_image_numb(idx_row) for idx_row in tqdm(tma_df.iterrows(), total=len(tma_df))]\n\n\n# output_df = pd.DataFrame(preds1+preds2)[[\"image_id\", \"label\"]]\n\n\n# # ==============================\n\n# print(output_df)\n\n# end_time = time.time()\n\n# output_df = pd.merge(test_df, output_df, on='image_id', suffixes=('_real', '_inf'))\n# accuracy = (output_df['label_real'] == output_df['label_inf']).mean()\n\n# print(\"accuracy: \", accuracy, \"; time: \", end_time-start_time)","metadata":{"execution":{"iopub.status.busy":"2024-01-02T17:31:33.219397Z","iopub.execute_input":"2024-01-02T17:31:33.219845Z","iopub.status.idle":"2024-01-02T17:40:35.296974Z","shell.execute_reply.started":"2024-01-02T17:31:33.21981Z","shell.execute_reply":"2024-01-02T17:40:35.295902Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Final Infer","metadata":{}},{"cell_type":"code","source":"# infer\n\n# ==============================\n\npreds1 = [infer_row_attention(idx_row, net, att_1, device=try_gpu())\n    for idx_row in tqdm(wsi_df.iterrows(), total=len(wsi_df))]\n\n# preds1 = [infer_single_image_numb(idx_row) for idx_row in tqdm(wsi_df.iterrows(), total=len(wsi_df))]\n\n\npreds2 = Parallel(n_jobs=2, backend='loky')(\n    delayed(infer_row_attention)\n    (idx_row, net, att_2, device=try_gpu())\n    for idx_row in tqdm(tma_df.iterrows(), total=len(tma_df))\n)\n\n# preds2 = [infer_single_image_numb(idx_row) for idx_row in tqdm(tma_df.iterrows(), total=len(tma_df))]\n\n\noutput_df = pd.DataFrame(preds1+preds2)[[\"image_id\", \"label\"]]\n\n# ==============================\n\nprint(output_df)","metadata":{"execution":{"iopub.status.busy":"2023-12-31T03:44:31.150276Z","iopub.execute_input":"2023-12-31T03:44:31.151399Z","iopub.status.idle":"2023-12-31T03:45:24.181951Z","shell.execute_reply.started":"2023-12-31T03:44:31.15136Z","shell.execute_reply":"2023-12-31T03:45:24.180652Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"output_df.to_csv(\"submission.csv\", index=False)\nprint(\"Your submission was successfully saved!\")","metadata":{"execution":{"iopub.status.busy":"2023-12-29T17:27:37.530874Z","iopub.execute_input":"2023-12-29T17:27:37.531978Z","iopub.status.idle":"2023-12-29T17:27:37.542031Z","shell.execute_reply.started":"2023-12-29T17:27:37.531928Z","shell.execute_reply":"2023-12-29T17:27:37.541083Z"},"trusted":true},"execution_count":null,"outputs":[]}]}