{"cells":[{"metadata":{},"cell_type":"markdown","source":"[Next notebook](https://www.kaggle.com/keremt/06-inference-multi-output)"},{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"cell_type":"code","source":"from fastai.vision.all import *","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"pd.options.display.max_columns = 100\npd.options.display.max_rows = 100","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"embspath = Path(\"/kaggle/input/rsnaperawembs256v2//\")","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"embspath.ls()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_feat_df = pd.read_csv(embspath/'train.csv')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# TODO: Make sure to normalize views for pe: left-right-central","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_feat_df.head(2)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"embs = torch.cat([torch.load(embspath/f'train_embs-{o}.pkl') for o in [0,1,2,'final']])\nembs.shape","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# add padding embedding and idx\nembs_len = len(embs)\npad_emb = torch.zeros_like(embs[:1])\nembs = torch.cat([embs, pad_emb])\ninput_pad_idx = embs_len","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"embs[input_pad_idx], embs.shape, input_pad_idx","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"seq_lens = train_feat_df.groupby(\"StudyInstanceUID\").apply(len)\nseq_lens.hist();","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"(seq_lens < 300).mean()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"### Read metadata"},{"metadata":{"trusted":true},"cell_type":"code","source":"metadata_path = Path(\"/kaggle/input/rsnastrpemetadata/\")","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_metadf = pd.concat([pd.read_parquet(metadata_path/f\"train_metadf_part{i}.pqt\") for i in range(1,3)]).reset_index(drop=True)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_metadf.head()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_meta_feats = train_metadf[['StudyInstanceUID', 'SOPInstanceUID', 'ImagePositionPatient2', 'img_min', 'img_max', 'img_mean', 'img_std', 'img_pct_window']]","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def minmax_scaler(o): return (o - min(o))/(max(o) - min(o))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"scaled_pos = train_meta_feats.groupby('StudyInstanceUID')['ImagePositionPatient2'].apply(minmax_scaler)\ntrain_meta_feats['scaled_position'] = scaled_pos","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_meta_feats.isna().sum()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_meta_feats.sort_values('ImagePositionPatient2')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"meta_feats = ['img_min', 'img_max', 'img_mean', 'img_std', 'scaled_position']\nmean_std = {}\nfor f in meta_feats:\n    mean,std = train_meta_feats[f].mean(), train_meta_feats[f].std()\n    train_meta_feats[f] = (train_meta_feats[f] - mean) / std\n    mean_std[f] = (mean, std)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_meta_feats","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"meta_feats_dict = dict(zip(train_meta_feats['SOPInstanceUID'], train_meta_feats[meta_feats].to_numpy()))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"meta_feats_dict['cc96a7a2e72c']","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"len(meta_feats_dict)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"meta_embs = np.vstack([meta_feats_dict[o] for o in train_feat_df['SOPInstanceUID'].values])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"meta_embs.shape","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"meta_embs = tensor(np.vstack((meta_embs, np.zeros((1,meta_embs.shape[1])))))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"meta_embs.shape","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"markdown","source":"### Data"},{"metadata":{"trusted":true},"cell_type":"code","source":"from fastai.text.all import *","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# TODO: Use saved pids\nunique_pids = train_feat_df.StudyInstanceUID.unique()\nn = len(unique_pids)\nnvalid = int(n*0.05); nvalid, n\nunique_pids = np.random.permutation(unique_pids)\ntrain_pids = unique_pids[nvalid:]\nvalid_pids = unique_pids[:nvalid]\nlen(train_pids), len(valid_pids)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"files_dict = defaultdict(list)\nfor i, (_, row) in enumerate(train_feat_df.iterrows()):\n    fn = Path(row['fname'])\n    slice_no = int(fn.stem.split(\"_\")[0])\n    sid = fn.parent.parent.name\n    files_dict[sid].append(({**dict(row), **{\"slice_no\":slice_no, \"embs_idx\":i}}))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"dict(row)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"image_targets = L(['pe_present_on_image'])\nexam_targets = L([\n#           'positive_exam_for_pe'\n            'negative_exam_for_pe',\n            'indeterminate',\n\n            'rv_lv_ratio_gte_1',\n            'rv_lv_ratio_lt_1',\n    # none\n\n            'leftsided_pe',\n            'rightsided_pe',\n            'central_pe',\n\n            'chronic_pe',\n            'acute_and_chronic_pe',           \n            # neither chronic or acute_and_chronic\n          \n    \n    \n#             'qa_motion',\n#             'qa_contrast',\n#             'flow_artifact',\n#             'true_filling_defect_not_pe',\n             ]); exam_targets","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"len(files_dict)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"trn_pid = np.random.choice(train_pids)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def get_x(pid):\n    o = files_dict[pid]    \n    l = sorted(o, key=lambda x: x['slice_no']) \n    return tensor([o['embs_idx'] for o in l])\n\ndef get_img_y(pid):\n    o = files_dict[pid]    \n    l = sorted(o, key=lambda x: x['slice_no']) \n    img_y = [o['pe_present_on_image'] for o in l]\n    exam_y = [max(img_y)] + [o[0][t] for t in exam_targets]\n    return tensor(img_y)\n\ndef get_exam_y(pid):\n    \"\"\"\n    'POSITIVE','negative_exam_for_pe','indeterminate',\n    'rv_lv_ratio_gte_1','rv_lv_ratio_lt_1', 'NEITHER'\n    'leftsided_pe','rightsided_pe','central_pe',\n    'chronic_pe','acute_and_chronic_pe','NEITHER'\n    \"\"\"\n    o = files_dict[pid]    \n    l = sorted(o, key=lambda x: x['slice_no']) \n    img_y = [o['pe_present_on_image'] for o in l]\n    \n    exam_y = [max(img_y)] + [o[0][t] for t in exam_targets]\n    \n    none_chro_acute = [exam_y[-1] == exam_y[-2]]\n    exam_y += none_chro_acute\n    \n    none_rv_lv = [exam_y[3] == exam_y[4]]\n    exam_y = exam_y[:4] + none_rv_lv + exam_y[4:]\n    \n    \n    return tensor(exam_y)\n\n# before_batch: after collecting samples before collating\ntarg_pad_idx = 666\ndef SequenceBlock():       return  TransformBlock(type_tfms=[get_x], dl_type=SortedDL, dls_kwargs={'before_batch':\n                                                       [partial(pad_input, pad_idx=input_pad_idx),\n                                                        partial(pad_input, pad_idx=targ_pad_idx, pad_fields=1)]})\ndef SequenceTargetBlock(): return TransformBlock(type_tfms=[get_img_y])\ndef TargetBlock():         return TransformBlock(type_tfms=[get_exam_y])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"data = DataBlock(blocks=(SequenceBlock,SequenceTargetBlock,TargetBlock), \n                 n_inp=1, \n                 splitter=FuncSplitter(lambda o: o in valid_pids))\ndls = data.dataloaders(list(train_pids)+list(valid_pids), bs=128)\nb = dls.one_batch()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"learner = Learner(dls, nn.Linear(10,10), loss_func=noop)\nlearner._split(b)\nlen(learner.xb), len(learner.yb)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"learner.xb[0].shape, learner.yb[0].shape, learner.yb[1].shape","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"embs[learner.xb[0]].shape, meta_embs[learner.xb[0]].shape","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"torch.cat([embs[learner.xb[0]], meta_embs[learner.xb[0]]], -1).shape","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"markdown","source":"### Model"},{"metadata":{"trusted":true},"cell_type":"code","source":"device = default_device()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class MultiHeadedSequenceClassifier(Module):\n    \"dim: input sequence feature dim\"\n    def __init__(self, input_pad_idx=input_pad_idx, dim=1024):\n        \n        store_attr('input_pad_idx')\n        self.lstm1 = nn.LSTM(dim+5, dim//16, bidirectional=True)\n#         self.lstm1 = nn.LSTM(dim+5, dim//8, bidirectional=True)\n#         self.lstm2 = nn.LSTM(dim//4, dim//16, bidirectional=True)\n        \n        # image level preds\n        self.seq_cls_head = nn.Linear(dim//8, 1)\n    \n        \n        # positive, negative, indeterminate\n        self.pe_head = nn.Linear(dim//4, 3) # softmax\n        # rv / lv >=,  < 1 or neither\n        self.rv_lv_head = nn.Linear(dim//4, 3) # softmax\n        # l,r,c pe\n        self.pe_position_head = nn.Linear(dim//4, 3) # sigmoid\n        # chronic, ac-chr or neither\n        self.chronic_pe_head = nn.Linear(dim//4, 3) # softmax\n        \n    \n    def forward(self, x):\n        \n        # get mask from non-pad idxs and then features\n        mask = x != self.input_pad_idx\n        x = torch.cat([embs[x], meta_embs[x]], dim=-1).to(device)\n        \n        # sequence outs\n        x, _ = self.lstm1(x) \n#         x, _ = self.lstm2(x)\n        seq_cls_out = self.seq_cls_head(x).squeeze(-1)\n        \n        \n        #masked concat pool\n        pooled_x = []\n        for i in range(x.size(0)):\n            xi = x[i, mask[i], :]\n            pooled_x.append(torch.cat([xi.mean(0), xi.max(0).values]).unsqueeze(0))\n        pooled_x = torch.cat(pooled_x)\n        \n\n        # 'POSITIVE','negative_exam_for_pe','indeterminate'\n        out1 = self.pe_head(pooled_x)\n\n        # 'rv_lv_ratio_gte_1','rv_lv_ratio_lt_1', 'NEITHER'\n        out2 = self.rv_lv_head(pooled_x)\n\n        # 'leftsided_pe','rightsided_pe','central_pe',\n        out3 = self.pe_position_head(pooled_x)\n\n        # 'chronic_pe','acute_and_chronic_pe','NEITHER'\n        out4 = self.chronic_pe_head(pooled_x)\n\n        return (seq_cls_out, out1, out2, out3, out4)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"model = MultiHeadedSequenceClassifier()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# outs = model(*learner.xb)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class MultiLoss(Module):\n    \n    def __init__(self, targ_pad_idx=666):\n        store_attr(\"targ_pad_idx\")\n        self.bce_loss = nn.BCEWithLogitsLoss()\n        self.ce_loss = CrossEntropyLossFlat()\n\n    def forward(self, inp, yb0, yb1):\n\n        seq_cls_out, out1, out2, out3, out4 = inp\n\n        mask = yb0 != self.targ_pad_idx\n        loss0 = self.bce_loss(seq_cls_out[mask], yb0[mask].float())\n\n\n        # loss 1 p/n/i\n        loss1 = self.ce_loss(out1, torch.where(yb1[:,:3])[1])\n\n        # loss 2 rv/lv/neither\n        loss2 = self.ce_loss(out2, torch.where(yb1[:,3:6])[1])\n\n        # loss 3 L/R/C\n        loss3 = self.bce_loss(out3, yb1[:,6:9].float())\n\n        # loss 3 chro/acute/neither\n        loss4 = self.ce_loss(out4, torch.where(yb1[:,9:])[1])\n\n        \n        return (loss0 + loss1 + loss2 + loss3 + loss4) / 5","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"loss_func = MultiLoss()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# loss = loss_func(outs, *learner.yb); loss","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"### Metric"},{"metadata":{"trusted":true},"cell_type":"code","source":"# 'POSITIVE','negative_exam_for_pe','indeterminate',\n# 'rv_lv_ratio_gte_1','rv_lv_ratio_lt_1', 'NEITHER'\n# 'leftsided_pe','rightsided_pe','central_pe',\n# 'chronic_pe','acute_and_chronic_pe','NEITHER'","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"bce_loss = BCEWithLogitsLossFlat()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# seq_cls_out, out1, out2, out3, out4 = outs\n# yb0, yb1 = learner.yb","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"neg_pe_wgt = 0.0736196319\nindeterminate_wgt = 0.09202453988\n\nrv_lv_gte_1_wgt = 0.2346625767\nrv_lv_lt_1_wgt = 0.0782208589\n\nleft_pe_wgt = 0.06257668712\nright_pe_wgt = 0.06257668712\ncentral_pe_wgt = 0.1877300613\n\nchronic_wgt = 0.1042944785\nacute_chronic_wgt = 0.1042944785\n\nbce_loss = BCEWithLogitsLossFlat()\n\ndef metric(preds, yb0, yb1):\n    \n    seq_cls_out, out1, out2, out3, out4 = preds\n    \n    bs = out1.size(0)\n    \n    out1 = F.softmax(out1, 1)\n    out2 = F.softmax(out2, 1)\n    out3 = torch.sigmoid(out3)\n    out4 = F.softmax(out4, 1)\n    \n    neg_pe_loss = F.binary_cross_entropy(out1[:,1], yb1[:,1].float())\n    indeterminate_loss = F.binary_cross_entropy(out1[:,2], yb1[:,2].float())\n\n    rv_lv_gte_1_loss = F.binary_cross_entropy(out2[:,0], yb1[:,3].float())\n    rv_lv_lt_1_loss = F.binary_cross_entropy(out2[:,1], yb1[:,4].float())\n\n    left_pe_wgt = F.binary_cross_entropy(out3[:,0], yb1[:,6].float())\n    right_pe_wgt = F.binary_cross_entropy(out3[:,0], yb1[:,7].float())\n    central_pe_wgt = F.binary_cross_entropy(out3[:,0], yb1[:,8].float())\n\n    chronic_loss = F.binary_cross_entropy(out4[:,0], yb1[:,9].float())\n    acute_chronic_loss = F.binary_cross_entropy(out4[:,1], yb1[:,10].float())\n\n    \n    tot_exam_loss = 0\n    tot_exam_wgts = 0\n\n    tot_exam_loss += neg_pe_loss*bs*neg_pe_wgt\n    tot_exam_loss += indeterminate_loss*bs*indeterminate_wgt\n    tot_exam_loss += rv_lv_gte_1_loss*bs*rv_lv_gte_1_wgt\n    tot_exam_loss += rv_lv_lt_1_loss*bs*rv_lv_lt_1_wgt\n    tot_exam_loss += left_pe_wgt*bs*left_pe_wgt\n    tot_exam_loss += right_pe_wgt*bs*right_pe_wgt\n    tot_exam_loss += central_pe_wgt*bs*central_pe_wgt\n    tot_exam_loss += chronic_loss*bs*chronic_wgt\n    tot_exam_loss += acute_chronic_loss*bs*acute_chronic_wgt\n\n    tot_exam_wgts += bs*neg_pe_wgt\n    tot_exam_wgts += bs*indeterminate_wgt\n    tot_exam_wgts += bs*rv_lv_gte_1_wgt\n    tot_exam_wgts += bs*rv_lv_lt_1_wgt\n    tot_exam_wgts += bs*left_pe_wgt\n    tot_exam_wgts += bs*right_pe_wgt\n    tot_exam_wgts += bs*central_pe_wgt\n    tot_exam_wgts += bs*chronic_wgt\n    tot_exam_wgts += bs*acute_chronic_wgt\n\n    tot_exam_loss, tot_exam_wgts = tot_exam_loss.item(), tot_exam_wgts.item()\n    \n    \n    \n    # Image-level weighted log loss (single batch)\n    w_img = 0.07361963\n    tot_img_loss = 0\n    tot_img_wgts = 0\n    for img_preds, img_targs in zip(seq_cls_out, yb0):\n        mask = img_targs != targ_pad_idx\n        img_preds = img_preds[mask]\n        img_targs = img_targs[mask]\n\n        n_imgs = sum(mask)\n\n        qi = img_targs.float().mean()    \n        img_loss = bce_loss(img_preds, img_targs)    \n        wgt = w_img*qi\n        tot_img_wgts += (wgt).item()*n_imgs\n        tot_img_loss += (wgt*img_loss).item()*n_imgs\n\n    \n    return (tot_exam_loss + tot_img_loss) / (tot_exam_wgts + tot_img_wgts)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"markdown","source":"### Train"},{"metadata":{"trusted":true},"cell_type":"code","source":"data = DataBlock(blocks=(SequenceBlock,SequenceTargetBlock,TargetBlock), \n                 n_inp=1, splitter=FuncSplitter(lambda o: o in valid_pids))\ndls = data.dataloaders(list(train_pids)+list(valid_pids), bs=64)\nmodel = MultiHeadedSequenceClassifier(dim=1024)\nloss_func = MultiLoss()\nlearner = Learner(dls, model, loss_func=loss_func, metrics=[metric], cbs=[SaveModelCallback(fname=\"best_seqmodel\")])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"learner.validate()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"learner.lr_find()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"learner.fit_flat_cos(20, lr=0.01)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"learner.save(\"init_seq_model\")","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"learner.export(\"best_seqmodel_export.pkl\")","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","collapsed":true,"_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","trusted":false},"cell_type":"code","source":"","execution_count":null,"outputs":[]}],"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":4,"nbformat_minor":4}