{"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","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**<h1>in this notebook I added the inception module + squeezenet as 3d convolutional network.\nI used an existing notebook and changed the architecture of the model.</h1>**","metadata":{}},{"cell_type":"code","source":"import os\nimport glob\nfrom tqdm import tqdm_notebook as tqdm\nimport random\nimport numpy as np\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.nn.functional as F\nfrom torchvision import transforms, utils\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport matplotlib\nimport pydicom\nfrom pydicom.pixel_data_handlers.util import apply_voi_lut\nimport cv2\n\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","metadata":{"execution":{"iopub.status.busy":"2021-08-27T07:41:51.311753Z","iopub.execute_input":"2021-08-27T07:41:51.312547Z","iopub.status.idle":"2021-08-27T07:41:53.051948Z","shell.execute_reply.started":"2021-08-27T07:41:51.31242Z","shell.execute_reply":"2021-08-27T07:41:53.051064Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from IPython.display import clear_output\n!pip install imutils\nclear_output()","metadata":{"execution":{"iopub.status.busy":"2021-08-27T07:42:01.221115Z","iopub.execute_input":"2021-08-27T07:42:01.221544Z","iopub.status.idle":"2021-08-27T07:42:11.944692Z","shell.execute_reply.started":"2021-08-27T07:42:01.221498Z","shell.execute_reply":"2021-08-27T07:42:11.94322Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import imutils","metadata":{"execution":{"iopub.status.busy":"2021-08-27T07:42:16.241446Z","iopub.execute_input":"2021-08-27T07:42:16.241773Z","iopub.status.idle":"2021-08-27T07:42:16.251686Z","shell.execute_reply.started":"2021-08-27T07:42:16.241744Z","shell.execute_reply":"2021-08-27T07:42:16.250681Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def set_seed(seed):\n    random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)\n        torch.backends.cudnn.deterministic = True\n\n\nset_seed(42)","metadata":{"execution":{"iopub.status.busy":"2021-08-27T07:42:23.45823Z","iopub.execute_input":"2021-08-27T07:42:23.458584Z","iopub.status.idle":"2021-08-27T07:42:23.467372Z","shell.execute_reply.started":"2021-08-27T07:42:23.458554Z","shell.execute_reply":"2021-08-27T07:42:23.466514Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"path = '/kaggle/input/rsna-miccai-brain-tumor-radiogenomic-classification'\ntrain_data = pd.read_csv(os.path.join(path, 'train_labels.csv'))\nprint('Num of train samples:', len(train_data))\ntrain_data.head()\nimg_size = 256","metadata":{"execution":{"iopub.status.busy":"2021-08-27T07:42:32.30713Z","iopub.execute_input":"2021-08-27T07:42:32.307485Z","iopub.status.idle":"2021-08-27T07:42:32.325539Z","shell.execute_reply.started":"2021-08-27T07:42:32.30745Z","shell.execute_reply":"2021-08-27T07:42:32.324625Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def crop_imgs(img, add_pixels_value=0):\n    \"\"\"\n    Finds the extreme points on the image and crops the rectangular out of them\n    \"\"\"\n    #set_new = []\n    #gray = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY)\n    gray = cv2.GaussianBlur(img, (5, 5), 0)\n\n    # threshold the image, then perform a series of erosions +\n    # dilations to remove any small regions of noise\n    thresh = cv2.threshold(gray, 45, 255, cv2.THRESH_BINARY)[1]\n    thresh = cv2.erode(thresh, None, iterations=2)\n    thresh = cv2.dilate(thresh, None, iterations=2)\n    #print(thresh.shape)\n\n    # find contours in thresholded image, then grab the largest one\n    \n    cnts = cv2.findContours(thresh.copy(), cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n    cnts = imutils.grab_contours(cnts)\n    try:\n        c = max(cnts, key=cv2.contourArea)\n    except:\n        return img\n    \n\n    # find the extreme points\n    extLeft = tuple(c[c[:, :, 0].argmin()][0])\n    extRight = tuple(c[c[:, :, 0].argmax()][0])\n    extTop = tuple(c[c[:, :, 1].argmin()][0])\n    extBot = tuple(c[c[:, :, 1].argmax()][0])\n\n    ADD_PIXELS = add_pixels_value\n    new_img = img[extTop[1]-ADD_PIXELS:extBot[1]+ADD_PIXELS, extLeft[0]-ADD_PIXELS:extRight[0]+ADD_PIXELS].copy()\n\n    return new_img","metadata":{"execution":{"iopub.status.busy":"2021-08-27T07:42:40.44166Z","iopub.execute_input":"2021-08-27T07:42:40.442047Z","iopub.status.idle":"2021-08-27T07:42:40.451317Z","shell.execute_reply.started":"2021-08-27T07:42:40.441974Z","shell.execute_reply":"2021-08-27T07:42:40.450206Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def dicom2array(path, voi_lut=True, fix_monochrome=True):\n    dicom = pydicom.read_file(path)\n    # VOI LUT (if available by DICOM device) is used to\n    # transform raw DICOM data to \"human-friendly\" view\n    if voi_lut:\n        data = apply_voi_lut(dicom.pixel_array, dicom)\n    else:\n        data = dicom.pixel_array\n    # depending on this value, X-ray may look inverted - fix that:\n    if fix_monochrome and dicom.PhotometricInterpretation == \"MONOCHROME1\":\n        data = np.amax(data) - data\n    data = data - np.min(data)\n    data = data / np.max(data)\n    data = (data * 255).astype(np.uint8)\n    data = crop_imgs(data)\n    data = cv2.resize(data, (img_size, img_size))\n    return data\n\ndef load_3d_dicom_images(scan_id, split = \"train\"):\n    \"\"\"\n    we will use some heuristics to choose the slices to avoid any numpy zero matrix (if possible)\n    \"\"\"\n    flair = sorted(glob.glob(f\"{path}/{split}/{scan_id}/FLAIR/*.dcm\"))\n    t1w = sorted(glob.glob(f\"{path}/{split}/{scan_id}/T1w/*.dcm\"))\n    t1wce = sorted(glob.glob(f\"{path}/{split}/{scan_id}/T1wCE/*.dcm\"))\n    t2w = sorted(glob.glob(f\"{path}/{split}/{scan_id}/T2w/*.dcm\"))\n    \n    \n    flair_img = np.array([dicom2array(a) for a in flair[len(flair)//2 - 10:len(flair)//2 + 10]]).T\n    \n    if flair_img.shape[-1] < 20:\n        n_zero = 20 - flair_img.shape[-1]\n        flair_img = np.concatenate((flair_img, np.zeros((img_size, img_size, n_zero))), axis = -1)\n    #print(flair_img.shape)\n        \n    \n    \n    t1w_img = np.array([dicom2array(a) for a in t1w[len(t1w)//2 - 10:len(t1w)//2 + 10]]).T\n    if t1w_img.shape[-1] < 20:\n        n_zero = 20 - t1w_img.shape[-1]\n        t1w_img = np.concatenate((t1w_img, np.zeros((img_size, img_size, n_zero))), axis = -1)\n    #print(t1w_img.shape)\n    \n    \n    t1wce_img = np.array([dicom2array(a) for a in t1wce[len(t1wce)//2 - 10:len(t1wce)//2 + 10]]).T\n    if t1wce_img.shape[-1] < 20:\n        n_zero = 20 - t1wce_img.shape[-1]\n        t1wce_img = np.concatenate((t1wce_img, np.zeros((img_size, img_size, n_zero))), axis = -1)\n    #print(t1wce_img.shape)\n    \n    \n    t2w_img = np.array([dicom2array(a) for a in t2w[len(t2w)//2 - 10:len(t2w)//2 + 10]]).T\n    if t2w_img.shape[-1] < 20:\n        n_zero = 20 - t2w_img.shape[-1]\n        t2w_img = np.concatenate((t2w_img, np.zeros((img_size, img_size, n_zero))), axis = -1)\n    #print(t2w_img.shape)\n    \n    return np.concatenate((flair_img, t1w_img, t1wce_img, t2w_img), axis = -1)","metadata":{"execution":{"iopub.status.busy":"2021-08-27T07:42:49.742557Z","iopub.execute_input":"2021-08-27T07:42:49.742908Z","iopub.status.idle":"2021-08-27T07:42:49.757993Z","shell.execute_reply.started":"2021-08-27T07:42:49.742863Z","shell.execute_reply":"2021-08-27T07:42:49.757215Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"load_3d_dicom_images(\"00000\").shape","metadata":{"execution":{"iopub.status.busy":"2021-08-27T07:43:00.404458Z","iopub.execute_input":"2021-08-27T07:43:00.405084Z","iopub.status.idle":"2021-08-27T07:43:02.04897Z","shell.execute_reply.started":"2021-08-27T07:43:00.405046Z","shell.execute_reply":"2021-08-27T07:43:02.047582Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"slices = load_3d_dicom_images(\"00000\")\nprint(slices.shape)","metadata":{"execution":{"iopub.status.busy":"2021-08-27T07:43:08.454735Z","iopub.execute_input":"2021-08-27T07:43:08.455081Z","iopub.status.idle":"2021-08-27T07:43:09.043082Z","shell.execute_reply.started":"2021-08-27T07:43:08.455052Z","shell.execute_reply":"2021-08-27T07:43:09.042072Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_imgs(imgs, cols=5, size=7, is_rgb=True, title=\"\", cmap='gray', img_size=(64,64)):\n    rows = len(imgs)//cols + 1\n    fig = plt.figure(figsize=(cols*size, rows*size))\n    for i in range(cols):\n        img = imgs[:,:,i]\n        fig.add_subplot(rows, cols, i+1)\n        plt.imshow(img, cmap=cmap)\n    plt.suptitle(title)\n    plt.show()\n    \nplot_imgs(slices)","metadata":{"execution":{"iopub.status.busy":"2021-08-27T07:43:19.348821Z","iopub.execute_input":"2021-08-27T07:43:19.349202Z","iopub.status.idle":"2021-08-27T07:43:20.399472Z","shell.execute_reply.started":"2021-08-27T07:43:19.349161Z","shell.execute_reply":"2021-08-27T07:43:20.398416Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class BrainTumor(Dataset):\n    def __init__(self, path = '/kaggle/input/rsna-miccai-brain-tumor-radiogenomic-classification', split = \"train\", validation_split = 0.20):\n        # labels\n        train_data = pd.read_csv(os.path.join(path, 'train_labels.csv'))\n        self.labels = {}\n        brats = list(train_data[\"BraTS21ID\"])\n        mgmt = list(train_data[\"MGMT_value\"])\n        for b, m in zip(brats, mgmt):\n            self.labels[str(b).zfill(5)] = m\n            \n        if split == \"valid\":\n            self.split = \"train\"\n            self.ids = [a.split(\"/\")[-1] for a in sorted(glob.glob(path + f\"/{'train'}/\" + \"/*\"))]\n            #print(len(self.ids))\n            self.ids = self.ids[:int(len(self.ids)* validation_split)] # first 20% as validation\n            #print(len(self.ids))\n            #print(self.ids[0])\n        elif split == \"train\":\n            self.split = split\n            self.ids = [a.split(\"/\")[-1] for a in sorted(glob.glob(path + f\"/{split}/\" + \"/*\"))]\n            #print(len(self.ids))\n            self.ids = self.ids[int(len(self.ids)* validation_split):] # last 80% as train\n            #print(len(self.ids))\n            #print(self.ids[0])\n        else:\n            self.split = split\n            self.ids = [a.split(\"/\")[-1] for a in sorted(glob.glob(path + f\"/{split}/\" + \"/*\"))]\n            \n    \n    def __len__(self):\n        return len(self.ids)\n    \n    def __getitem__(self, idx):\n        imgs = load_3d_dicom_images(self.ids[idx], self.split)\n        transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,) * 80, (0.5,) * 80)])\n        imgs = transform(imgs)\n        imgs = torch.unsqueeze(imgs, 0)\n        \n        if self.split != \"test\":\n            label = self.labels[self.ids[idx]]\n            return torch.tensor(imgs, dtype = torch.float32), torch.tensor(label, dtype = torch.long)\n        else:\n            return torch.tensor(imgs, dtype = torch.float32)","metadata":{"execution":{"iopub.status.busy":"2021-08-27T07:43:28.235225Z","iopub.execute_input":"2021-08-27T07:43:28.235851Z","iopub.status.idle":"2021-08-27T07:43:28.248618Z","shell.execute_reply.started":"2021-08-27T07:43:28.235804Z","shell.execute_reply":"2021-08-27T07:43:28.247869Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = BrainTumor()\ntrain_loader = DataLoader(train_dataset, batch_size=2, shuffle=True)\nval_dataset = BrainTumor(split=\"valid\")\nval_loader = DataLoader(val_dataset, batch_size=2, shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2021-08-27T07:43:37.785629Z","iopub.execute_input":"2021-08-27T07:43:37.785986Z","iopub.status.idle":"2021-08-27T07:43:37.882328Z","shell.execute_reply.started":"2021-08-27T07:43:37.785954Z","shell.execute_reply":"2021-08-27T07:43:37.881298Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"x,y = train_dataset[0]\nprint(x.shape)\nprint(y)\nfor img, label in train_loader:\n    print(img.shape)\n    print(label.shape)\n    break","metadata":{"execution":{"iopub.status.busy":"2021-08-27T07:43:53.108504Z","iopub.execute_input":"2021-08-27T07:43:53.108839Z","iopub.status.idle":"2021-08-27T07:43:55.804592Z","shell.execute_reply.started":"2021-08-27T07:43:53.108809Z","shell.execute_reply":"2021-08-27T07:43:55.802544Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"x,y = val_dataset[1]\nprint(x.shape)\nprint(y)\nfor img, label in val_loader:\n    print(img.shape)\n    print(label.shape)\n    break","metadata":{"execution":{"iopub.status.busy":"2021-08-27T07:44:02.801869Z","iopub.execute_input":"2021-08-27T07:44:02.802545Z","iopub.status.idle":"2021-08-27T07:44:05.551561Z","shell.execute_reply.started":"2021-08-27T07:44:02.802495Z","shell.execute_reply":"2021-08-27T07:44:05.550708Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nclass SELayer(nn.Module):\n    def __init__(self, channel, reduction=16):\n        super(SELayer, self).__init__()\n        self.avg_pool = nn.AdaptiveAvgPool3d((1,1,1))\n        self.fc = nn.Sequential(\n            nn.Linear(channel, channel // reduction, bias=False),\n            nn.ReLU(inplace=True),\n            nn.Linear(channel // reduction, channel, bias=False),\n            nn.Sigmoid()\n        )\n\n    def forward(self, x):\n        b, c, _, _, _ = x.size()\n        y = self.avg_pool(x).view(b, c)\n        y = self.fc(y).view(b, c, 1, 1, 1)\n        return x * y.expand_as(x)\n\n\n\nclass Inception(nn.Module):\n\tdef __init__(\n\t\tself,\n\t\tin_channels,\n\t\tch1x1,\n\t\tch3x3red,\n\t\tch3x3,\n\t\tch5x5red,\n\t\tch5x5,\n\t\tpooling\n\t\t):\n\t\tsuper(Inception, self).__init__()\n\n\t\t# 1x1 conv branch\n\t\tself.branch1 = nn.Sequential(\n\t\t\tnn.Conv3d(in_channels, ch1x1, kernel_size=(1, 1, 1), bias=False),\n\t\t\tnn.BatchNorm3d(ch1x1),\n\t\t\tnn.ReLU()\n\t\t\t)\n\n\t\t# 1x1 conv + 3x3 conv branch\n\t\tself.branch2 = nn.Sequential(\n\t\t\tnn.Conv3d(in_channels, ch3x3red, kernel_size=(1, 1, 1), bias=False),\n\t\t\tnn.BatchNorm3d(ch3x3red),\n\t\t\tnn.ReLU(),\n\t\t\tnn.Conv3d(ch3x3red, ch3x3, kernel_size=(3, 3, 3), padding=(1, 1, 1), bias=False), \n\t\t\tnn.BatchNorm3d(ch3x3),\n\t\t\tnn.ReLU()\n\t\t\t)\n\n\t\t# 1x1 conv + 5x5 conv branch\n\t\tself.branch3 = nn.Sequential(\n\t\t\tnn.Conv3d(in_channels, ch5x5red, kernel_size=(1, 1, 1), bias=False),\n\t\t \tnn.BatchNorm3d(ch5x5red),\n\t\t \tnn.ReLU(),\n\t\t \tnn.Conv3d(ch5x5red, ch5x5, kernel_size=(5, 5, 5), padding=(2, 2, 2), bias=False), \n\t\t \tnn.BatchNorm3d(ch5x5),\n\t\t \tnn.ReLU()\n\t\t \t)\n\n\t\t# 3x3 pool + 1x1 conv branch\n\t\tself.branch4 = nn.Sequential(\n\t\t\tnn.MaxPool3d(kernel_size=(1, 3, 3), stride=(1, 1, 1), padding=(0, 1, 1), ceil_mode=True),\n\t\t  \tnn.Conv3d(in_channels, pooling, kernel_size=(1, 1, 1), bias=False),\n\t        nn.BatchNorm3d(pooling),\n\t        nn.ReLU()\n\t\t  \t)\n\n\tdef forward(self, x):\n\t\tbranch1 = self.branch1(x)\n\t\t#print(branch1.shape)\n\t\tbranch2 = self.branch2(x)\n\t\t#print(branch2.shape)\n\t\tbranch3 = self.branch3(x)\n\t\t#print(branch3.shape)\n\t\tbranch4 = self.branch4(x)\n\t\t#print(branch4.shape)\n\t\treturn torch.cat([branch1, branch2, branch3, branch4], 1)\n\nclass GoogLeNet(nn.Module):\n\tdef __init__(self, num_classes):\n\t\tsuper(GoogLeNet, self).__init__()\n\n\t\t#conv layers before inception\n\t\tself.pre_inception = nn.Sequential(\n\t\t\tnn.Conv3d(1, 64, kernel_size=(7, 7, 7), stride=(1, 2, 2), padding=(3, 3, 3)),\n\t\t\tnn.MaxPool3d(kernel_size=(1, 3, 3), stride=(2, 2, 2), ceil_mode=True),\n\t\t\tnn.ReLU(),\n\t\t\tnn.Conv3d(64, 64, kernel_size=(1, 1, 1)),\n\t\t\tnn.ReLU(),\n\t\t\tnn.Conv3d(64, 192, kernel_size=(3, 3, 3), padding=(1, 1, 1)),\n\t\t\tnn.MaxPool3d(kernel_size=(1, 3, 3), stride=(1, 2, 2), ceil_mode=True),\n\t\t\tnn.ReLU()\n\t\t\t)\n\n\t\tself.inception3a = Inception(192, 64, 96, 128, 16, 32, 32)\n\t\tself.se1 = SELayer(256)\n\t\tself.inception3b = Inception(256, 128, 128, 192, 32, 96, 64)\n\t\tself.se2 = SELayer(480)\n\t\tself.maxpool = nn.MaxPool3d(kernel_size=(1, 3, 3), stride=(1, 2, 2), ceil_mode=True)\n\t\tself.inception4a = Inception(480, 192, 96, 208, 16, 48, 64)\n\t\tself.se3 = SELayer(512)\n\t\tself.inception4b = Inception(512, 160, 112, 224, 24, 64, 64)\n\t\tself.se4 = SELayer(512)\n\t\tself.inception4c = Inception(512, 128, 128, 256, 24, 64, 64)\n\t\tself.se5 = SELayer(512)\n\t\tself.inception4d = Inception(512, 112, 144, 288, 32, 64, 64)\n\t\tself.se6 = SELayer(528)\n\t\tself.inception4e = Inception(528, 256, 160, 320, 32, 128, 128)\n\t\tself.se7 = SELayer(832)\n\t\tself.inception5a = Inception(832, 256, 160, 320, 32, 128, 128)\n\t\tself.se8 = SELayer(832)\n\t\t#self.inception5b = Inception(832, 384, 192, 384, 48, 128, 128)\n\t\t#self.se9 = SELayer(1024)\n\t\tself.avgpool = nn.AdaptiveAvgPool3d((1, 1, 1))\n\t\tself.dropout = nn.Dropout(0.4)\n\t\tself.fc1 = nn.Linear(832, 512)\n\t\tself.fc2 = nn.Linear(512, 128)\n\t\tself.fc3 = nn.Linear(128, num_classes)\n\n\tdef forward(self, x):\n\t\tx = self.pre_inception(x)\n\t\t#print(x.shape)\n\t\tx = self.inception3a(x)\n\t\tx = self.se1(x)\n\t\tx = self.inception3b(x)\n\t\tx = self.se2(x)\n\t\tx = self.maxpool(x)\n\t\tx = self.inception4a(x)\n\t\tx = self.se3(x)\n\t\tx = self.inception4b(x)\n\t\tx = self.se4(x)\n\t\tx = self.inception4c(x)\n\t\tx = self.se5(x)\n\t\tx = self.inception4d(x)\n\t\tx = self.se6(x)\n\t\tx = self.inception4e(x)\n\t\tx = self.se7(x)\n\t\tx = self.maxpool(x)\n\t\tx = self.inception5a(x)\n\t\tx = self.se8(x)\n\t\t#x = self.inception5b(x)\n\t\t#x = self.se9(x)\n\t\t#print(x.shape)\n\t\tx = self.avgpool(x)\n\t\t#print(x.shape)\n\t\tx = x.view(x.size(0), -1)\n\t\tx = F.elu(self.fc1(x))\n\t\tx = F.elu(self.fc2(x))\n\t\tx = self.fc3(x)\n\t\treturn F.log_softmax(x, dim=1)\n\n\nmodel = GoogLeNet(2)\n#x = torch.randn(2,3, 80,256,256)\n#y = model(x)\n#print(y.size())","metadata":{"execution":{"iopub.status.busy":"2021-08-27T07:46:08.334093Z","iopub.execute_input":"2021-08-27T07:46:08.334418Z","iopub.status.idle":"2021-08-27T07:46:08.477208Z","shell.execute_reply.started":"2021-08-27T07:46:08.334389Z","shell.execute_reply":"2021-08-27T07:46:08.476155Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(device)\nx = torch.randn(2,1, 80,256,256).to(device)\ny = model(x)\nprint(y.size())","metadata":{"execution":{"iopub.status.busy":"2021-08-27T07:46:09.754627Z","iopub.execute_input":"2021-08-27T07:46:09.755Z","iopub.status.idle":"2021-08-27T07:46:23.763539Z","shell.execute_reply.started":"2021-08-27T07:46:09.754963Z","shell.execute_reply":"2021-08-27T07:46:23.762792Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_parameters = filter(lambda p: p.requires_grad, model.parameters())\nparams = sum([np.prod(p.size()) for p in model_parameters])\nparams","metadata":{"execution":{"iopub.status.busy":"2021-08-27T07:47:44.834434Z","iopub.execute_input":"2021-08-27T07:47:44.834999Z","iopub.status.idle":"2021-08-27T07:47:44.848201Z","shell.execute_reply.started":"2021-08-27T07:47:44.834949Z","shell.execute_reply":"2021-08-27T07:47:44.846771Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}