{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":91249,"databundleVersionId":11294684,"sourceType":"competition"},{"sourceId":9862305,"sourceType":"datasetVersion","datasetId":6052780},{"sourceId":10482787,"sourceType":"datasetVersion","datasetId":6490555},{"sourceId":10482884,"sourceType":"datasetVersion","datasetId":6490567},{"sourceId":10588947,"sourceType":"datasetVersion","datasetId":6553343},{"sourceId":11750653,"sourceType":"datasetVersion","datasetId":6847111},{"sourceId":11965964,"sourceType":"datasetVersion","datasetId":6968147},{"sourceId":234929223,"sourceType":"kernelVersion"}],"dockerImageVersionId":30805,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install git+https://github.com/kai-coder/lossmap.git","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T21:30:53.676677Z","iopub.execute_input":"2025-05-26T21:30:53.677114Z","iopub.status.idle":"2025-05-26T21:31:13.020799Z","shell.execute_reply.started":"2025-05-26T21:30:53.677054Z","shell.execute_reply":"2025-05-26T21:31:13.019689Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom matplotlib import animation, rc\nimport cv2\nimport json\nfrom numpy.typing import NDArray\nfrom scipy import ndimage\nimport seaborn as sns\nimport scipy\nfrom torch.optim import Optimizer\nfrom functools import partial\nfrom torch.utils.data import Dataset\nfrom skimage import measure\nfrom tqdm import tqdm\nfrom PIL import Image\nimport plotly.graph_objs as go\nfrom scipy.fft import ifftn, fftn\nfrom PIL import Image\nimport os\nfrom scipy import misc\nimport gc\nimport torch\nimport warnings\nfrom multiprocessing import Pool\nimport threading\nfrom torch import nn\nimport lossmap\nimport pickle\n\nwarnings.simplefilter(action='ignore', category=FutureWarning)\nrc('animation', html='jshtml')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-05-26T21:31:13.023623Z","iopub.execute_input":"2025-05-26T21:31:13.023928Z","iopub.status.idle":"2025-05-26T21:31:17.622201Z","shell.execute_reply.started":"2025-05-26T21:31:13.023901Z","shell.execute_reply":"2025-05-26T21:31:17.621286Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class DatasetAttr(Dataset):\n    def __init__(self, batch_size: int, shuffle: bool) -> None:\n        self.batch_size = batch_size\n        self.shuffle = shuffle\n\n    def on_epoch_end(self) -> None:\n        print(\"user must define on_epoch_end\")\n\n    def __len__(self) -> int:\n        print(\"user must define __len__\")\n        return -1\n\nclass DatasetModel(DatasetAttr):\n    def __init__(self, x: np.ndarray, y: np.ndarray, batch_size: int, shuffle: bool) -> None:\n        super().__init__(batch_size, shuffle)\n        self.x = x\n        self.y = y\n\n    def __len__(self) -> int:\n        return len(self.x) // self.batch_size\n\n    def on_epoch_end(self, must_shuffle: bool=False) -> None:\n        index_arr = np.arange(len(self.x))\n\n        if self.shuffle or must_shuffle:\n            np.random.shuffle(index_arr)\n\n        self.x = self.x[index_arr]\n        self.y = self.y[index_arr]\n\n    def __getitem__(self, idx: int) -> tuple[torch.Tensor, torch.Tensor]:\n        if idx >= self.__len__():\n            raise IndexError()\n\n        next_idx = (idx + 1) * self.batch_size\n\n        x = torch.tensor(self.x[idx:next_idx],  dtype=torch.float32)\n        y = torch.tensor(self.y[idx:next_idx, None],  dtype=torch.float32)\n\n        return x, y\n\ndef put_augment_image(idx: int, arr: NDArray[np.int_], data: NDArray[np.int_], patch_size: int\n                   , coordinate: NDArray[np.int_], shuffle: bool, random_array: NDArray[np.float_]) -> None:\n    low_coord = coordinate - patch_size // 2\n    high_coord = coordinate + patch_size//2\n    arr[idx, 0] = data[low_coord[0]:high_coord[0], low_coord[1]:high_coord[1], low_coord[2]:high_coord[2]]\n    if shuffle:\n        if random_array[idx][0] < 0.5:\n            arr[idx] = arr[idx].transpose(0, 1, 3, 2)\n        if random_array[idx][1] < 0.5:\n            arr[idx] = np.flip(arr[idx], 1)\n        if random_array[idx][2] < 0.5:\n            arr[idx] = np.flip(arr[idx], 2)\n        if random_array[idx][3] < 0.5:\n            arr[idx] = np.flip(arr[idx], 3)\n\n\nclass ImageLoader(DatasetModel):\n    def __init__(self, x: np.ndarray, y: np.ndarray, batch_size: int, shuffle: bool, patch_size: int\n                 , full_data_amount: int, radius_var: float, radius_thresh: float) -> None:\n        super().__init__(x, y, batch_size, shuffle)\n        self.patch_size = patch_size\n        self.full_data_amount = full_data_amount\n        self.shapes = np.array([i.shape for i in self.x])\n        self.radius_var = radius_var\n        self.radius_thresh = radius_thresh\n        self.chosen_idx = None\n        self.chosen_coordinates = None\n        self.on_epoch_end(True)\n\n    def __len__(self) -> int:\n        if self.shuffle:\n            return 48 * 3\n        else:\n            return 48 * 4\n\n    def on_epoch_end(self, must_shuffle: bool = False) -> None:\n        if self.shuffle or must_shuffle:\n            patch_amount = self.batch_size * self.__len__()\n\n            full_patch_amount = int(patch_amount * 0.5)\n            full_data_idx = np.arange(self.full_data_amount)\n\n            chosen_full_idx = np.random.choice(full_data_idx, full_patch_amount)\n\n            low_idx = (self.patch_size // 2, self.patch_size // 2, self.patch_size // 2)\n            high_idx = self.shapes[chosen_full_idx] - np.array(low_idx)\n\n            chosen_full_coordinates = np.random.randint(low_idx, high_idx, (full_patch_amount, 3))\n\n            flagella_idx = np.arange(self.full_data_amount, len(self.x))\n\n            flagella_patch_amount = patch_amount - full_patch_amount\n            chosen_flagella_idx = np.random.choice(flagella_idx, flagella_patch_amount)\n\n            self.chosen_idx = np.append(chosen_full_idx, chosen_flagella_idx).astype(int)\n\n            high_idx = self.shapes[chosen_flagella_idx] - np.array(low_idx)\n\n            z_angle = np.pi * np.random.uniform(size=flagella_patch_amount)[:, None]\n            x_y_angle = 2 * np.pi * np.random.uniform(size=flagella_patch_amount)[:, None]\n\n            unit_z = np.cos(z_angle)\n            unit_y = np.sin(z_angle) * np.sin(x_y_angle)\n            unit_x = np.sin(z_angle) * np.cos(x_y_angle)\n            unit_vector = np.concatenate((unit_z, unit_y, unit_x), axis=1)\n\n            radius = np.random.uniform(size=flagella_patch_amount)[:, None] * (self.patch_size // 2 * self.radius_var)\n\n            vector = unit_vector * radius\n\n            center_coordinate = self.y[chosen_flagella_idx] + vector\n\n            chosen_flagella_coordinate = np.clip(center_coordinate, low_idx, high_idx)\n\n            self.chosen_coordinates = np.concatenate((chosen_full_coordinates, chosen_flagella_coordinate)).astype(int)\n\n            index_arr = np.arange(len(self.chosen_idx))\n            np.random.shuffle(index_arr)\n\n            self.chosen_idx = self.chosen_idx[index_arr]\n            self.chosen_coordinates = self.chosen_coordinates[index_arr]\n\n\n    def getBatch(self, idx):\n        if idx >= self.__len__():\n            raise IndexError()\n\n        first_idx = idx * self.batch_size\n        next_idx = (idx + 1) * self.batch_size\n\n        return self.chosen_idx[first_idx:next_idx], self.chosen_coordinates[first_idx:next_idx]\n\n    def __getitem__(self, idx):\n        batch = self.getBatch(idx)\n\n        x = np.zeros((self.batch_size, 1, self.patch_size, self.patch_size, self.patch_size), dtype=np.uint8)\n\n        thread_array = []\n\n        random_array = np.random.rand(self.batch_size, 4)\n        for idx, (imgIdx, coord) in enumerate(zip(batch[0], batch[1])):\n            args = (idx, x, self.x[imgIdx], self.patch_size, coord, self.shuffle, random_array)\n            process = threading.Thread(target=put_augment_image, args=args)\n            thread_array.append(process)\n\n            process.start()\n\n        for process in thread_array:\n            process.join()\n\n        x = torch.tensor(x, dtype=torch.float32).to(DEVICE)\n\n        lower = torch.quantile(x.flatten(1), 0.02, dim=1)[:, None, None, None, None]\n        upper = torch.quantile(x.flatten(1), 0.8, dim=1)[:, None, None, None, None]\n        x = torch.clip(x, lower, upper)\n        \n        x = (x - lower) / (upper - lower)\n        \n\n        contains_flagella = (self.y[batch[0]] > 0).all(axis=1)\n\n        squared_sum = np.sum(np.square(self.y[batch[0]] - batch[1]), axis=1)\n        distances = np.sqrt(squared_sum)\n        within_range = distances < self.patch_size // 2 * self.radius_thresh\n\n        y = contains_flagella & within_range\n\n        y = torch.tensor(y[:, None], dtype=torch.float32).to(DEVICE)\n        \n        return x, y","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T21:31:17.623361Z","iopub.execute_input":"2025-05-26T21:31:17.623747Z","iopub.status.idle":"2025-05-26T21:31:17.647082Z","shell.execute_reply.started":"2025-05-26T21:31:17.623722Z","shell.execute_reply":"2025-05-26T21:31:17.646355Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class TrainModelClass:\n    def __init__(self, optimizer: Optimizer, loss_fn: nn.Module) -> None:\n        self.optimizer = optimizer\n        self.loss_fn = loss_fn\n\n    def run_epoch(self, dataset: DatasetAttr, train: bool, model: nn.Module) -> float:\n        total_loss = 0\n        for data_point in dataset:\n            x, y = data_point\n\n            if train:\n                model.train()\n            else:\n                model.eval()\n\n            with torch.set_grad_enabled(train):\n                out = model(x)\n                loss = self.loss_fn(out, y)\n\n            if train:\n                self.optimizer.zero_grad()\n                loss.backward()\n                self.optimizer.step()\n\n            total_loss += loss.item()\n\n        return total_loss / dataset.__len__()\n\n    def train(self, epochs: int, train_data: DatasetAttr, test_data: DatasetAttr, model: nn.Module) -> None:\n        for epoch_num in range(epochs):\n            train_loss = self.run_epoch(train_data, True, model)\n            test_loss = self.run_epoch(test_data, False, model)\n\n            print(\"Epoch: {0}; Train Loss: {1:.3f}; Test Loss: {2:.3f}\".format(epoch_num, train_loss, test_loss))\n\n            train_data.on_epoch_end()\n            test_data.on_epoch_end()\n\n\nclass FancyTrainModelClass(TrainModelClass):\n    def __init__(self, optimizer: Optimizer, loss_fn: nn.Module, scheduler: nn.Module) -> None:\n        super().__init__(optimizer, loss_fn)\n        self.scaler = torch.GradScaler(\"cuda\")\n        self.scheduler = scheduler\n\n    def run_epoch(self, dataset: DatasetAttr, train: bool, model: nn.Module) -> float:\n        total_loss = 0\n        for data_point in dataset:\n            x, y = data_point\n\n            if train:\n                model.train()\n            else:\n                model.eval()\n\n            with torch.set_grad_enabled(train):\n                with torch.amp.autocast('cuda'):\n                    out = model(x)\n                    loss = self.loss_fn(out, y)\n\n            if train:\n                self.optimizer.zero_grad()\n                self.scaler.scale(loss).backward()\n                self.scaler.step(self.optimizer)\n                self.scaler.update()\n\n            total_loss += loss.item()\n\n        return total_loss / dataset.__len__()\n\n    def train(self, epochs: int, train_data: DatasetAttr, test_data: DatasetAttr, model: nn.Module)\\\n            -> tuple[list[float], list[float]]:\n        train_loss_arr = []\n        test_loss_arr = []\n        for epoch_num in range(epochs):\n            train_loss = self.run_epoch(train_data, True, model)\n            test_loss = self.run_epoch(test_data, False, model)\n\n            print(\"Epoch: {0}; Learning Rate: {1:.5f}\".format(epoch_num, self.scheduler.get_last_lr()[0]))\n            print(\"Train Loss: {0:.3f}; Test Loss: {1:.3f}\".format(train_loss, test_loss))\n            print()\n\n            train_loss_arr.append(train_loss)\n            test_loss_arr.append(test_loss)\n\n            self.scheduler.step()\n            train_data.on_epoch_end()\n            test_data.on_epoch_end()\n\n        return train_loss_arr, test_loss_arr","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T21:31:17.648250Z","iopub.execute_input":"2025-05-26T21:31:17.648525Z","iopub.status.idle":"2025-05-26T21:31:17.662226Z","shell.execute_reply.started":"2025-05-26T21:31:17.648499Z","shell.execute_reply":"2025-05-26T21:31:17.661368Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"TRAIN_PATH = \"/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/train/\"\n\ntrainDf = pd.read_csv(\"/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/train_labels.csv\", index_col='row_id')\narrShapeCols = ['Array shape (axis 0)', 'Array shape (axis 1)', 'Array shape (axis 2)']\ntrainDf['imgSize'] = trainDf[arrShapeCols].apply(lambda x:np.array([x.iloc[i] for i in range(3)]),axis=1)\ntrainDf = trainDf.drop(arrShapeCols, axis=1)\narrShapeCols = ['tomo_id', 'Motor axis 0', 'Motor axis 1', 'Motor axis 2', \"Voxel spacing\", \"Number of motors\"]\ntrainDf = trainDf.rename({k:v for k, v in zip(arrShapeCols, ['id', 'z', 'y', 'x', 'space', 'num'])}, axis=1)\ntrainDf = trainDf[trainDf.num <= 1]\ntrainDf = trainDf[trainDf.space>10]\ntrainDf[['z', 'x', 'y', 'num']] = trainDf[['z', 'x', 'y', 'num']].astype(int)\n\nseed = 0\nwhile True:\n    np.random.seed(seed)\n    \n    ps = 1 / np.e**((np.stack(trainDf.imgSize).prod(axis=1)/10**9))[np.stack(trainDf.imgSize).prod(axis=1).argsort()]\n    arr = np.stack(np.random.choice(trainDf.id, 50, p=np.e**ps/np.sum(np.e**ps), replace=False))\n\n    size = np.sum(np.prod(np.stack(trainDf.set_index(\"id\").loc[arr].imgSize.values), axis=1)) / 10**9\n    if size > 17 and size < 18:\n        print(size, seed)\n        break\n    seed += 1\n\ntrainDf2 = trainDf[trainDf.num == 1]\ntrainDf2 = trainDf2[~trainDf2.id.isin(arr)]\n\narr2 = list(trainDf2.id.values)\narr = list(arr) + arr2\n\ncoords = trainDf.set_index(\"id\").loc[arr2][['z', 'y', 'x']].values\nminPlace = np.clip(coords - 96, np.repeat(0, len(coords) * 3).reshape(-1, 3), np.stack(trainDf.set_index(\"id\").loc[arr2].imgSize.values)) \nmaxPlace = np.clip(coords + 96, np.repeat(0, len(coords) * 3).reshape(-1, 3), np.stack(trainDf.set_index(\"id\").loc[arr2].imgSize.values)) \ncenters = coords - minPlace","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T21:31:17.663235Z","iopub.execute_input":"2025-05-26T21:31:17.663486Z","iopub.status.idle":"2025-05-26T21:31:17.733818Z","shell.execute_reply.started":"2025-05-26T21:31:17.663463Z","shell.execute_reply":"2025-05-26T21:31:17.733022Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def readAllFiles(filename):\n    with open(filename, \"rb\") as file:\n        while True:\n            try:\n                yield pickle.load(file)\n            except:\n                break","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T21:31:17.734802Z","iopub.execute_input":"2025-05-26T21:31:17.735053Z","iopub.status.idle":"2025-05-26T21:31:17.739176Z","shell.execute_reply.started":"2025-05-26T21:31:17.735030Z","shell.execute_reply":"2025-05-26T21:31:17.738454Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"imgArr = list(readAllFiles(\"/kaggle/input/byu-dataset-creator/BYUData.pkl\"))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T21:31:17.741876Z","iopub.execute_input":"2025-05-26T21:31:17.742120Z","iopub.status.idle":"2025-05-26T21:32:45.671608Z","shell.execute_reply.started":"2025-05-26T21:31:17.742097Z","shell.execute_reply":"2025-05-26T21:32:45.670740Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else 'cpu')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T21:32:45.672689Z","iopub.execute_input":"2025-05-26T21:32:45.672979Z","iopub.status.idle":"2025-05-26T21:32:45.757486Z","shell.execute_reply.started":"2025-05-26T21:32:45.672952Z","shell.execute_reply":"2025-05-26T21:32:45.756575Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class resConv(torch.nn.Module):\n    def __init__(self, inChannel, newAmnt, stride, identity=False, ReLU=True, small=True, attention=True):\n        super().__init__()\n\n\n        if ReLU:\n            self.conv1 = nn.Sequential(nn.Conv3d(inChannel, newAmnt, 3, stride=stride, padding=1),\n                                       nn.InstanceNorm3d(newAmnt)\n                                            )\n        else:\n            self.conv1 = nn.Sequential(nn.Conv3d(inChannel, newAmnt, 3, stride=stride, padding=1)\n                                            )\n        if attention:\n            self.CBAM = CBAM(newAmnt)\n        else:\n            self.CBAM = nn.Identity()\n\n        self.ReLU = ReLU\n            \n        if identity:\n            self.conv2 = nn.Identity()\n        else:\n            if small:\n                self.conv2 = nn.Conv3d(inChannel, newAmnt, 3, stride=stride, padding=1)\n            else:\n                self.conv2 = nn.Conv3d(inChannel, newAmnt, 1, stride=1)\n\n        self.relu = nn.ReLU()\n    \n    def forward(self, x):\n\n        \n        x2 = self.conv1(x)\n\n        x2 = self.CBAM(x2)\n\n        x3 = x2 + self.conv2(x)\n\n        if self.ReLU:\n            return self.relu(x3)\n        else:\n            return x3","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T21:32:45.758796Z","iopub.execute_input":"2025-05-26T21:32:45.759169Z","iopub.status.idle":"2025-05-26T21:32:45.768101Z","shell.execute_reply.started":"2025-05-26T21:32:45.759128Z","shell.execute_reply":"2025-05-26T21:32:45.767393Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CBAM(nn.Module):\n    def __init__(self, channels):\n        super().__init__()\n        self.channels = channels\n        inNeurons = channels\n        outNeurons = int(channels)\n\n        self.neuralNet = nn.Sequential(nn.Flatten(),\n                                      nn.Linear(inNeurons, outNeurons),\n                                      #nn.BatchNorm1d(outNeurons),\n                                      #nn.GELU(),\n                                      nn.Linear(outNeurons, inNeurons),\n                                     )\n\n        self.convoltution = nn.Conv3d(2, 1, 5, padding='same')\n\n        self.sigmoid = nn.Sigmoid()\n        self.sigmoid2 = nn.Sigmoid()\n\n    def forward(self, x):\n        m = torch.amax(x, dim=(2, 3, 4), keepdim=True)\n        m = self.neuralNet(m)\n        m = m.view(x.size(0), self.channels, 1, 1, 1)\n        \n        a = torch.mean(x, dim=(2, 3, 4), keepdim=True)\n        a = self.neuralNet(a)\n        a = a.view(x.size(0), self.channels, 1, 1, 1)\n        \n\n        a = self.sigmoid(m + a)\n\n        x = x * a\n\n        m = torch.max(x, dim=1, keepdim=True)\n        m = m[0]\n        \n        a = torch.mean(x, dim=1, keepdim=True)\n\n        a = self.convoltution(torch.concat((m, a), axis=1))\n\n        a = self.sigmoid2(a)\n\n        x = x * a\n        \n        return x\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T21:32:45.769104Z","iopub.execute_input":"2025-05-26T21:32:45.769409Z","iopub.status.idle":"2025-05-26T21:32:45.778738Z","shell.execute_reply.started":"2025-05-26T21:32:45.769374Z","shell.execute_reply":"2025-05-26T21:32:45.778058Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class attentionLayer(torch.nn.Module):\n    def __init__(self, inChannel, newAmnt):\n        super().__init__()\n            \n        self.query = nn.Conv3d(inChannel, newAmnt, 1, padding='same')\n        self.key = nn.Conv3d(inChannel, newAmnt, 1, padding='same')\n        self.value = nn.Conv3d(inChannel, inChannel, 1, padding='same')\n        self.soft = torch.nn.Softmax(dim=1)\n        self.p1 = nn.Parameter(torch.tensor(0.0))\n    \n    def forward(self, x):\n        q = self.query(x).flatten(2)\n        k = self.key(x).flatten(2).transpose(1, 2)\n        \n        similarityScores = torch.bmm(k, q)\n        mag = torch.sqrt((k**2).sum(axis=2)).unsqueeze(2)\n        \n        similarityScores /= mag\n        similarityScores = self.soft(similarityScores)\n        \n        v = self.value(x).flatten(2)\n        similarityScores = torch.bmm(v, similarityScores)\n        similarityScores = torch.reshape(similarityScores, x.shape)\n        \n        return similarityScores * self.p1 + x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T21:32:45.779810Z","iopub.execute_input":"2025-05-26T21:32:45.780471Z","iopub.status.idle":"2025-05-26T21:32:45.792777Z","shell.execute_reply.started":"2025-05-26T21:32:45.780432Z","shell.execute_reply":"2025-05-26T21:32:45.792114Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class attentionSkipLayer(torch.nn.Module):\n    def __init__(self, skipChannels, upChannels, stride=1):\n        super().__init__()\n            \n        self.upChannels = nn.Conv3d(upChannels, skipChannels, 1)\n        \n        self.skipChannels = nn.Conv3d(skipChannels, skipChannels, 1, stride=stride)\n        \n        self.weightChannel = nn.Conv3d(skipChannels, 1, 1)\n\n        if stride==1:\n            self.up = nn.Identity()\n        else:\n            self.up = nn.Upsample(scale_factor=2, mode='trilinear')\n        \n        self.gelu = nn.ReLU()\n        self.sigmoid = nn.Sigmoid()\n    \n    def forward(self, x, x2):\n        x2 = self.upChannels(x2)\n        \n        x3 = self.skipChannels(x)\n\n        x3 = x2 + x3\n\n        x3 = self.gelu(x3)\n\n        x3 = self.weightChannel(x3)\n        x3 = self.sigmoid(x3)\n\n        x3 = self.up(x3)\n        \n        return x * x3","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T21:32:45.793627Z","iopub.execute_input":"2025-05-26T21:32:45.793890Z","iopub.status.idle":"2025-05-26T21:32:45.804980Z","shell.execute_reply.started":"2025-05-26T21:32:45.793868Z","shell.execute_reply":"2025-05-26T21:32:45.804194Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class transConv(torch.nn.Module):\n    def __init__(self, inChannel, newAmnt, stride):\n        super().__init__()\n\n        if stride == 1:\n            outPad = 0\n        else:\n            outPad = 1\n            \n        self.conv1 = nn.Sequential(nn.ConvTranspose3d(inChannel, newAmnt, 3, stride=stride, padding=1, output_padding=outPad),\n                                         nn.InstanceNorm3d(newAmnt),\n                                         nn.GELU()\n                                        )\n    \n    def forward(self, x):\n\n        \n        x2 = self.conv1(x)\n        \n        return x2","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T21:32:45.805840Z","iopub.execute_input":"2025-05-26T21:32:45.806109Z","iopub.status.idle":"2025-05-26T21:32:45.814671Z","shell.execute_reply.started":"2025-05-26T21:32:45.806086Z","shell.execute_reply":"2025-05-26T21:32:45.813862Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class upAndNext(torch.nn.Module):\n    def __init__(self, skipChannels, upChannels, nextChannels):\n        super().__init__()\n            \n        self.upChannels = transConv(upChannels + skipChannels, nextChannels, 2)\n        self.skipChannels = resConv(nextChannels, nextChannels, 1, identity=True, small=False)\n    \n    def forward(self, x, x2):\n        x2 = torch.concat((x, x2), axis=1)\n        x2 = self.upChannels(x2)\n        \n        x2 = self.skipChannels(x2)\n        \n        return x2","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T21:32:45.815500Z","iopub.execute_input":"2025-05-26T21:32:45.815785Z","iopub.status.idle":"2025-05-26T21:32:45.825215Z","shell.execute_reply.started":"2025-05-26T21:32:45.815760Z","shell.execute_reply":"2025-05-26T21:32:45.824638Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class upAndNext(torch.nn.Module):\n    def __init__(self, skipChannels, upChannels, nextChannels):\n        super().__init__()\n            \n        self.upChannels = transConv(upChannels + skipChannels, nextChannels, 2)\n        self.skipChannels = resConv(nextChannels, nextChannels, 1, identity=True, small=False)\n    \n    def forward(self, x, x2):\n        x2 = torch.concat((x, x2), axis=1)\n        x2 = self.upChannels(x2)\n        \n        x2 = self.skipChannels(x2)\n        \n        return x2","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T21:32:45.826049Z","iopub.execute_input":"2025-05-26T21:32:45.826277Z","iopub.status.idle":"2025-05-26T21:32:45.834076Z","shell.execute_reply.started":"2025-05-26T21:32:45.826255Z","shell.execute_reply":"2025-05-26T21:32:45.833265Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class scoreGuesser(torch.nn.Module):\n    def __init__(self):\n        super().__init__()\n\n        size = 5\n        sigma = 1\n        \n        ax = np.linspace(-(size // 2), size // 2, size)\n        xx, yy, zz = np.meshgrid(ax, ax, ax, indexing='ij')\n    \n        kernel = np.exp(-(xx**2 + yy**2 + zz**2) / (2. * sigma**2))\n    \n        kernel /= np.sum(kernel)\n\n        kernel = torch.tensor(kernel, dtype=torch.float32).to(DEVICE)\n\n        self.kernel = nn.Parameter(kernel.view(1, 1, size, size, size), requires_grad=True).to(DEVICE)\n        \n\n        d = np.array([32, 64, 128, 256])\n\n        d = [int(i) for i in d]\n\n        self.Conv1 = resConv(1, d[0], 2, identity=False,attention=False)\n\n        self.Conv2 = resConv(d[0], d[1], 2, identity=False,attention=True)\n\n        self.Conv3 = resConv(d[1], d[2], 2, identity=False, attention=True)\n\n        self.Conv4 = nn.Sequential(resConv(d[2], d[3], 2, identity=False, attention=True), \n                                  resConv(d[3], d[3], 1, identity=True, small=False,attention=True))\n        \n        #self.Conv5 = resConv(d[3], d[4], 2, identity=False, attention=True)\n\n        #self.Conv6 = resConv(d[4], d[4], 1, identity=True, small=False,attention=True)\n\n        self.pool = nn.Sequential(nn.AdaptiveAvgPool3d(1),\n                                  nn.Flatten(),\n                                  nn.Linear(d[-1], 128),\n                                  #nn.Dropout(0.2),\n                                  nn.InstanceNorm1d(128),\n                                  nn.ReLU(),\n                                  nn.Linear(128, 32),\n                                  #nn.Dropout(0.2),\n                                  nn.InstanceNorm1d(32),\n                                  nn.ReLU(),\n                                  nn.Linear(32, 16),\n                                  #nn.Dropout(0.05),\n                                  nn.InstanceNorm1d(16),\n                                \n                                  nn.Linear(16, 1))\n\n\n    def forward(self, x):\n        x = nn.functional.conv3d(x, self.kernel, padding='same')\n        x = self.Conv1(x)\n        \n        x = self.Conv2(x)\n        \n        x = self.Conv3(x)\n\n        x = self.Conv4(x)\n        #x = self.Conv5(x)\n        #x = self.Conv6(x)\n\n        x = self.pool(x)\n        \n    \n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T21:32:45.835083Z","iopub.execute_input":"2025-05-26T21:32:45.835410Z","iopub.status.idle":"2025-05-26T21:32:45.845388Z","shell.execute_reply.started":"2025-05-26T21:32:45.835376Z","shell.execute_reply":"2025-05-26T21:32:45.844410Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"trainX = imgArr[:37] + imgArr[50:50 + int(len(imgArr[50:]) * 0.75)]\ntrainY = trainDf.set_index(\"id\").loc[arr[:37] + arr[50:50 + int(len(imgArr[50:]) * 0.75)]][['z','y','x']].values\ntrainY[37:37 + int(len(imgArr[50:]) * 0.75)] = centers[:int(len(imgArr[50:]) * 0.75)]\n\ntestX = imgArr[37:50] + imgArr[50 + int(len(imgArr[50:]) * 0.75):]\ntestY = trainDf.set_index(\"id\").loc[arr[37:50] + arr[50 + int(len(imgArr[50:]) * 0.75):]][['z','y','x']].values\ntestY[13:] = centers[int(len(imgArr[50:]) * 0.75):]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T21:32:45.846100Z","iopub.execute_input":"2025-05-26T21:32:45.846395Z","iopub.status.idle":"2025-05-26T21:32:45.860160Z","shell.execute_reply.started":"2025-05-26T21:32:45.846337Z","shell.execute_reply":"2025-05-26T21:32:45.859541Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = scoreGuesser()\n\n\nmodel = nn.DataParallel(model)\n\nmodel.load_state_dict(torch.load('/kaggle/input/byu-epochs/epoch_21 (1).pt', map_location=torch.device('cpu')))\n\nmodel = model.to(DEVICE)\n\n\ngc.collect()\ntorch.cuda.empty_cache()\n\n#lossFunct = nn.BCELoss(reduction='mean')\nloss_fn = nn.BCEWithLogitsLoss()\n\noptimizer = torch.optim.Adam(model.parameters(), lr=2e-4)\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, 25, 1e-5, verbose=True)\n\ngc.collect()\ntorch.cuda.empty_cache()\n\n\ntrainData = ImageLoader(trainX, trainY, 32, True, 96, 37, 0.8, 0.8)\ntestData = ImageLoader(testX, testY, 32, False, 96, 13, 0.8, 0.8)\n\ngc.collect()\ntorch.cuda.empty_cache()\n\ntrain_model_class = FancyTrainModelClass(optimizer, loss_fn, scheduler)\n\n#trainLossArr, testLossArr = train_model_class.train(25, trainData, testData, model)\nmodel.eval()\ntorch.cuda.empty_cache()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T21:32:45.861218Z","iopub.execute_input":"2025-05-26T21:32:45.861540Z","iopub.status.idle":"2025-05-26T21:32:47.676710Z","shell.execute_reply.started":"2025-05-26T21:32:45.861504Z","shell.execute_reply":"2025-05-26T21:32:47.675805Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"loss_map = lossmap.LossMap(model, DEVICE)\nx, y, loss = loss_map.get_loss_landscape(-100, 100, 100, partial(train_model_class.run_epoch, trainData, False))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T21:39:21.656703Z","iopub.execute_input":"2025-05-26T21:39:21.657528Z","iopub.status.idle":"2025-05-26T21:40:49.455080Z","shell.execute_reply.started":"2025-05-26T21:39:21.657492Z","shell.execute_reply":"2025-05-26T21:40:49.453745Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"fig, ax = plt.subplots(subplot_kw={\"projection\": \"3d\"})\nax.plot_wireframe(x, y, loss, color='C0')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T21:40:49.455844Z","iopub.status.idle":"2025-05-26T21:40:49.456145Z","shell.execute_reply.started":"2025-05-26T21:40:49.456005Z","shell.execute_reply":"2025-05-26T21:40:49.456020Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pickle\n\nwith open('data.pkl', 'wb') as file:\n    pickle.dump((x, y, loss), file)","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}