{"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":"gpu","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":9126231,"sourceType":"datasetVersion","datasetId":5480747},{"sourceId":201026335,"sourceType":"kernelVersion"}],"dockerImageVersionId":30747,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install '/kaggle/input/rsna2024-demo-workflow/natsort-8.4.0-py3-none-any.whl'","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys, os\nsys.path.append('/kaggle/input/rsna2024-demo-workflow')\n\nfrom _dir_setting_ import *\nprint('NOT_KAGGLE:', NOT_KAGGLE)\nprint('DATA_KAGGLE_DIR:', DATA_KAGGLE_DIR)\n\nimport matplotlib\nimport matplotlib.pyplot as plt\n\nfrom helper import *\nfrom data import *\nfrom model import *\n\nprint('import ok!')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# STEP0 : SETUP DATA ==========================\n\ncfg=dotdict(\n    point_net=dotdict(\n        checkpoint='/kaggle/input/rsna2024-demo-workflow/00002484.pth',\n        image_size=160,\n    ),\n)\n\nlevel_to_label={\n    'l1_l2':1,\n    'l2_l3':2,\n    'l3_l4':3,\n    'l4_l5':4,\n    'l5_s1':5,\n}\n\n#################################################\n\n# study id used for demo\nid_df = pd.read_csv(f'{DATA_KAGGLE_DIR}/train_series_descriptions.csv')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Basic flow:\n- We use affine transformation for projection of dicom array's 2d points into 3d world corrdinates. See our previous notebook: https://www.kaggle.com/code/hengck23/2d-to-3d-projection-for-dicom\n\n- Given sagittal volume, choose central slice z.\n- Use a 2d unet to predict the (x,y) coordinates of the 5 level key points (spinal_canal_stenosis label points)\n- Project x,y,z inpto word coordinates xx,yy,zz.\n- Now given the axial volume, compute the level (l1_l2, ..., l5_s1) for each slice, using their distances from xx,yy,zz points.\n\n","metadata":{}},{"cell_type":"code","source":"# STEP1 : DETECT KEYPOINT ==========================\n\n\npoint_net = Net(pretrained=False)\nf = torch.load(cfg.point_net.checkpoint, map_location=lambda storage, loc: storage)\nstate_dict = f['state_dict']\nprint(point_net.load_state_dict(state_dict, strict=False))\npoint_net.cuda()\npoint_net.eval()\npoint_net.output_type = ['infer']\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from _dir_setting_ import *\nfrom natsort import natsorted\nimport pandas as pd\npd.set_option('mode.chained_assignment', None) # disable SettingWithCopyWarning\n\nimport numpy as np\nimport pydicom\nimport glob\n\nimport cv2\n\nfrom matplotlib.patches import FancyArrowPatch\nfrom mpl_toolkits.mplot3d import proj3d\nfrom pdb import set_trace as st\n\nclass dotdict(dict):\n    __setattr__ = dict.__setitem__\n    __delattr__ = dict.__delitem__\n\n    def __getattr__(self, name):\n        try:\n            return self[name]\n        except KeyError:\n            raise AttributeError(name)\n\n################################################################################3\n\n\n# read into volume (np_array) + dicom_header(df)\n'''\nimport notes:\n- instance_number may not be sequentially (can have missing num)\n- instance_number is 1-indexed\n\n'''\n\n## 3d/2d processing #########################################################\ndef np_dot(a,b):\n    return np.sum(a * b, 1)\n\ndef project_to_3d(x,y,z, df):\n    d = df.iloc[z]\n    H, W = d.H, d.W\n    sx, sy, sz = [float(v) for v in d.ImagePositionPatient]\n    o0, o1, o2, o3, o4, o5, = [float(v) for v in d.ImageOrientationPatient]\n    delx, dely = d.PixelSpacing\n\n    xx = o0 * delx * x + o3 * dely * y + sx\n    yy = o1 * delx * x + o4 * dely * y + sy\n    zz = o2 * delx * x + o5 * dely * y + sz\n    return xx,yy,zz\n\n## read data #########################################################\ndef resize_volume(volume, image_size):\n    image = volume.copy()\n    image = np.ascontiguousarray(image.transpose((1, 2, 0)))\n    image = cv2.resize(image, (image_size, image_size), interpolation=cv2.INTER_LINEAR)\n    image = np.ascontiguousarray(image.transpose((2, 0, 1)))  # cv2.INTER_LINEAR=1\n    return image\n\ndef normalise_to_8bit(x, lower=0.1, upper=99.9): # 1, 99 #0.05, 99.5 #0, 100\n    lower, upper = np.percentile(x, (lower, upper))\n    x = np.clip(x, lower, upper)\n    x = x - np.min(x)\n    x = x / np.max(x)\n    return (x * 255).astype(np.uint8)\ndef read_series(study_id,series_id,series_description):\n    data_kaggle_dir = DATA_KAGGLE_DIR\n    dicom_dir = f'{data_kaggle_dir}/train_images/{study_id}/{series_id}'\n\n    # read dicom file\n    dicom_file = natsorted(glob.glob(f'{dicom_dir}/*.dcm'))\n    instance_number = [int(f.split('/')[-1].split('.')[0]) for f in dicom_file]\n    dicom = [pydicom.dcmread(f) for f in dicom_file]\n\n    # make dicom header df\n    H, W = dicom[0].pixel_array.shape\n    dicom_df = []\n    for i, d in zip(instance_number, dicom):  # d__.dict__\n        dicom_df.append(\n            dotdict(\n                study_id=study_id,\n                series_id=series_id,\n                series_description=series_description,\n                instance_number=i,\n                # InstanceNumber = d.InstanceNumber,\n                ImagePositionPatient=[float(v) for v in d.ImagePositionPatient],\n                ImageOrientationPatient=[float(v) for v in d.ImageOrientationPatient],\n                PixelSpacing=[float(v) for v in d.PixelSpacing],\n                SpacingBetweenSlices=float(d.SpacingBetweenSlices),\n                SliceThickness=float(d.SliceThickness),\n                grouping=str([round(float(v), 3) for v in d.ImageOrientationPatient]),\n                H=H,\n                W=W,\n            )\n        )\n    dicom_df = pd.DataFrame(dicom_df)\n\n    # sort slices\n    dicom_df = [d for _, d in dicom_df.groupby('grouping')]\n\n    data = []\n    sort_data_by_group = []\n    for df in dicom_df:\n        position = np.array(df['ImagePositionPatient'].values.tolist())\n        orientation = np.array(df['ImageOrientationPatient'].values.tolist())\n        normal = np.cross(orientation[:, :3], orientation[:, 3:])\n        projection = np_dot(normal, position)\n        df.loc[:, 'projection'] = projection\n        df = df.sort_values('projection')\n\n\n        # todo: assert all slices are continous ??\n        # use  (position[-1]-position[0])/N = SpacingBetweenSlices ??\n        assert len(df.SliceThickness.unique()) == 1\n        assert len(df.SpacingBetweenSlices.unique()) == 1\n\n\n        volume = [\n            dicom[instance_number.index(i)].pixel_array for i in df.instance_number\n        ]\n        volume = np.stack(volume)\n        volume = normalise_to_8bit(volume)\n        data.append(dotdict(\n            df=df,\n            volume=volume,\n        ))\n\n        if 'sagittal' in series_description.lower():\n            sort_data_by_group.append(position[0, 0])  # x\n        if 'axial' in series_description.lower():\n            sort_data_by_group.append(position[0, 2])  # z\n\n    data = [r for _, r in sorted(zip(sort_data_by_group, data))]\n    for i, r in enumerate(data):\n        r.df.loc[:, 'group'] = i\n\n    df = pd.concat([r.df for r in data])\n    df.loc[:, 'z'] = np.arange(len(df))\n    try:    \n        volume = np.concatenate([r.volume for r in data])\n    except:\n        volume = []\n        for r in data:\n            # st()\n            volume.append([cv2.resize(im, (608, 608)) for im in r.volume])\n        volume = np.concatenate(volume)\n    data = dotdict(\n        series_id=series_id,\n        df=df,\n        volume=volume,\n    )\n    return data\ndef read_axial_df(study_id,series_id,series_description):\n    data_kaggle_dir = DATA_KAGGLE_DIR\n    dicom_dir = f'{data_kaggle_dir}/train_images/{study_id}/{series_id}'\n\n    # read dicom file\n    dicom_file = natsorted(glob.glob(f'{dicom_dir}/*.dcm'))\n    instance_number = [int(f.split('/')[-1].split('.')[0]) for f in dicom_file]\n    dicom = [pydicom.dcmread(f) for f in dicom_file]\n\n    # make dicom header df\n    H, W = dicom[0].pixel_array.shape\n    dicom_df = []\n    for i, d in zip(instance_number, dicom):  # d__.dict__\n        dicom_df.append(\n            dotdict(\n                study_id=study_id,\n                series_id=series_id,\n                series_description=series_description,\n                instance_number=i,\n                # InstanceNumber = d.InstanceNumber,\n                ImagePositionPatient=[float(v) for v in d.ImagePositionPatient],\n                ImageOrientationPatient=[float(v) for v in d.ImageOrientationPatient],\n                PixelSpacing=[float(v) for v in d.PixelSpacing],\n                SpacingBetweenSlices=float(d.SpacingBetweenSlices),\n                SliceThickness=float(d.SliceThickness),\n                grouping=str([round(float(v), 3) for v in d.ImageOrientationPatient]),\n                H=H,\n                W=W,\n            )\n        )\n    dicom_df = pd.DataFrame(dicom_df)\n\n    # sort slices\n    dicom_df = [d for _, d in dicom_df.groupby('grouping')]\n\n    data = []\n    sort_data_by_group = []\n    for df in dicom_df:\n        position = np.array(df['ImagePositionPatient'].values.tolist())\n        orientation = np.array(df['ImageOrientationPatient'].values.tolist())\n        normal = np.cross(orientation[:, :3], orientation[:, 3:])\n        projection = np_dot(normal, position)\n        df.loc[:, 'projection'] = projection\n        df = df.sort_values('projection')\n\n\n        # todo: assert all slices are continous ??\n        # use  (position[-1]-position[0])/N = SpacingBetweenSlices ??\n        assert len(df.SliceThickness.unique()) == 1\n        if len(df.SpacingBetweenSlices.unique()) != 1:\n            import pdb;pdb.set_trace()\n        assert len(df.SpacingBetweenSlices.unique()) == 1\n\n\n        volume = [\n            dicom[instance_number.index(i)].pixel_array for i in df.instance_number\n        ]\n        volume = np.stack(volume)\n        volume = normalise_to_8bit(volume)\n        data.append(dotdict(\n            df=df,\n            volume=volume,\n        ))\n\n        if 'sagittal' in series_description.lower():\n            sort_data_by_group.append(position[0, 0])  # x\n        if 'axial' in series_description.lower():\n            sort_data_by_group.append(position[0, 2])  # z\n\n    data = [r for _, r in sorted(zip(sort_data_by_group, data))]\n    for i, r in enumerate(data):\n        r.df.loc[:, 'group'] = i\n\n    df = pd.concat([r.df for r in data])\n    df.loc[:, 'z'] = np.arange(len(df))\n    return df\n\ndef read_study(study_id, sagittal_t2_id, axial_t2_id):\n    return dotdict(\n        study_id = study_id,\n        sagittal_t2 =read_series(study_id, sagittal_t2_id, 'sagittal_t2'),\n        axial_t2 =read_series(study_id, axial_t2_id, 'axial_t2'),\n    )\ndef get_true_sagittal_t2_point(study_id, sagittal_t2_df):\n    series_id = sagittal_t2_df.iloc[0].series_id\n    label_coord_df = pd.read_csv(f'{DATA_KAGGLE_DIR}/train_label_coordinates.csv')\n    label_df = label_coord_df[\n         (label_coord_df.study_id == study_id)\n       & (label_coord_df.series_id == series_id)\n    ]\n    label_df = label_df.sort_values('level')\n    point=label_df[['x','y']].values\n    instance_number = label_df.instance_number.values\n\n    #mapping from instance num to z (array index)\n    map_instance_number, map_z = sagittal_t2_df[['instance_number','z',]].values.T\n    map = {n:z for n,z in zip(map_instance_number,map_z)}\n    z = [map[n] for n in instance_number]\n    return point,z\n\n############################################################3\n#post processing (sagittal_t2 point net)\ndef probability_to_point(probability, threshold=0.5):\n    #todo: handle mssing point\n    point=[]\n    for l in range(1, 6):\n        y, x = np.where(probability[l] > threshold)\n        y = round(y.mean())\n        x = round(x.mean())\n        point.append((x, y))\n    return point\n\n\ndef view_to_world(sagittal_t2_point, z, sagittal_t2_df, image_size):\n\n    H = sagittal_t2_df.iloc[0].H\n    W = sagittal_t2_df.iloc[0].W\n    scale_x = W / image_size\n    scale_y = H / image_size\n\n    xxyyzz = []\n    for l in range(1, 6):\n        x,y = sagittal_t2_point[l-1]\n        xx,yy,zz = project_to_3d(x*scale_x, y*scale_y, z, sagittal_t2_df)\n        xxyyzz.append((xx, yy, zz))\n\n    xxyyzz = np.array(xxyyzz)\n    return xxyyzz\n\ndef point_to_level(world_point, axial_t2_df):\n\n    # we get closest axial slices (z) to the CSC world points\n\n    xxyyzz = world_point\n    orientation = np.array(axial_t2_df.ImageOrientationPatient.values.tolist())\n    position = np.array(axial_t2_df.ImagePositionPatient.values.tolist())\n    ox = orientation[:, :3]\n    oy = orientation[:, 3:]\n    oz = np.cross(ox, oy)\n    t = xxyyzz.reshape(-1, 1, 3) - position.reshape(1, -1, 3)\n    dis = (oz.reshape(1, -1, 3) * t).sum(-1)  # np.dot(point-s,oz)\n    fdis = np.fabs(dis)\n    closest_z = fdis.argmin(-1)\n    closest_fdis = fdis.min(-1)\n    closest_df = axial_t2_df.iloc[closest_z]\n\n    if 1:\n        #<todo> hard/soft assigment, multi/single assigment\n        # no assignment based on distance\n\n        # allow point found in multi group\n        num_group   = len(axial_t2_df['group'].unique())\n        point_group = axial_t2_df.group.values[fdis.argsort(-1)[:, :3]].tolist()\n        point_group = [list(set(g)) for g in point_group]\n        group_point = [[] for g in range(num_group)]\n        for l in range(5):\n            for k in point_group[l]:\n                group_point[k].append(l)\n                # print(k)\n                # print(group_point[k])\n                # print(group_point)\n        group_point = [sorted(list(set(g))) for g in group_point]\n\n    D = len(axial_t2_df)\n    assigned_level=np.full(D,fill_value=0, dtype=int)\n    for group in range(num_group):\n        point_in_this_group = np.array(group_point[group])  # np.where(closest_df['group'] == group)[0]\n        slice_in_this_group = np.where(axial_t2_df['group'] == group)[0]\n        if len(point_in_this_group) == 0:\n            continue # unassigned, level=0\n\n        level = point_in_this_group[fdis[point_in_this_group][:, slice_in_this_group].argmin(0)] + 1\n        assigned_level[slice_in_this_group] = level\n\n    # sor =  (fdis.argmin(0)+1)[closest_z]\n    # closest_z= [ fdis.argmin(0)+1]\n    return assigned_level, closest_z, closest_fdis #dis is soft assignment\n\n\n#########################################################################3\n#visualisation\nlevel_color = [\n    [0, 0, 0],\n    [255, 0, 0],\n    [0, 255, 0],\n    [0, 0, 255],\n    [255, 255, 0],\n    [0, 255, 255],\n]\n\ndef probability_to_rgb(probability):\n    _6_,H,W = probability.shape\n    rgb = np.zeros((H, W, 3))\n    for i in range(1, 6):\n        rgb += probability[i].reshape(H, W, 1) * [[level_color[i]]]\n    rgb = rgb.astype(np.uint8)\n    return rgb\n\n\nclass Arrow3D(FancyArrowPatch):\n    def __init__(self, xs, ys, zs, *args, **kwargs):\n        super().__init__((0,0), (0,0), *args, **kwargs)\n        self._verts3d = xs, ys, zs\n\n    def do_3d_projection(self, renderer=None):\n        xs3d, ys3d, zs3d = self._verts3d\n        xs, ys, zs = proj3d.proj_transform(xs3d, ys3d, zs3d, self.axes.M)\n        self.set_positions((xs[0],ys[0]),(xs[1],ys[1]))\n        return np.min(zs)\n\ndef draw_slice(\n    ax, df,\n    is_slice =True, scolor=[[1,0,0]], salpha=[0.1],\n    is_border=True, bcolor=[[1,0,0]], balpha=[0.1],\n    is_origin=True, ocolor=[[1,0,0]], oalpha=[0.1],\n    is_arrow=True,\n):\n    df = df.copy()\n    df = df.reset_index(drop=True)\n\n    D = len(df)\n    if len(scolor)==1: scolor = scolor*D\n    if len(salpha)==1: salpha = salpha*D\n    if len(bcolor)==1: bcolor = bcolor*D\n    if len(balpha)==1: balpha = balpha*D\n    if len(ocolor)==1: ocolor = bcolor*D\n    if len(oalpha)==1: oalpha = balpha*D\n\n\n    #for i,d in df.iterrows():\n    for i in range(D):\n        d = df.iloc[i]\n        W, H = d.W, d.H\n        o0, o1, o2, o3, o4, o5 = d.ImageOrientationPatient\n        ox = np.array([o0, o1, o2])\n        oy = np.array([o3, o4, o5])\n        sx, sy, sz = d.ImagePositionPatient\n        s = np.array([sx, sy, sz])\n        delx, dely = d.PixelSpacing\n\n        p0 = s\n        p1 = s + W * delx * ox\n        p2 = s + H * dely * oy\n        p3 = s + H * dely * oy + W * delx * ox\n\n        grid = np.stack([p0, p1, p2, p3]).reshape(2, 2, 3)\n        gx = grid[:, :, 0]\n        gy = grid[:, :, 1]\n        gz = grid[:, :, 2]\n\n        #outline\n        if is_slice:\n            ax.plot_surface(gx, gy, gz, color=scolor[i], alpha=salpha[i])\n\n        if is_border:\n            line = np.stack([p0, p1, p3, p2] )\n            ax.plot(line[:,0], line[:,1], zs=line[:,2], color=ocolor[i], alpha=oalpha[i])\n\n        if is_origin:\n            ax.scatter([sx], [sy], [sz],   color=ocolor[i], alpha=oalpha[i])\n\n    #check ordering of slice\n    if is_arrow :\n        sx0, sy0, sz0 = df.iloc[0].ImagePositionPatient\n        sx1, sy1, sz1 = df.iloc[-1].ImagePositionPatient\n        arrow_prop_dict = dict(mutation_scale=20, arrowstyle='-|>', color='k', shrinkA=0, shrinkB=0)\n        a = Arrow3D([sx0, sx1], [sy0, sy1], [sz0, sz1], **arrow_prop_dict)\n        ax.add_artist(a)\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\n\ndef add_axial_angle(df):\n    def calculate_angle(orientation):\n        # 行ベクトルと列ベクトルを抽出\n        row_vector = np.array(orientation[:3])\n        col_vector = np.array(orientation[3:])\n        \n        # スライスの法線ベクトルを計算\n        normal_vector = np.cross(row_vector, col_vector)\n        \n        # 純粋なaxial方向の法線ベクトル（頭尾方向）\n        true_axial_normal = np.array([0, 0, 1])\n        \n        # 角度を計算（ラジアン）\n        angle_rad = np.arccos(np.clip(np.dot(normal_vector, true_axial_normal), -1.0, 1.0))\n        \n        # ラジアンから度に変換\n        angle_deg = np.degrees(angle_rad)\n        \n        # 傾きの方向を決定\n        # 法線ベクトルのy成分の符号を使用\n        direction = np.sign(normal_vector[1])\n        \n        # 方向に基づいて角度の符号を調整\n        angle_deg *= direction\n        \n        return angle_deg\n\n    # ImageOrientationPatientから角度を計算し、新しい列として追加\n    df['axial_angle'] = df['ImageOrientationPatient'].apply(calculate_angle)\n    \n    return df\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_series_descriptions.csv')\ndf = df[df.series_description=='Axial T2']","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tqdm import tqdm\ndfs = []\nerror_ids = []\nall_dfs = []\nfor study_id_n, study_id in enumerate(tqdm(id_df[id_df.study_id.isin([3429410502])].study_id.unique())):\n    for axial_t2_id in id_df[(id_df.study_id == study_id) & (id_df.series_description=='Axial T2')].series_id.values:\n\n        sagittal_t2_id = id_df[(id_df.study_id == study_id) & (id_df.series_description=='Sagittal T2/STIR')].iloc[0].series_id\n        data = read_study(study_id, axial_t2_id=axial_t2_id, sagittal_t2_id=sagittal_t2_id)\n\n\n        #--- step.1 : detect 2d point in sagittal_t2\n        sagittal_t2 = data.sagittal_t2.volume\n        sagittal_t2_df = data.sagittal_t2.df\n        axial_t2_df = data.axial_t2.df\n        all_dfs.append(axial_t2_df)\n\n        D,H,W = sagittal_t2.shape\n        image = resize_volume(sagittal_t2,cfg.point_net.image_size)\n\n        sagittal_t2_z = D//2\n        image = image[sagittal_t2_z] #we use only center image #todo: better selection\n\n        batch = dotdict(\n            sagittal=torch.from_numpy(image).unsqueeze(0).unsqueeze(0).byte()\n        )\n        with torch.cuda.amp.autocast(enabled=True):\n            with torch.no_grad():\n                output = point_net(batch)\n\n        probability = output['probability'][0].float().data.cpu().numpy()\n        sagittal_t2_point = probability_to_point(probability) #5 level 2d points, todo: check invalid output point\n\n        #for debug and development\n        point_hat, z_hat = sagittal_t2_point_hat = get_true_sagittal_t2_point(study_id, sagittal_t2_df)\n        point_hat = point_hat*[[cfg.point_net.image_size/W, cfg.point_net.image_size/H]]\n\n\n        #--- step.2 : perdict slice level of axial_t2\n        world_point = world_point = view_to_world(sagittal_t2_point, sagittal_t2_z, sagittal_t2_df, cfg.point_net.image_size)\n        assigned_level, closest_z, closest_fdis = axial_t2_level = point_to_level(world_point, axial_t2_df)\n\n        axial_t2_df['level'] = assigned_level\n#             break\n        exist_closest_fdis = closest_fdis[np.unique([c for c in assigned_level if c != 0])-1]\n        exist_closest_z = closest_z[np.unique([c for c in assigned_level if c != 0])-1]\n        ns = axial_t2_df.iloc[exist_closest_z].instance_number    \n        axial_t2_df.loc[~axial_t2_df.instance_number.isin(ns), 'closest'] = 0\n        axial_t2_df.loc[axial_t2_df.instance_number.isin(ns), 'closest'] = 1\n        axial_t2_df['closest'] = axial_t2_df['closest'].astype(int)\n        for level in range(1, 6):\n            axial_t2_df.loc[axial_t2_df.level == level, 'dis'] = closest_fdis[level-1]            \n\n        for n, dis in zip(ns, exist_closest_fdis):\n            axial_t2_df.loc[axial_t2_df.instance_number == n, 'dis'] = dis\n\n        assert len(assigned_level)==len(axial_t2_df)\n        axial_t2_df = add_axial_angle(axial_t2_df)\n\n        dfs.append(axial_t2_df)        \n\n#             if exist_closest_fdis.max() > 5:                \n        print(study_id, axial_t2_id, exist_closest_fdis.max(), assigned_level)\n        tmp = axial_t2_df[axial_t2_df.closest==1]\n        if len(tmp[tmp['level'] == 5]) > 0:\n            print(tmp[tmp['level'] == 5].axial_angle.values[0])\n        ###################################################################\n        #visualisation\n\n        # https://matplotlib.org/stable/gallery/mplot3d/mixed_subplots.html\n        fig = plt.figure(figsize=(23, 6))\n        ax1 = fig.add_subplot(1, 1, 1)\n        ax2 = fig.add_subplot(1, 2, 2, projection='3d')\n\n        # detection result\n        p = probability_to_rgb(probability)\n        m = cv2.cvtColor(image, cv2.COLOR_GRAY2BGR)\n        m = 255 - (255 - m * 0.8) * (1 - p / 255)\n\n        ax1.imshow(m / 255)\n        ax1.set_title(f'sagittal keypoint detection (unet)\\n series_id: {axial_t2_id}')\n\n\n        # draw  assigned_level\n        level_ncolor = np.array(level_color) / 255\n        coloring = level_ncolor[assigned_level].tolist()\n        draw_slice(\n            ax2, axial_t2_df,\n            is_slice=True,   scolor=coloring, salpha=[0.1],\n            is_border=True,  bcolor=coloring, balpha=[0.2],\n            is_origin=False, ocolor=[[0, 0, 0]], oalpha=[0.0],\n            is_arrow=True\n        )\n\n    #     draw world_point\n        ax2.scatter(world_point[:, 0], world_point[:, 1], world_point[:, 2], alpha=1, color='black')\n\n\n        ### draw closest slice\n        coloring = level_ncolor[1:].tolist()\n        draw_slice(\n            ax2, axial_t2_df.iloc[closest_z],\n            is_slice=True, scolor=coloring, salpha=[0.1],\n            is_border=True, bcolor=coloring, balpha=[1],\n            is_origin=False, ocolor=[[1, 0, 0]], oalpha=[0],\n            is_arrow=False\n        )\n\n        ax2.set_aspect('equal')\n        ax2.set_title(f'axial slice assignment\\n series_id:{sagittal_t2_id}')\n        ax2.set_xlabel('x')\n        ax2.set_ylabel('y')\n        ax2.set_zlabel('z')\n        ax2.view_init(elev=0, azim=-10, roll=0)\n        plt.tight_layout(pad=2)\n        plt.show()\n\nresult_df = pd.concat(dfs)\nall_axial_df = pd.concat(all_dfs)\nall_axial_df.to_csv('axial_direction.csv', index=False)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tqdm import tqdm\ndfs = []\nerror_ids = []\nall_dfs = []\nfor study_id_n, study_id in enumerate(tqdm(id_df.study_id.unique())):\n#     print(study_id)\n    for axial_t2_id in id_df[(id_df.study_id == study_id) & (id_df.series_description=='Axial T2')].series_id.values:\n        try:\n\n            sagittal_t2_id = id_df[(id_df.study_id == study_id) & (id_df.series_description=='Sagittal T2/STIR')].iloc[0].series_id\n            data = read_study(study_id, axial_t2_id=axial_t2_id, sagittal_t2_id=sagittal_t2_id)\n\n            #--- step.1 : detect 2d point in sagittal_t2\n            sagittal_t2 = data.sagittal_t2.volume\n            sagittal_t2_df = data.sagittal_t2.df\n            axial_t2_df = data.axial_t2.df\n            all_dfs.append(axial_t2_df)\n\n            D,H,W = sagittal_t2.shape\n            image = resize_volume(sagittal_t2,cfg.point_net.image_size)\n\n            sagittal_t2_z = D//2\n            image = image[sagittal_t2_z] #we use only center image #todo: better selection\n\n            batch = dotdict(\n                sagittal=torch.from_numpy(image).unsqueeze(0).unsqueeze(0).byte()\n            )\n            with torch.cuda.amp.autocast(enabled=True):\n                with torch.no_grad():\n                    output = point_net(batch)\n\n            probability = output['probability'][0].float().data.cpu().numpy()\n            sagittal_t2_point = probability_to_point(probability) #5 level 2d points, todo: check invalid output point\n\n            #for debug and development\n            point_hat, z_hat = sagittal_t2_point_hat = get_true_sagittal_t2_point(study_id, sagittal_t2_df)\n            point_hat = point_hat*[[cfg.point_net.image_size/W, cfg.point_net.image_size/H]]\n\n\n            #--- step.2 : perdict slice level of axial_t2\n            world_point = world_point = view_to_world(sagittal_t2_point, sagittal_t2_z, sagittal_t2_df, cfg.point_net.image_size)\n            assigned_level, closest_z, closest_fdis = axial_t2_level = point_to_level(world_point, axial_t2_df)\n            \n            axial_t2_df['level'] = assigned_level\n#             break\n            exist_closest_fdis = closest_fdis[np.unique([c for c in assigned_level if c != 0])-1]\n            exist_closest_z = closest_z[np.unique([c for c in assigned_level if c != 0])-1]\n            ns = axial_t2_df.iloc[exist_closest_z].instance_number    \n            axial_t2_df.loc[~axial_t2_df.instance_number.isin(ns), 'closest'] = 0\n            axial_t2_df.loc[axial_t2_df.instance_number.isin(ns), 'closest'] = 1\n            axial_t2_df['closest'] = axial_t2_df['closest'].astype(int)\n            for level in range(1, 6):\n                axial_t2_df.loc[axial_t2_df.level == level, 'dis'] = closest_fdis[level-1]            \n                \n            for n, dis in zip(ns, exist_closest_fdis):\n                axial_t2_df.loc[axial_t2_df.instance_number == n, 'dis'] = dis\n                \n            assert len(assigned_level)==len(axial_t2_df)\n            axial_t2_df = add_axial_angle(axial_t2_df)\n            \n            dfs.append(axial_t2_df)        \n\n    \n        except:\n            error_ids.append(axial_t2_id)\ndf = pd.concat(dfs)\nall_axial_df = pd.concat(all_dfs)\nall_axial_df.to_csv('axial_direction.csv', index=False)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"level_df = pd.read_csv('/kaggle/input/rsna2024-csv/axial_level_pred_keroppi_v3.csv')\nlevel_df['closest'] = 1\nlevel_df['dis'] = 1\nlevel_df['level'] = level_df.pred_level.values\nlevel_df = level_df[~level_df.series_id.isin(df.series_id)][list(set(list(df)) & set(list(level_df)))]\nlevel_df['series_description'] = 'axial_t2'\ndf = pd.concat([df, level_df])\ndf.to_csv('axial_closest_df.csv', index=False)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}