{"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":"!pip install torchviz\n!pip install -q timm","metadata":{"execution":{"iopub.status.busy":"2023-01-03T16:05:03.82042Z","iopub.execute_input":"2023-01-03T16:05:03.820884Z","iopub.status.idle":"2023-01-03T16:05:25.529434Z","shell.execute_reply.started":"2023-01-03T16:05:03.820842Z","shell.execute_reply":"2023-01-03T16:05:25.528218Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import h5py\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport random\nimport os\nimport math\nimport time\nimport gc\n\nfrom sklearn.preprocessing import MinMaxScaler\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import DataLoader, Dataset, random_split\n\nimport torchvision\nfrom torchvision import transforms\nfrom torchvision.models import efficientnet_b0\n\nfrom torchviz import make_dot\n\nimport timm","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-output":true,"_kg_hide-input":false,"execution":{"iopub.status.busy":"2023-01-03T16:30:11.05057Z","iopub.execute_input":"2023-01-03T16:30:11.05099Z","iopub.status.idle":"2023-01-03T16:30:11.435563Z","shell.execute_reply.started":"2023-01-03T16:30:11.050955Z","shell.execute_reply":"2023-01-03T16:30:11.434344Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"timm.list_models(pretrained=True)","metadata":{"execution":{"iopub.status.busy":"2023-01-03T16:07:45.290866Z","iopub.execute_input":"2023-01-03T16:07:45.29133Z","iopub.status.idle":"2023-01-03T16:07:45.3149Z","shell.execute_reply.started":"2023-01-03T16:07:45.291294Z","shell.execute_reply":"2023-01-03T16:07:45.313775Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class config():\n    DEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    print(\"Using: \",DEVICE)\n    \n    KAGGLE = True\n    DS_PATH = '../input/g2net-detecting-continuous-gravitational-waves' if KAGGLE else \"./\"\n    TRAIN_PATH = os.path.join(DS_PATH, 'train')\n    TEST_PATH = os.path.join(DS_PATH, 'test')\n\n    TARGET_SIZE = [224,224]\n    \n    VAL_SPLIT = 0.2\n    BATCH_SIZE=32\n\n    EPOCHS = 20\n    LEARNING_RATE = 0.05\n    \n    \ndef seed_everything(seed):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    #the following line gives ~10% speedup\n    #but may lead to some stochasticity in the results \n    torch.backends.cudnn.benchmark = True\n    \nconfig()\nseed_everything(353)","metadata":{"execution":{"iopub.status.busy":"2023-01-03T15:36:24.392818Z","iopub.execute_input":"2023-01-03T15:36:24.39317Z","iopub.status.idle":"2023-01-03T15:36:24.407657Z","shell.execute_reply.started":"2023-01-03T15:36:24.393135Z","shell.execute_reply":"2023-01-03T15:36:24.406204Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f\"Train len: {len(os.listdir(config.TRAIN_PATH))}\\nTest len: {len(os.listdir(config.TEST_PATH))}\")","metadata":{"execution":{"iopub.status.busy":"2023-01-03T15:36:24.410052Z","iopub.execute_input":"2023-01-03T15:36:24.410426Z","iopub.status.idle":"2023-01-03T15:36:24.421385Z","shell.execute_reply.started":"2023-01-03T15:36:24.410392Z","shell.execute_reply":"2023-01-03T15:36:24.420249Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_labels_df = pd.read_csv(os.path.join(config.DS_PATH, \"train_labels.csv\"))\ntrain_labels_df = train_labels_df[train_labels_df.target >= 0]# 3 items are -1 label\n\nsub_df = pd.read_csv(os.path.join(config.DS_PATH, \"sample_submission.csv\"))\n\ntrain_labels_df","metadata":{"execution":{"iopub.status.busy":"2023-01-03T15:36:24.423339Z","iopub.execute_input":"2023-01-03T15:36:24.423769Z","iopub.status.idle":"2023-01-03T15:36:24.449155Z","shell.execute_reply.started":"2023-01-03T15:36:24.423729Z","shell.execute_reply":"2023-01-03T15:36:24.44835Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def print_hdf5(id_sample, isTrain = True):\n    h5_sample = h5py.File(os.path.join(config.TRAIN_PATH if isTrain else config.TEST_PATH,id_sample + \".hdf5\"), \"r\")[id_sample]\n    print(f\"HDF5 Obj Keys: {list(h5_sample.keys())}\\n\")\n    for k in [\"H1\", \"L1\"]:\n        print(f\"{k}:\")\n        for h in list(h5_sample[k].keys()):\n            print(f\"\\t{h}: {h5_sample[k][h].shape}\")\n    print(f\"\\nfrequency_Hz: {h5_sample['frequency_Hz'].shape}\")\n    \nprint_hdf5(\"010a387db\")","metadata":{"execution":{"iopub.status.busy":"2023-01-03T15:36:24.451109Z","iopub.execute_input":"2023-01-03T15:36:24.451694Z","iopub.status.idle":"2023-01-03T15:36:24.469623Z","shell.execute_reply.started":"2023-01-03T15:36:24.451659Z","shell.execute_reply":"2023-01-03T15:36:24.46851Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def augment(image, hflip=False, vflip=False, vshift=0):\n    \"\"\"\n    Params\n    \n    image: np.array(2, 360, 128): Apply the transformation to this image\n    vflip: Boolean for vertical flip\n    hflip: Boolean for horizontal flip\n    vshift: Shift amount in the vertical direction\n    \"\"\"\n    new_img = image.copy()\n    \n    if hflip:\n        new_img[0] = np.fliplr(new_img[0])\n        new_img[1] = np.fliplr(new_img[1])\n        \n    if vflip:\n        new_img[0] = np.flipud(new_img[0])\n        new_img[1] = np.flipud(new_img[1])\n    \n    if isinstance(vshift, int):\n        if vshift != 0:\n            new_img[0] = np.roll(new_img[0], vshift, axis=0)\n            new_img[1] = np.roll(new_img[1], vshift, axis=0)\n    if not isinstance(vshift, int):\n        print(\"Error: vshift expects an int or tuple of ints\")\n        return\n        \n    return new_img","metadata":{"execution":{"iopub.status.busy":"2023-01-03T15:36:24.471153Z","iopub.execute_input":"2023-01-03T15:36:24.4716Z","iopub.status.idle":"2023-01-03T15:36:24.481828Z","shell.execute_reply.started":"2023-01-03T15:36:24.471567Z","shell.execute_reply":"2023-01-03T15:36:24.480874Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class G2NetDataset(Dataset):\n    \n    def __init__(self, labels, transforms = None):\n        self.labels = labels\n        self.transforms = transforms\n    \n    def __len__(self):\n        return len(self.labels)\n    \n    def __getitem__(self, index):\n        data = self.labels.iloc[index]\n        file_id = str(data['id'])\n        \n        y = np.asarray(data['target'], dtype = np.uint8)\n        \n        img = np.empty((2, config.TARGET_SIZE[0], config.TARGET_SIZE[1]))\n        \n        filename = file_id + \".hdf5\"\n        hdf5_filepath = os.path.join(config.TRAIN_PATH, filename)\n        with h5py.File(hdf5_filepath, 'r') as f:\n            group = f[file_id]\n            \n            for i, obs in enumerate(['H1', 'L1']):\n                sft = group[obs]['SFTs'][:, :4096] * 1e22\n                m = sft.real ** 2 + sft.imag ** 2\n                m /= np.mean(m)\n                m = np.mean(m.reshape(360, 128, 32), axis=2)\n                    \n                if self.transforms:\n                    m = self.transforms(m)\n                    \n                img[i] = m.float()\n            \n            if random.random() < .5:\n                img = augment(img, hflip = True)\n            if random.random() < .5:\n                img = augment(img, vflip = True)    \n            if random.random() < .5:\n                img = augment(img, vshift = 2)\n            #img[2] = np.mean(np.stack([img[0], img[1]]), axis = 0)\n            \n        return img, y","metadata":{"execution":{"iopub.status.busy":"2023-01-03T15:36:24.483235Z","iopub.execute_input":"2023-01-03T15:36:24.48428Z","iopub.status.idle":"2023-01-03T15:36:24.496236Z","shell.execute_reply.started":"2023-01-03T15:36:24.484245Z","shell.execute_reply":"2023-01-03T15:36:24.495125Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"a = imgs_to_plot[0]\nprint(a.max(), a.min())\n\nprint(a.max(), a.min())","metadata":{"execution":{"iopub.status.busy":"2023-01-03T16:26:41.059822Z","iopub.execute_input":"2023-01-03T16:26:41.060237Z","iopub.status.idle":"2023-01-03T16:26:41.071578Z","shell.execute_reply.started":"2023-01-03T16:26:41.060179Z","shell.execute_reply":"2023-01-03T16:26:41.070111Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transform =  transforms.Compose([\n    transforms.ToPILImage(),\n    transforms.Resize(256, interpolation = transforms.InterpolationMode.BICUBIC),\n    transforms.CenterCrop(224),\n    transforms.ToTensor(),\n    transforms.Normalize((0.5), (0.5))])\n\ndataset = G2NetDataset(train_labels_df, transforms=transform)\n\nsample_img, sample_y = dataset[0]\nprint(f\"sample_img shape : {sample_img.shape}, sample_y: {sample_y}\")t","metadata":{"execution":{"iopub.status.busy":"2023-01-03T15:36:24.515813Z","iopub.execute_input":"2023-01-03T15:36:24.516138Z","iopub.status.idle":"2023-01-03T15:36:24.903414Z","shell.execute_reply.started":"2023-01-03T15:36:24.516107Z","shell.execute_reply":"2023-01-03T15:36:24.902397Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_len=math.floor(len(dataset)*config.VAL_SPLIT)\ntrain_len=len(dataset) - val_len\n\ntrain_ds,val_ds = torch.utils.data.random_split(dataset,[train_len,val_len])\nprint(f\"Training dataset length : {len(train_ds)}\\nValidation dataset length: {len(val_ds)}\")","metadata":{"execution":{"iopub.status.busy":"2023-01-03T15:36:24.90481Z","iopub.execute_input":"2023-01-03T15:36:24.90584Z","iopub.status.idle":"2023-01-03T15:36:24.914428Z","shell.execute_reply.started":"2023-01-03T15:36:24.905798Z","shell.execute_reply":"2023-01-03T15:36:24.913249Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_loader=DataLoader(\n    train_ds,\n    batch_size=config.BATCH_SIZE,\n    shuffle=False,\n    pin_memory=True,\n    num_workers=os.cpu_count()\n)\n\nval_loader=DataLoader(\n    val_ds,\n    batch_size=config.BATCH_SIZE,\n    shuffle=False,\n    pin_memory=True,\n    num_workers=os.cpu_count()\n)","metadata":{"execution":{"iopub.status.busy":"2023-01-03T15:36:24.916625Z","iopub.execute_input":"2023-01-03T15:36:24.917397Z","iopub.status.idle":"2023-01-03T15:36:24.925607Z","shell.execute_reply.started":"2023-01-03T15:36:24.91735Z","shell.execute_reply":"2023-01-03T15:36:24.92467Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_SFTs(imgs, labels, cm = [\"inferno\", \"viridis\"]):\n    n = len(imgs)\n    fig, ax = plt.subplots(n, 2, figsize=(int(n/2), int(5.4*n)))\n    titles = ['H1', 'L1']\n    \n    k = 0\n    for i in range(n):\n        for j in range(2):\n            sft_img = imgs[k]\n            \n            ax[i,j].set_title(f\"{titles[j]}, Label: {labels[k]}\")\n            p = ax[i,j].imshow(sft_img[j], cmap=cm[j])\n            fig.colorbar(p, ax=ax[i,j])\n            ax[i, j].set_xlabel(\"Timestamp\")\n            ax[i, j].grid(False)\n            \n        k += 1","metadata":{"execution":{"iopub.status.busy":"2023-01-03T15:36:24.92842Z","iopub.execute_input":"2023-01-03T15:36:24.928963Z","iopub.status.idle":"2023-01-03T15:36:24.93876Z","shell.execute_reply.started":"2023-01-03T15:36:24.928924Z","shell.execute_reply":"2023-01-03T15:36:24.937549Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"imgs_to_plot,labels_to_plot = next(iter(train_loader))\nimgs_to_plot.numpy().shape","metadata":{"execution":{"iopub.status.busy":"2023-01-03T15:36:24.940801Z","iopub.execute_input":"2023-01-03T15:36:24.941243Z","iopub.status.idle":"2023-01-03T15:36:36.426212Z","shell.execute_reply.started":"2023-01-03T15:36:24.941205Z","shell.execute_reply":"2023-01-03T15:36:36.425202Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f\"SFTs data of the first batch (batch_size={config.BATCH_SIZE}) for L1 and H1 observatories:\\n\")\nplot_SFTs(imgs_to_plot.numpy(), labels_to_plot.numpy())\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-01-03T15:36:36.427861Z","iopub.execute_input":"2023-01-03T15:36:36.428245Z","iopub.status.idle":"2023-01-03T15:36:50.067445Z","shell.execute_reply.started":"2023-01-03T15:36:36.428208Z","shell.execute_reply":"2023-01-03T15:36:50.066277Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Efficentnet","metadata":{}},{"cell_type":"code","source":"model = efficientnet_b0(pretrained=True)\n\nfor param in model.parameters():\n    param.requires_grad = False\n    \n# Change the number of output units\nfc = nn.Sequential(\n    nn.Dropout(p=0.4, inplace=True),\n    nn.Linear(1280, 1),\n    nn.Sigmoid()\n)\nmodel.classifier = fc","metadata":{"execution":{"iopub.status.busy":"2023-01-03T15:36:50.068871Z","iopub.execute_input":"2023-01-03T15:36:50.069344Z","iopub.status.idle":"2023-01-03T15:36:52.539858Z","shell.execute_reply.started":"2023-01-03T15:36:50.069292Z","shell.execute_reply":"2023-01-03T15:36:52.53885Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# My Model","metadata":{}},{"cell_type":"code","source":"class G2Net(nn.Module):\n    def __init__(self):\n        super(G2Net, self).__init__()\n        \n        h1_eff = efficientnet_b0(pretrained=True)\n        self.h1_eff = nn.Sequential(*list(h1_eff.children())[:-1])\n        h1_eff = list(h1_eff.children())\n        w_h1 = h1_eff[0][0][0].weight\n        h1_eff[0][0][0] = nn.Conv2d(1, 32, kernel_size=3, stride=2, padding=1, bias=False)\n        h1_eff[0][0][0].weight = nn.Parameter(torch.mean(w_h1, dim=1, keepdim=True))\n        h1_eff = nn.Sequential(*h1_eff)\n        \n        for param in h1_eff.parameters():\n            param.requires_grad = False\n\n            \n        l1_eff = efficientnet_b0(pretrained=True)\n        self.l1_eff = nn.Sequential(*list(l1_eff.children())[:-1])\n        l1_eff = list(l1_eff.children())\n        w_l1 = l1_eff[0][0][0].weight\n        l1_eff[0][0][0] = nn.Conv2d(1, 32, kernel_size=3, stride=2, padding=1, bias=False)\n        l1_eff[0][0][0].weight = nn.Parameter(torch.mean(w_l1, dim=1, keepdim=True))\n        l1_eff = nn.Sequential(*l1_eff)\n        \n        for param in self.l1_eff.parameters():\n            param.requires_grad = False\n            \n        # Change the number of output units\n        self.out = nn.Sequential(\n            nn.Dropout(p=0.4, inplace=True),\n            nn.Linear(2560, 256),\n            nn.Dropout(p=0.4, inplace=True),\n            nn.Linear(256, 1),\n            nn.Sigmoid()\n        )\n        \n    def forward(self, x):\n        h1_img, l1_img = x[:, 0, :, :], x[:, 1, :, :]\n        \n        #print(x.shape, h1_img.shape, l1_img.shape)\n        #print(h1_img.reshape((config.BATCH_SIZE, 1, config.TARGET_SIZE[0], config.TARGET_SIZE[1])).shape)\n        h1_out = self.h1_eff(h1_img.reshape((x.shape[0], 1, config.TARGET_SIZE[0], config.TARGET_SIZE[1])))\n        l1_out = self.l1_eff(l1_img.reshape((x.shape[0], 1, config.TARGET_SIZE[0], config.TARGET_SIZE[1])))\n        \n        x = torch.cat([h1_out.view(h1_out.size(0), -1), l1_out.view(l1_out.size(0), -1)], dim=1)\n        #print(x.shape)\n        x = self.out(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2023-01-03T15:36:52.575142Z","iopub.execute_input":"2023-01-03T15:36:52.575763Z","iopub.status.idle":"2023-01-03T15:36:52.591611Z","shell.execute_reply.started":"2023-01-03T15:36:52.575727Z","shell.execute_reply":"2023-01-03T15:36:52.590478Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = G2Net().to(config.DEVICE)","metadata":{"execution":{"iopub.status.busy":"2023-01-03T15:36:52.593822Z","iopub.execute_input":"2023-01-03T15:36:52.594135Z","iopub.status.idle":"2023-01-03T15:36:55.831684Z","shell.execute_reply.started":"2023-01-03T15:36:52.594108Z","shell.execute_reply":"2023-01-03T15:36:55.83067Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# TIMM","metadata":{}},{"cell_type":"code","source":"model = timm.create_model(\"efficientnetv2_rw_s\", pretrained=True, num_classes=1, in_chans=2, drop_rate=0.3).to(config.DEVICE)\n\nparams = sum([np.prod(p.size()) for p in model.parameters()])\nprint(\"Network params:\", params)","metadata":{"execution":{"iopub.status.busy":"2023-01-03T16:11:50.079289Z","iopub.execute_input":"2023-01-03T16:11:50.079669Z","iopub.status.idle":"2023-01-03T16:11:50.890717Z","shell.execute_reply.started":"2023-01-03T16:11:50.079636Z","shell.execute_reply":"2023-01-03T16:11:50.889734Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"loss_fn = nn.BCELoss()\noptimizer = torch.optim.Adam(model.parameters(),lr=config.LEARNING_RATE)\nscheduler=torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer,mode=\"min\",patience=5,verbose=True)\n\nbest_val_loss=float(1e6)\nsince = time.time()\n\nfor n_epoch in range(1,config.EPOCHS):\n    \n    print(\"EPOCH : \"+str(n_epoch)+\"/\"+str(config.EPOCHS))\n    \n    \n    running_train_loss=0.0\n    running_val_loss=0.0\n    \n    \n    model.train()\n    for train_batch_idx,train_batch in enumerate(train_loader):\n        optimizer.zero_grad()\n\n        #PREDICT\n        images,labels = train_batch\n        images,labels = images.to(config.DEVICE),labels.to(config.DEVICE)\n        \n        preds=model(images.float())\n        train_loss=loss_fn(preds.squeeze(),labels.float())\n        \n        gc.collect()\n        del train_batch\n        del images\n        \n        #BACKPROPAGATION\n        train_loss.backward()\n        optimizer.step()\n        running_train_loss += train_loss.item()\n        \n        gc.collect()\n        if torch.cuda.is_available():\n            torch.cuda.empty_cache()\n\n\n    model.eval()\n    #VALIDATION\n    with torch.no_grad():\n        for val_batch_idx,val_batch in enumerate (val_loader):\n\n            #Predict\n            images,labels = val_batch\n            images,labels=images.to(config.DEVICE),labels.to(config.DEVICE)\n            val_preds=model(images.float())\n            val_loss=loss_fn(val_preds.squeeze(),labels.float())\n\n            gc.collect()\n            del val_batch\n            del images\n\n            running_val_loss+=val_loss.item()\n        \n    running_train_loss /= train_batch_idx+1\n    running_val_loss /= val_batch_idx+1\n    \n    #Reduce LR on Plateau\n    scheduler.step(running_val_loss) \n\n    print(f\"EPOCH : {n_epoch} Train Loss : {running_train_loss:.5f}, Val Loss : {running_val_loss:.5f}\")\n    if(running_val_loss < best_val_loss):\n        torch.save(model.state_dict(), \"/kaggle/working/best_model.pth\")\n        print(\"Model Saved\")\n        best_val_loss=running_val_loss\n        \ntime_elapsed = time.time() - since\nprint('Training complete in {:.0f}m {:.0f}s'.format(\ntime_elapsed // 60, time_elapsed % 60))\nprint('Best Val Loss: {:4f}'.format(best_val_loss))","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-01-03T15:36:55.833043Z","iopub.execute_input":"2023-01-03T15:36:55.833422Z","iopub.status.idle":"2023-01-03T15:50:54.753588Z","shell.execute_reply.started":"2023-01-03T15:36:55.833386Z","shell.execute_reply":"2023-01-03T15:50:54.750626Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_loss.item()","metadata":{"execution":{"iopub.status.busy":"2023-01-03T15:50:54.756965Z","iopub.status.idle":"2023-01-03T15:50:54.757368Z","shell.execute_reply.started":"2023-01-03T15:50:54.757167Z","shell.execute_reply":"2023-01-03T15:50:54.757197Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"sleeping\")\ntime.sleep(1000000)","metadata":{"execution":{"iopub.status.busy":"2023-01-03T15:50:54.759083Z","iopub.status.idle":"2023-01-03T15:50:54.759826Z","shell.execute_reply.started":"2023-01-03T15:50:54.759569Z","shell.execute_reply":"2023-01-03T15:50:54.759594Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"batch = next(iter(train_loader))\nyhat = model(batch.text)\nmake_dot(yhat, params=dict(list(model.named_parameters()))).render(\"model\", format=\"png\")","metadata":{"execution":{"iopub.status.busy":"2023-01-01T14:34:21.641541Z","iopub.execute_input":"2023-01-01T14:34:21.642017Z","iopub.status.idle":"2023-01-01T14:34:33.932083Z","shell.execute_reply.started":"2023-01-01T14:34:21.641979Z","shell.execute_reply":"2023-01-01T14:34:33.929923Z"}}}]}