{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":8983556,"sourceType":"datasetVersion","datasetId":5405382}],"dockerImageVersionId":30746,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"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)\nimport matplotlib as mpl\nfrom matplotlib import pyplot as plt\nimport matplotlib.animation as anime\nfrom IPython.display import HTML,display\nimport os\nimport tqdm\nroot=\"/kaggle/input/rsna-train-images-256x256-pfgk/\"\nrooot=\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/\"\npd.set_option('future.no_silent_downcasting', True)\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","execution":{"iopub.status.busy":"2024-07-19T12:38:29.382913Z","iopub.execute_input":"2024-07-19T12:38:29.383336Z","iopub.status.idle":"2024-07-19T12:38:30.528004Z","shell.execute_reply.started":"2024-07-19T12:38:29.383279Z","shell.execute_reply":"2024-07-19T12:38:30.526874Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"conds=[\n    \"Spinal Canal Stenosis\",\n    \"Left Neural Foraminal Narrowing\",\n    \"Right Neural Foraminal Narrowing\",\n    \"Left Subarticular Stenosis\",\n    \"Right Subarticular Stenosis\",\n]\nlevels=[\"L1/L2\",\"L2/L3\",\"L3/L4\",\"L4/L5\",\"L5/S1\"]\nseries=[\n    \"Axial T2\",\n    \"Sagittal T1\",\n    \"Sagittal T2/STIR\",\n]\ndiags=[\n    \"Normal/Mild\",\n    \"Moderate\",\n    \"Severe\",\n]","metadata":{"execution":{"iopub.status.busy":"2024-07-19T12:38:32.972808Z","iopub.execute_input":"2024-07-19T12:38:32.973892Z","iopub.status.idle":"2024-07-19T12:38:32.980470Z","shell.execute_reply.started":"2024-07-19T12:38:32.973841Z","shell.execute_reply":"2024-07-19T12:38:32.979170Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_meta=pd.read_csv(f\"{root}meta.csv\")\ndf_meta","metadata":{"execution":{"iopub.status.busy":"2024-07-19T12:38:35.847789Z","iopub.execute_input":"2024-07-19T12:38:35.848180Z","iopub.status.idle":"2024-07-19T12:38:36.040252Z","shell.execute_reply.started":"2024-07-19T12:38:35.848148Z","shell.execute_reply":"2024-07-19T12:38:36.039207Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_series=pd.read_csv(f\"{rooot}train_series_descriptions.csv\")\ndf_series","metadata":{"execution":{"iopub.status.busy":"2024-07-19T12:38:38.317078Z","iopub.execute_input":"2024-07-19T12:38:38.317531Z","iopub.status.idle":"2024-07-19T12:38:38.345010Z","shell.execute_reply.started":"2024-07-19T12:38:38.317499Z","shell.execute_reply":"2024-07-19T12:38:38.343960Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"n=3","metadata":{"execution":{"iopub.status.busy":"2024-07-19T12:38:41.331161Z","iopub.execute_input":"2024-07-19T12:38:41.331565Z","iopub.status.idle":"2024-07-19T12:38:41.336786Z","shell.execute_reply.started":"2024-07-19T12:38:41.331533Z","shell.execute_reply":"2024-07-19T12:38:41.335372Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"indexed_df_meta=df_meta.set_index([\"study_id\",\"series_id\"])\nlabels=[\"SCS\",\"LNFN\",\"RNFN\",\"LSS\",\"RSS\"]\nlabels={k:v for k,v in zip(conds,labels)}\ncolor={\"Normal/Mild\":\"#A0C0FF\",\"Moderate\":\"#70FF70\",\"Severe\":\"#FFFF50\"}\ndef plot(i):\n    ax.cla()\n    ax.imshow(arr[i,:,:],cmap=mpl.cm.magma,vmin=0,vmax=arr.max())\n    if i in i_diags:\n        for j,r in df_diags[df_diags[\"instance_number\"]==i].iterrows():\n            label=f\"{labels[r['condition']]}\\n{r['level']}\"\n            x,y=float(r[\"x\"]),float(r[\"y\"])\n            ax.scatter([x],[y],s=500,marker=\"o\",facecolors='none',edgecolors=color[r[\"diagnosis\"]])\n            ax.text(x+10,y,label,fontsize=7,c=\"white\",weight=\"bold\",)\n            \n\nfor serie in series:\n    print(serie)\n    df_serie=df_series[df_series[\"series_description\"]==serie].sample(n=3)\n    for i,row in df_serie.iterrows():\n        fn=f\"{root}images/{row['study_id']}/{row['series_id']}.dat\"\n        df_diags=indexed_df_meta.loc[row['study_id'],row['series_id']]\n        i_diags=set(df_diags[\"instance_number\"])\n        arr=np.fromfile(fn,\"uint16\").reshape(-1,256,256)\n        fig=plt.figure()\n        ax=plt.axes()\n        anim=anime.FuncAnimation(fig,plot,frames=len(arr))\n        plt.close(fig)\n        display(HTML(anim.to_jshtml(fps=5)))\n        ","metadata":{"execution":{"iopub.status.busy":"2024-07-19T12:38:43.585977Z","iopub.execute_input":"2024-07-19T12:38:43.586385Z","iopub.status.idle":"2024-07-19T12:38:58.770755Z","shell.execute_reply.started":"2024-07-19T12:38:43.586350Z","shell.execute_reply":"2024-07-19T12:38:58.768566Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"winsize=16\nw,h=4,4","metadata":{"execution":{"iopub.status.busy":"2024-07-19T12:39:11.524457Z","iopub.execute_input":"2024-07-19T12:39:11.524831Z","iopub.status.idle":"2024-07-19T12:39:11.530001Z","shell.execute_reply.started":"2024-07-19T12:39:11.524805Z","shell.execute_reply":"2024-07-19T12:39:11.528834Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cond=\"Spinal Canal Stenosis\"\nfor diag in diags:\n    df_diag=df_meta[(df_meta[\"condition\"]==cond)&(df_meta[\"diagnosis\"]==diag)].sample(n=w*h)\n    fig,axs=plt.subplots(w,h)\n    plt.suptitle(f\"{cond}/{diag}\")\n    axs.shape=-1\n    for i in range(16):\n        img=np.zeros((winsize*2,winsize*2))\n        row=df_diag.iloc[i]\n        fn=f\"{root}images/{row['study_id']}/{int(row['series_id'])}.dat\"\n        arr=np.fromfile(fn,\"uint16\").reshape(-1,256,256)\n        arr=arr[row[\"instance_number\"]]\n        x0,x1=max(0,round(row[\"x\"])-winsize),min(round(row[\"x\"])+winsize,255)\n        y0,y1=max(0,round(row[\"y\"])-winsize),min(round(row[\"y\"])+winsize,255)\n        img[:y1-y0,:x1-x0]=arr[y0:y1,x0:x1]\n        axs[i].imshow(img/img.max(),cmap=mpl.cm.magma)\n        axs[i].set_title(row[\"level\"])\n        axs[i].set_xticks([])\n        axs[i].set_yticks([])\n    plt.show()\n","metadata":{"execution":{"iopub.status.busy":"2024-07-19T08:40:04.370696Z","iopub.status.idle":"2024-07-19T08:40:04.371259Z","shell.execute_reply.started":"2024-07-19T08:40:04.371022Z","shell.execute_reply":"2024-07-19T08:40:04.371042Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for cond in conds[1:3]:\n    for diag in diags:\n        df_diag=df_meta[(df_meta[\"condition\"]==cond)&(df_meta[\"diagnosis\"]==diag)].sample(n=w*h)\n        fig,axs=plt.subplots(w,h)\n        plt.suptitle(f\"{cond}/{diag}\")\n        axs.shape=-1\n        for i in range(16):\n            img=np.zeros((winsize*2,winsize*2))\n            row=df_diag.iloc[i]\n            fn=f\"{root}images/{row['study_id']}/{int(row['series_id'])}.dat\"\n            arr=np.fromfile(fn,\"uint16\").reshape(-1,256,256)\n#             arr=arr[row[\"instance_number\"]-1:row[\"instance_number\"]+2]\n#             arr=arr.transpose(1,2,0)\n            arr=arr[row[\"instance_number\"]]\n            x0,x1=max(0,round(row[\"x\"])-winsize),min(round(row[\"x\"])+winsize,255)\n            y0,y1=max(0,round(row[\"y\"])-winsize),min(round(row[\"y\"])+winsize,255)\n            img[:y1-y0,:x1-x0]=arr[y0:y1,x0:x1]\n            axs[i].imshow(img/img.max(),cmap=mpl.cm.magma)\n            axs[i].set_title(row[\"level\"])\n            axs[i].set_xticks([])\n            axs[i].set_yticks([])\n        plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-07-19T08:40:04.373649Z","iopub.status.idle":"2024-07-19T08:40:04.374206Z","shell.execute_reply.started":"2024-07-19T08:40:04.373967Z","shell.execute_reply":"2024-07-19T08:40:04.373987Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for cond in conds[3:]:\n    for diag in diags:\n        df_diag=df_meta[(df_meta[\"condition\"]==cond)&(df_meta[\"diagnosis\"]==diag)].sample(n=w*h)\n        fig,axs=plt.subplots(w,h)\n        plt.suptitle(f\"{cond}/{diag}\")\n        axs.shape=-1\n        for i in range(16):\n            img=np.zeros((winsize*2,winsize*2))\n            row=df_diag.iloc[i]\n            fn=f\"{root}images/{row['study_id']}/{int(row['series_id'])}.dat\"\n            arr=np.fromfile(fn,\"uint16\").reshape(-1,256,256)\n            arr=arr[row[\"instance_number\"]]\n            x0,x1=max(0,round(row[\"x\"])-winsize),min(round(row[\"x\"])+winsize,255)\n            y0,y1=max(0,round(row[\"y\"])-winsize),min(round(row[\"y\"])+winsize,255)\n            img[:y1-y0,:x1-x0]=arr[y0:y1,x0:x1]\n            axs[i].imshow(img/img.max(),cmap=mpl.cm.magma)\n            axs[i].set_title(row[\"level\"])\n            axs[i].set_xticks([])\n            axs[i].set_yticks([])\n        plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-07-19T08:40:04.376294Z","iopub.status.idle":"2024-07-19T08:40:04.376749Z","shell.execute_reply.started":"2024-07-19T08:40:04.376541Z","shell.execute_reply":"2024-07-19T08:40:04.376558Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"input_arr=[]\nlabel_arr=[]\ndf_scs=df_meta[df_meta[\"condition\"]==conds[0]]\nfor i,row in tqdm.tqdm(df_scs.iterrows()):\n    img=np.zeros((1,winsize*2,winsize*2))\n    fn=f\"{root}images/{row['study_id']}/{int(row['series_id'])}.dat\"\n    arr=np.fromfile(fn,\"uint16\").reshape(-1,256,256)\n    arr=arr[row[\"instance_number\"]]\n    x0,x1=max(0,round(row[\"x\"])-winsize),min(round(row[\"x\"])+winsize,255)\n    y0,y1=max(0,round(row[\"y\"])-winsize),min(round(row[\"y\"])+winsize,255)\n    img[0,:y1-y0,:x1-x0]=arr[y0:y1,x0:x1]\n    input_arr.append(img/float(img.max()+0.0001))\n    if row[\"diagnosis\"]==\"Normal/Mild\":\n        label_arr.append([1,0,0])\n    elif row[\"diagnosis\"]==\"Moderate\":\n        label_arr.append([0,1,0])\n    else:\n        label_arr.append([0,0,1])\ninput_arr=np.array(input_arr)\nlabel_arr=np.array(label_arr)    ","metadata":{"execution":{"iopub.status.busy":"2024-07-19T12:39:16.986636Z","iopub.execute_input":"2024-07-19T12:39:16.987050Z","iopub.status.idle":"2024-07-19T12:40:23.235038Z","shell.execute_reply.started":"2024-07-19T12:39:16.987017Z","shell.execute_reply":"2024-07-19T12:40:23.233906Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchvision\nfrom torch.utils.data import DataLoader,random_split,TensorDataset","metadata":{"execution":{"iopub.status.busy":"2024-07-19T08:40:04.382594Z","iopub.status.idle":"2024-07-19T08:40:04.383236Z","shell.execute_reply.started":"2024-07-19T08:40:04.382922Z","shell.execute_reply":"2024-07-19T08:40:04.382949Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_ratio=0.8\nbatch_size=1000\n\"\"\"\nバッチごとのサンプル数と学習データの比率を指定してDataLoaderを作成する\n\"\"\"\n#make torch.tensor object from numpy.ndarray object\ninput_tensor=torch.from_numpy(input_arr.astype(\"float64\")).clone()\nlabel_tensor=torch.from_numpy(label_arr.astype(\"float64\")).clone()\n\n#make torch TensorDataset object from torch.tensor object\nfull_dataset=TensorDataset(input_tensor,label_tensor)\n\n#calculate dataset size for each \ntrain_size=int(train_ratio*len(full_dataset))\ntest_size=len(full_dataset)-train_size\n\n#split to train and test randomly\ntrain_dataset,test_dataset=random_split(full_dataset,[train_size,test_size])\n\n#make dataloaders\ntrain_dataloader=DataLoader(dataset=train_dataset,batch_size=batch_size,shuffle=True)\ntest_dataloader=DataLoader(dataset=test_dataset,batch_size=test_size,shuffle=True)    ","metadata":{"execution":{"iopub.status.busy":"2024-07-19T08:40:04.386122Z","iopub.status.idle":"2024-07-19T08:40:04.386747Z","shell.execute_reply.started":"2024-07-19T08:40:04.386480Z","shell.execute_reply":"2024-07-19T08:40:04.386514Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class EarlyStopping:\n    def __init__(self,patience=7,verbose=False,delta=0,path='checkpoint.pt',trace_func=print,score_or_loss=\"score\"):\n        self.patience=patience\n        self.verbose=verbose\n        self.counter=0\n        self.best_score=None\n        self.early_stop=False\n        self.delta=delta\n        self.path=path\n        self.trace_func=trace_func\n        self.score_or_loss=score_or_loss\n        \n    def __call__(self,eval,model):\n\n        if self.score_or_loss==\"loss\":score=-eval\n        else:score=eval\n\n        if self.best_score is None:\n            self.best_score=score\n            self.save_checkpoint(eval,model)\n        elif score<self.best_score+self.delta:\n            self.counter+=1\n            self.trace_func(f'EarlyStopping counter:{self.counter} out of{self.patience}')\n            if self.counter>=self.patience:\n                self.early_stop=True\n        else:\n            self.save_checkpoint(eval,model)\n            self.best_score=score\n            self.counter=0\n\n    def save_checkpoint(self,eval,model):\n        if self.verbose:\n            self.trace_func(\n                f'Validation loss decreased ({-self.best_score:.6f}',\n                            f' --> {-eval:.6f}).  Saving model ...')\n        torch.save(model.state_dict(),self.path)\n        \n    def load_checkpoint(self,model):\n        model.load_state_dict(torch.load(self.path))","metadata":{"execution":{"iopub.status.busy":"2024-07-19T08:40:04.389105Z","iopub.status.idle":"2024-07-19T08:40:04.389609Z","shell.execute_reply.started":"2024-07-19T08:40:04.389389Z","shell.execute_reply":"2024-07-19T08:40:04.389407Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Trainer:\n    \n    def __init__(self,net,evaluation):\n        \"\"\"\n        eval:callable evaluation fuction \n        arguments:y...np.ndarray shape=(num samples,num components)\n                  t...np.ndarray shape=(num samples,num components)\n        return:sum of evaluation\n        \"\"\"\n        self.net=net\n        self.eval=evaluation#\n        \n        \n    def define_process(self,loss_func,learning_rate,optimizer,plot=None):       \n        \"\"\"\n        plot:plot function\n        arguments:train_loss list[float]\n                  train_eval list[float]\n                  test_loss list[float]\n                  test_eval list[float]\n                  -->None\n        \"\"\"\n\n        self.device=torch.device('cuda:0' if torch.cuda.is_available() else 'cpu')\n        self.net=self.net.to(device=self.device)\n        \n        self.optimizer=optimizer(self.net.parameters(),lr=learning_rate)\n        self.criterion=loss_func\n        \n        print(\"network is defined\")\n        \n        self.plot=plot\n        \n    def train(self,dataloader):\n        self.net.train()\n        train_loss=0.0\n        train_eval=0.0\n        cnt=0\n        nsmpl=0\n        for inputs,labels in dataloader:\n            cnt+=1\n            nsmpl+=inputs.shape[0]\n            \n            inputs=inputs.to(self.device)\n            labels=labels.to(self.device)\n            labels=labels.float().to(self.device)\n            \n            self.optimizer.zero_grad()\n            outputs=self.net(inputs)\n            outputs=outputs.squeeze()\n\n            y=outputs.detach().cpu().numpy()\n            t=labels.detach().cpu().numpy()\n\n            train_eval+=self.eval(y,t)\n            \n            loss=self.criterion(outputs,labels)\n            loss.backward()\n            self.optimizer.step()\n            \n            train_loss+=loss.detach().cpu().numpy()*inputs.shape[0]\n            \n            del loss\n            \n            if cnt%1==0:\n                print(\n                    f\"train_loss:{train_loss/nsmpl},\",\n                    f\"train_eval:{train_eval/nsmpl},\",\n                    f\"batch_cnt:{cnt}/{len(dataloader)}\"\n                )\n        return train_loss/nsmpl,train_eval/nsmpl\n\n    def test(self,dataloader):\n        self.net.eval()\n        test_loss=0.0\n        test_eval=0.0\n        nsmpl=0\n        with torch.no_grad():\n            for inputs,labels in dataloader:\n                nsmpl+=inputs.shape[0]\n                inputs=inputs.to(self.device)\n                labels=labels.to(self.device)\n                labels=labels.float().to(self.device)\n                \n                outputs=self.net(inputs)\n                outputs=outputs.squeeze()\n\n                y=outputs.detach().cpu().numpy()\n                t=labels.detach().cpu().numpy()\n                \n                test_eval+=self.eval(y,t)\n                \n                loss=self.criterion(outputs,labels)\n                \n                test_loss+=loss.detach().cpu().numpy()*inputs.shape[0]\n                \n        print(\n            \"---------------------------------------------\",\n            f\"\\ntest_loss:{test_loss/nsmpl},test_eval:{test_eval/nsmpl}\"\n        )\n        return test_loss/nsmpl,test_eval/nsmpl\n    \n    def learn(self,dataset,num_epochs,early_stopping):\n        print(\"start training\")\n        train_loss_list=[]\n        test_loss_list=[]\n        train_eval_list=[]\n        test_eval_list=[]\n        \n        for epoch in range(num_epochs):\n            print(\n                \"=============================================\",\n                f\"\\n epoch:{epoch}/{num_epochs}\"\n            )\n            \n            train_loss,train_eval=self.train(dataset[0])\n            train_loss_list.append(train_loss)\n            train_eval_list.append(train_eval)\n            \n            test_loss,test_eval=self.test(dataset[1])\n            test_loss_list.append(test_loss)\n            test_eval_list.append(test_eval)\n            \n            if self.plot is not None:self.plot(train_loss_list,test_loss_list,train_eval_list,test_eval_list)\n            \n            early_stopping(test_eval,self.net)\n            if early_stopping.early_stop:\n                print(\"Early Stopping\")\n                break\n            print(\n                \"=============================================\",\n            )\n            \n        early_stopping.load_checkpoint(self.net)\n            \n        return train_loss_list,test_loss_list,train_eval_list,test_eval_list","metadata":{"execution":{"iopub.status.busy":"2024-07-19T11:22:17.283284Z","iopub.execute_input":"2024-07-19T11:22:17.284843Z","iopub.status.idle":"2024-07-19T11:22:17.312439Z","shell.execute_reply.started":"2024-07-19T11:22:17.284759Z","shell.execute_reply":"2024-07-19T11:22:17.310909Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"loss_func=nn.CrossEntropyLoss()\n\ndef evaluate(y,t):\n    return (y.argmax(axis=1)==t.argmax(axis=1)).sum()\n\nclass Net(nn.Module):\n    def __init__(self):\n        w_init_mean=0\n        w_init_std=0.5\n        b_init_mean=0\n        b_init_std=0.5\n        super().__init__()\n        self.conv1=nn.Conv2d(1,10,5,padding=2,dtype=torch.float64)\n        self.conv2=nn.Conv2d(10,10,5,padding=2,dtype=torch.float64)\n        self.pool1=nn.MaxPool2d(2,2)\n        self.conv3=nn.Conv2d(10,10,3,padding=1,dtype=torch.float64)\n        self.conv4=nn.Conv2d(10,10,3,padding=1,dtype=torch.float64)\n        self.conv5=nn.Conv2d(10,2,3,padding=1,dtype=torch.float64)\n        self.pool2=nn.MaxPool2d(2,2)\n        self.fc1=nn.Linear(128,50,dtype=torch.float64)\n        self.fc2=nn.Linear(50,12,dtype=torch.float64)\n        self.fc3=nn.Linear(12,3,dtype=torch.float64)\n        nn.init.normal_(self.fc1.weight,mean=w_init_mean,std=w_init_std)\n        nn.init.normal_(self.fc1.bias,mean=b_init_mean,std=b_init_std)\n        nn.init.normal_(self.fc2.weight,mean=w_init_mean,std=w_init_std)\n        nn.init.normal_(self.fc2.bias,mean=b_init_mean,std=b_init_std)\n        nn.init.normal_(self.fc3.weight,mean=w_init_mean,std=w_init_std)\n        nn.init.normal_(self.fc3.bias,mean=b_init_mean,std=b_init_std)\n    def forward(self,x):\n        x=F.silu(self.conv1(x))\n        x=F.silu(self.conv2(x))\n        x=_x=self.pool1(x)\n        x=F.silu(self.conv3(x))\n        x=F.silu(self.conv4(x))+_x\n        x=F.silu(self.conv5(x))\n        x=self.pool2(x)\n        x=torch.flatten(x,1)\n        x=F.silu(self.fc1(x))\n        x=F.silu(self.fc2(x))\n        x=torch.tanh(self.fc3(x))\n        return x","metadata":{"execution":{"iopub.status.busy":"2024-07-19T08:40:04.393935Z","iopub.status.idle":"2024-07-19T08:40:04.394337Z","shell.execute_reply.started":"2024-07-19T08:40:04.394146Z","shell.execute_reply":"2024-07-19T08:40:04.394162Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot(trl,tel,tre,tee):\n    plt.plot(trl,label=\"train\")\n    plt.plot(tel,label=\"test\")\n    plt.title(\"loss\")\n    plt.legend()\n    plt.show()\n    plt.plot(tre,label=\"train\")\n    plt.plot(tee,label=\"test\")\n    plt.title(\"eval\")\n    plt.legend()\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-07-19T11:24:06.321496Z","iopub.execute_input":"2024-07-19T11:24:06.321990Z","iopub.status.idle":"2024-07-19T11:24:06.330574Z","shell.execute_reply.started":"2024-07-19T11:24:06.321956Z","shell.execute_reply":"2024-07-19T11:24:06.328857Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"net=Net()\ntrainer=Trainer(net,evaluate)\ntrainer.define_process(loss_func,0.001,torch.optim.Adam,plot)\ntrainer.learn((train_dataloader,test_dataloader),20,EarlyStopping())","metadata":{"execution":{"iopub.status.busy":"2024-07-19T08:40:04.398756Z","iopub.status.idle":"2024-07-19T08:40:04.399255Z","shell.execute_reply.started":"2024-07-19T08:40:04.399019Z","shell.execute_reply":"2024-07-19T08:40:04.399036Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"input_arr2=np.array([]).reshape(0,1,winsize*2,winsize*2)\nlabel_arr2=np.array([]).reshape(0,2)\ndf_st2=df_series[df_series[\"series_description\"]==series[2]]\nidxd_df_meta=df_meta.set_index([\"study_id\",\"series_id\"])    \ndef make_input(crd,arr):\n    im=arr[crd[0],crd[1]:crd[1]+2*winsize,crd[2]:crd[2]+2*winsize]\n    return np.tile(im/(im.max()+0.001),(1,1,1))\ndef make_label(crd,dcrds):\n    crd[1:]+=winsize\n    return np.exp(-((dcrds-crd)**2).sum(axis=1)/6).sum()\nfor i,row in tqdm.tqdm(df_st2.iterrows()):\n    try:dcrds=np.array([[r[\"instance_number\"],r[\"x\"],r[\"y\"]] for i,r in idxd_df_meta.loc[row[\"study_id\"],row[\"series_id\"]].iterrows()])\n    except:continue\n    if len(dcrds)<5:continue\n    fn=f\"{root}images/{row['study_id']}/{int(row['series_id'])}.dat\"\n    arr=np.fromfile(fn,\"uint16\").reshape(-1,256,256)\n    vlsz=256-winsize*2\n    crds=np.random.choice(len(arr)*vlsz*2,size=10,replace=False)\n    crds=np.array([(crds/vlsz**2).astype(\"int\"),((crds%vlsz**2)/vlsz).astype(\"int\"),(crds%vlsz).astype(\"int\")]).T\n    input=np.apply_along_axis(lambda crd:make_input(crd,arr),axis=1,arr=crds)\n    label=np.apply_along_axis(lambda crd:make_label(crd,dcrds),axis=1,arr=crds)\n    label=np.array([label,1-label]).T\n#     print(label.shape)\n    input_arr2=np.append(input_arr2,input,axis=0)\n    label_arr2=np.append(label_arr2,label,axis=0)\ninput_arr=np.append(input_arr,input_arr2,axis=0)\nlabel_arr=np.append(np.tile(np.arange(1,-1,-1),(len(label_arr),1)),label_arr2,axis=0)","metadata":{"execution":{"iopub.status.busy":"2024-07-19T12:40:44.648493Z","iopub.execute_input":"2024-07-19T12:40:44.649547Z","iopub.status.idle":"2024-07-19T12:41:49.063608Z","shell.execute_reply.started":"2024-07-19T12:40:44.649496Z","shell.execute_reply":"2024-07-19T12:41:49.062465Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_ratio=0.8\nbatch_size=1000\n\"\"\"\nバッチごとのサンプル数と学習データの比率を指定してDataLoaderを作成する\n\"\"\"\n#make torch.tensor object from numpy.ndarray object\ninput_tensor=torch.from_numpy(input_arr.astype(\"float64\")).clone()\nlabel_tensor=torch.from_numpy(label_arr.astype(\"float64\")).clone()\n\n#make torch TensorDataset object from torch.tensor object\nfull_dataset=TensorDataset(input_tensor,label_tensor)\n\n#calculate dataset size for each \ntrain_size=int(train_ratio*len(full_dataset))\ntest_size=len(full_dataset)-train_size\n\n#split to train and test randomly\ntrain_dataset,test_dataset=random_split(full_dataset,[train_size,test_size])\n\n#make dataloaders\ntrain_dataloader=DataLoader(dataset=train_dataset,batch_size=batch_size,shuffle=True)\ntest_dataloader=DataLoader(dataset=test_dataset,batch_size=test_size,shuffle=True)    ","metadata":{"execution":{"iopub.status.busy":"2024-07-19T11:18:41.596888Z","iopub.execute_input":"2024-07-19T11:18:41.597340Z","iopub.status.idle":"2024-07-19T11:18:41.786459Z","shell.execute_reply.started":"2024-07-19T11:18:41.597306Z","shell.execute_reply":"2024-07-19T11:18:41.785155Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"loss_func=nn.CrossEntropyLoss()\n\ndef evaluate(y,t):\n    return (y.argmax(axis=1)==t.argmax(axis=1)).sum()\n\nclass Net(nn.Module):\n    def __init__(self):\n        w_init_mean=0\n        w_init_std=0.5\n        b_init_mean=0\n        b_init_std=0.5\n        super().__init__()\n        self.conv1=nn.Conv2d(1,10,5,padding=2,dtype=torch.float64)\n        self.conv2=nn.Conv2d(10,10,5,padding=2,dtype=torch.float64)\n        self.pool1=nn.MaxPool2d(2,2)\n        self.conv3=nn.Conv2d(10,10,3,padding=1,dtype=torch.float64)\n        self.conv4=nn.Conv2d(10,10,3,padding=1,dtype=torch.float64)\n        self.conv5=nn.Conv2d(10,2,3,padding=1,dtype=torch.float64)\n        self.pool2=nn.MaxPool2d(2,2)\n        self.fc1=nn.Linear(128,50,dtype=torch.float64)\n        self.fc2=nn.Linear(50,12,dtype=torch.float64)\n        self.fc3=nn.Linear(12,2,dtype=torch.float64)\n        nn.init.normal_(self.fc1.weight,mean=w_init_mean,std=w_init_std)\n        nn.init.normal_(self.fc1.bias,mean=b_init_mean,std=b_init_std)\n        nn.init.normal_(self.fc2.weight,mean=w_init_mean,std=w_init_std)\n        nn.init.normal_(self.fc2.bias,mean=b_init_mean,std=b_init_std)\n        nn.init.normal_(self.fc3.weight,mean=w_init_mean,std=w_init_std)\n        nn.init.normal_(self.fc3.bias,mean=b_init_mean,std=b_init_std)\n    def forward(self,x):\n        x=F.silu(self.conv1(x))\n        x=F.silu(self.conv2(x))\n        x=_x=self.pool1(x)\n        x=F.silu(self.conv3(x))\n        x=F.silu(self.conv4(x))+_x\n        x=F.silu(self.conv5(x))\n        x=self.pool2(x)\n        x=torch.flatten(x,1)\n        x=F.silu(self.fc1(x))\n        x=F.silu(self.fc2(x))\n        x=torch.tanh(self.fc3(x))\n        return x","metadata":{"execution":{"iopub.status.busy":"2024-07-19T11:22:22.095996Z","iopub.execute_input":"2024-07-19T11:22:22.096444Z","iopub.status.idle":"2024-07-19T11:22:22.115857Z","shell.execute_reply.started":"2024-07-19T11:22:22.096413Z","shell.execute_reply":"2024-07-19T11:22:22.114048Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"net=Net()\ntrainer=Trainer(net,evaluate)\ntrainer.define_process(loss_func,0.001,torch.optim.Adam,plot)\ntrainer.learn((train_dataloader,test_dataloader),20,EarlyStopping())","metadata":{"execution":{"iopub.status.busy":"2024-07-19T11:24:13.629550Z","iopub.execute_input":"2024-07-19T11:24:13.630010Z","iopub.status.idle":"2024-07-19T11:25:45.571892Z","shell.execute_reply.started":"2024-07-19T11:24:13.629977Z","shell.execute_reply":"2024-07-19T11:25:45.569868Z"},"trusted":true},"execution_count":null,"outputs":[]}]}