{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":84969,"databundleVersionId":10033515,"sourceType":"competition"},{"sourceId":9965452,"sourceType":"datasetVersion","datasetId":6074268},{"sourceId":9982740,"sourceType":"datasetVersion","datasetId":6040928}],"dockerImageVersionId":30787,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"#!pip download connected-components-3d\n#!pip download zarr\n\ntry:\n    import zarr\nexcept: \n    !cp -r '/kaggle/input/hengck-czii-cryo-et-02/wheel_file' '/kaggle/working/'\n    !pip install /kaggle/working/wheel_file/asciitree-0.3.3/asciitree-0.3.3\n    !pip install --no-index --find-links=/kaggle/working/wheel_file zarr\n    !pip install --no-index --find-links=/kaggle/working/wheel_file connected-components-3d\n\nprint('PIP INSTALL OK!!!')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-22T19:24:02.604665Z","iopub.execute_input":"2024-11-22T19:24:02.605658Z","iopub.status.idle":"2024-11-22T19:24:02.613333Z","shell.execute_reply.started":"2024-11-22T19:24:02.605618Z","shell.execute_reply":"2024-11-22T19:24:02.612294Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from datetime import datetime\n\nimport pandas as pd\nimport pytz\nprint('LOGGING TIME OF START:',  datetime.strftime(datetime.now(pytz.timezone('Asia/Singapore')), \"%Y-%m-%d %H:%M:%S\"))\n\nimport sys\nsys.path.append('/kaggle/input/hengck-czii-cryo-et-02')\n\nfrom czii_helper import *\nfrom dataset import *\nfrom model import *\nimport numpy as np\nfrom scipy.optimize import linear_sum_assignment\nimport glob\nimport cc3d\nimport cv2\n\nimport matplotlib.pyplot as plt\nfrom mpl_toolkits.mplot3d import Axes3D\n\nimport warnings\nwarnings.filterwarnings(\"ignore\", category=FutureWarning, module=\"torch.nn.parallel\")\n\nprint('IMPORT OK!!!')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2024-11-22T19:24:02.617601Z","iopub.execute_input":"2024-11-22T19:24:02.617875Z","iopub.status.idle":"2024-11-22T19:24:02.627464Z","shell.execute_reply.started":"2024-11-22T19:24:02.617849Z","shell.execute_reply":"2024-11-22T19:24:02.626635Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DATA_KAGGLE_DIR = '/kaggle/input/czii-cryo-et-object-identification'\n\nMODE='submit'\n\n \nif MODE == 'local':\n    valid_dir = f'{DATA_KAGGLE_DIR}/train'\n    valid_id = ['TS_5_4', 'TS_73_6','TS_99_9'] # \n    #valid_id = ['TS_73_6', ]  # fold2\n    #valid_id = ['TS_5_4', 'TS_5_4','TS_5_4',]  # fold0\n\nif MODE == 'submit':\n    valid_dir = f'{DATA_KAGGLE_DIR}/test'\n    valid_id = glob.glob(f'{valid_dir}/static/ExperimentRuns/*')\n    valid_id = [f.split('/')[-1] for f in valid_id]\n\nprint('valid_id:', len(valid_id), valid_id)\n\ncfg = dotdict(\n    arch='resnet34d',\n    checkpoint= \\\n    '/kaggle/input/hengck-czii-cryo-et-weights-01/resnet34d-scan640-fold0-00000154.pth',\n        #'/kaggle/input/hengck-czii-cryo-et-weights-01/resnet18d-aug-rot-scan320-fold0-00005610.pth',\n        #'/kaggle/input/hengck-czii-cryo-et-weights-01/resnet34d-scan320-fold0-00003168.pth',\n        #'/kaggle/input/hengck-czii-cryo-et-weights-01/resnet34d-simple3.0-00002300.pth',\n\n    threshold={\n        'apo-ferritin': 0.05,\n        'beta-amylase': 0.05,\n        'beta-galactosidase': 0.05,\n        'ribosome': 0.05,\n        'thyroglobulin': 0.05,\n        'virus-like-particle': 0.05,\n    },\n)\n\nprint('MODE:', MODE)\nprint('SETTING OK!!!')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-22T19:24:02.628726Z","iopub.execute_input":"2024-11-22T19:24:02.628977Z","iopub.status.idle":"2024-11-22T19:24:02.644103Z","shell.execute_reply.started":"2024-11-22T19:24:02.628952Z","shell.execute_reply":"2024-11-22T19:24:02.643217Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"net = Net(pretrained=False, cfg=cfg)\nstate_dict = torch.load(cfg.checkpoint, map_location=lambda storage, loc: storage, weights_only=True)['state_dict']\nprint(net.load_state_dict(state_dict, strict=False))\nprint(net.arch)\nprint('MODEL OK!!!')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-22T19:24:02.645366Z","iopub.execute_input":"2024-11-22T19:24:02.645596Z","iopub.status.idle":"2024-11-22T19:24:03.136053Z","shell.execute_reply.started":"2024-11-22T19:24:02.645570Z","shell.execute_reply":"2024-11-22T19:24:03.135147Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def make_weight(h, w, top=10,left=0,bottom=0,right=0):\n    weight = np.full((h, w),fill_value=1)\n\n    if top>0:\n        wt = np.ones((h, w))\n        wt[:top]=np.linspace(0,1,top+1)[1:].reshape(-1,1)\n        weight = np.minimum(weight, wt)\n\n    if left>0:\n        wt = np.ones((h, w))\n        wt[:,:left] = np.linspace(0, 1, left + 1)[1:].reshape(1, -1)\n        weight = np.minimum(weight, wt)\n\n    if bottom>0:\n        wt = np.ones((h, w))\n        wt[-bottom:]=np.linspace(0,1,bottom+1)[1:][::-1].reshape(-1,1)\n        weight = np.minimum(weight, wt)\n\n    if right>0:\n        wt = np.ones((h, w))\n        wt[:,-right:] = np.linspace(0, 1, right + 1)[1:][::-1].reshape(1, -1)\n        weight = np.minimum(weight, wt)\n\n    return weight\n\n\nclass Scanner:\n    def __init__(self, w,h,d, overlap ):\n        csum=np.cumsum(overlap)\n        self.xyz = [\n            (0, 0, 0),\n            (0, 0,   d-csum[0]),\n            (0, 0, 2*d-csum[1]),\n            (0, 0, 3*d-csum[2]),\n            (0, 0, 4*d-csum[3]),\n        ]\n        self.length = len(self.xyz)\n        self.weight = [None for i in range(len(self.xyz))]\n        self.weight[0] = make_weight(h=d, w=1, top=0,          left=0, bottom=overlap[0], right=0)\n        self.weight[1] = make_weight(h=d, w=1, top=overlap[0], left=0, bottom=overlap[1], right=0)\n        self.weight[2] = make_weight(h=d, w=1, top=overlap[1], left=0, bottom=overlap[2], right=0)\n        self.weight[3] = make_weight(h=d, w=1, top=overlap[2], left=0, bottom=overlap[3], right=0)\n        self.weight[4] = make_weight(h=d, w=1, top=overlap[3], left=0, bottom=0,          right=0)\n        self.generator = self.create_generator()\n\n    def create_generator(self):\n        for i in range(self.length):\n            yield self.weight[i], self.xyz[i]\n\n    def __iter__(self):\n        self.generator = self.create_generator()\n        return self\n\n    def __next__(self):\n        return next(self.generator)\n\n\ndef normalise_by_percentile(data, min=5, max=99):\n    min = np.percentile(data,min)\n    max = np.percentile(data,max)\n    data = (data-min)/(max-min)\n    return data\n    \ndef draw_probability(probability, color):\n\t_6_, D, H, W = probability.shape\n\tpcolor = np.zeros((_6_, D, H, W, 3), dtype=np.float32)\n\tfor i in range(_6_):\n\t\tpcolor[i] += probability[i][..., None] * [[[color[i]]]]\n\n\t# ----\n\tp_max = pcolor.max(0)\n\tp0 = p_max.max(0)\n\tp1 = p_max.max(1)\n\tp2 = p_max.max(2)\n\tall = np.zeros((H + D, W + D, 3), dtype=np.uint8)\n\tall[:H, :W] = p0\n\tall[H:, :W] = p1\n\tall[:H, W:] = p2.transpose(1, 0, 2)\n\tall[H] = 255\n\tall[:, W] = 255\n\tall = np.clip(all, 0, 255)\n\treturn all\n\n\n\n#start here !!!! -------------------------------------------------------\ndef run_submit(net):  \n    \n    net.output_type = ['infer']\n    net = torch.nn.DataParallel(net, device_ids=[0, 1])\n    net.cuda()\n    net.eval()\n    \n    num_slice = 48\n    D, H, W = (184, 630, 630)\n    scanner = Scanner(w=1, h=1, d=48, overlap=[14, 14, 14, 14, 14])\n    threshold = list(cfg.threshold.values())\n    \n    with torch.no_grad():\n        probability = torch.zeros((7, D, H, W), device='cuda')\n        count = torch.zeros((7, D, H, W), device='cuda') \n        scanner.weight=[\n            torch.from_numpy(wt.reshape(1, num_slice, 1, 1)).float().cuda() for wt in scanner.weight\n        ] \n        threshold = torch.tensor(threshold, device='cuda').reshape(6, 1, 1, 1)\n\n    submit_df = []\n    total_time = 0  \n    for i,id in enumerate(valid_id):\n        start_timer = timer()\n        torch.cuda.empty_cache() \n        \n        print(i, id, '---------------')\n        volume, scale = read_one_data(id, static_dir=f'{valid_dir}/static') \n        volume = normalise_by_percentile(volume)\n        D, H, W = volume.shape\n        assert ((D, H, W)== (184, 630, 630))\n \n        probability.zero_()\n        count.zero_()\n\n        \n        with torch.amp.autocast('cuda', enabled=True):\n            with torch.no_grad():\n \n                for weight, (x, y, z) in scanner:\n                    print('\\r', f'{id}:{(x, y, z)}', end='', flush=True)\n\n                    image = volume[z:z + num_slice]\n                    batch = dotdict(\n                        image=torch.from_numpy(\n                            np.stack([\n                                image,\n                                np.rot90(image, k=1, axes=(1,2)),\n                            ])\n                        ),\n                    )\n                    batch['image']=F.pad(batch['image'],[0,10,0,10])\n                    output = net(batch)\n                    prob = output['particle'][...,:H,:W]\n\n                    prob0 = prob[0]\n                    prob1 = torch.rot90(prob[1], k=-1, dims=(2,3))\n                    prob  = (prob0 + prob1)/2\n\n                    probability[:, z:z + num_slice] += weight * prob\n                    count[:, z:z + num_slice] += weight\n                print('')\n                probability = probability / count\n                probability0 = probability[1:]\n                probability1 = F.interpolate(probability0, scale_factor=0.5, mode='bilinear', align_corners=False)\n                #smaller for faster post-processing\n        \n        binary0 = (probability0 > threshold).data.cpu().numpy()\n        binary1 = (probability1 > threshold).data.cpu().numpy()\n        location = [np.empty((0,3)) for i in range(6)]\n\n        #1: apo-ferritin, radius=60\n        for c in [0]:\n            componet = cc3d.connected_components(binary0[c])\n            stats = cc3d.statistics(componet)\n            zyx = stats['centroids'][1:] * [scale]\n            xyz = np.ascontiguousarray(zyx[:, ::-1])\n            location[c] = xyz\n\n        # beta-amylase is ignored, rest of particles have radius>=90\n        for c in [2,3,4,5]:\n            componet = cc3d.connected_components(binary1[c])\n            stats = cc3d.statistics(componet)\n            zyx = stats['centroids'][1:] * [scale]*[[1,2,2]]\n            xyz = np.ascontiguousarray(zyx[:, ::-1])\n            location[c] = xyz\n        print('location', np.concatenate(location).shape)\n        \n        for name,xyz in zip(PARTICLE_NAME,location):\n            if len(xyz)==0: continue\n            submit_df.append(\n                pd.DataFrame({'experiment': id, 'particle_type':name,'x':xyz[:,0],'y':xyz[:,1],'z':xyz[:,2]})\n            )\n        time_taken = timer() - start_timer\n        total_time += time_taken\n        print(time_to_str(time_taken, 'sec'))\n\n        #debug\n        if i==0:\n            p = probability0.data.cpu().numpy()\n            \n            m0 = np.clip(volume,0,1)\n            m0 = m0.mean(0) \n            m0 = np.dstack([m0, m0, m0])\n\n            g0 = np.zeros((H, W, 3), dtype=np.float32)\n            for c in [0,1,2]:\n                color=PARTICLE[c]['color']\n                q = p[c].max(0)[...,None]\n                g0 += q*[color]\n            g0 = np.clip(g0/255 > 0.1,0,1)\n            g0 = 1-(1-m0)*(1-g0)\n\n            g1 = np.zeros((H, W, 3), dtype=np.float32)\n            for c in [3,4, 5,]:\n                color = PARTICLE[c]['color']\n                q = p[c].max(0)[..., None]\n                g1 += q * [color]\n            g1 = np.clip(g1/255 > 0.1,0,1)\n            g1 = 1 - (1 - m0) * (1 - g1)\n\n            m0g0g1 =np.hstack([m0,g1,g0])\n            plt.imshow(m0g0g1)\n            plt.show()\n            #plt.waitforbuttonpress()\n\n    \n            color = [PARTICLE[c]['color'] for c in range(6)]\n            all = draw_probability(p, color)\n            plt.imshow(all)\n            plt.show()\n\n\n    torch.cuda.empty_cache() \n    print('\\ndone!') \n    num_volume = len(valid_id)\n    print(f'Total time for {num_volume} volumes:', time_to_str(total_time, 'min'))\n    print(f'Total time for 500 volumes:', time_to_str(total_time/num_volume*500, 'min'))\n    print('')\n    submit_df = pd.concat(submit_df)\n    submit_df.insert(loc=0, column='id', value=np.arange(len(submit_df)))\n    return submit_df\n\n\n\nif 1:\n    submit_df = run_submit(net)\n    print('submit_df', submit_df.shape)\n    print(submit_df)\n    submit_df.to_csv('submission.csv', index=False)\n\nprint('MODE:', MODE)\nprint('SUBMIT OK!!!')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-22T19:24:03.138139Z","iopub.execute_input":"2024-11-22T19:24:03.138536Z","iopub.status.idle":"2024-11-22T19:24:52.098709Z","shell.execute_reply.started":"2024-11-22T19:24:03.138495Z","shell.execute_reply":"2024-11-22T19:24:52.097804Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if 1:\n    if MODE=='local':\n        pd.set_option('display.max_columns', 500)\n        pd.set_option('display.width', 1000)\n        \n        submit_df=pd.read_csv(\n           'submission.csv'\n            # '/kaggle/input/hengck-czii-cryo-et-weights-01/submission.csv'\n        )\n        gb, lb_score = compute_lb(submit_df, f'{valid_dir}/overlay')\n        print('lb_score:',lb_score)\n        print(gb)\n        print('')\n\n\n        #--------------------------------------------\n        #visualisation\n\n        fig = plt.figure(figsize=(18, 8))\n\n        id = valid_id[0]\n        truth = read_one_truth(id,overlay_dir=f'{valid_dir}/overlay')\n\n        submit_df=pd.read_csv('submission.csv')\n        submit_df = submit_df[submit_df['experiment']==id]\n        for p in PARTICLE:\n            p = dotdict(p)\n            xyz_truth = truth[p.name]\n            xyz_predict = submit_df[submit_df['particle_type']==p.name][['x','y','z']].values\n            hit, fp, miss, metric = do_one_eval(xyz_truth, xyz_predict, p.radius)\n            # print(id, p.name)\n            # print('\\t num truth   :',len(xyz_truth) )\n            # print('\\t num predict :',len(xyz_predict) )\n            # print('\\t num hit  :',len(hit[0]) )\n            # print('\\t num fp   :',len(fp) )\n            # print('\\t num miss :',len(miss) )\n            ax = fig.add_subplot(2, 3, p.label, projection='3d')\n\n            if 0:\n                pt = xyz_predict\n                ax.scatter(pt[:, 0], pt[:, 1], pt[:, 2], alpha=0.25, color='b', label='predict')\n                pt = xyz_truth\n                ax.scatter(pt[:, 0], pt[:, 1], pt[:, 2], s=160, alpha=0.25, color='r', label='truth')\n                ax.scatter(pt[:, 0], pt[:, 1], pt[:, 2], s=160, facecolors='none', edgecolors='r')\n            if 1:\n                if hit[0]:\n                    pt = xyz_predict[hit[0]]\n                    ax.scatter(pt[:, 0], pt[:, 1], pt[:, 2], alpha=0.5, color='r', label='predict')\n                    pt = xyz_truth[hit[1]]\n                    ax.scatter(pt[:, 0], pt[:, 1], pt[:, 2], s=80, facecolors='none', edgecolors='r', label='truth')\n                if fp:\n                    pt = xyz_predict[fp]\n                    ax.scatter(pt[:, 0], pt[:, 1], pt[:, 2], alpha=0.5, color='k', label='fp')\n                if miss:\n                    pt = xyz_truth[miss]\n                    ax.scatter(pt[:, 0], pt[:, 1], pt[:, 2], s=160, alpha=0.5, facecolors='none', edgecolors='k', label='miss')\n            ax.legend()\n            ax.set_title(\n                f'{id}:{p.name} ({p.difficulty})\\npredict={metric[0]}, truth={metric[1]}, hit={metric[2]}, miss={metric[3]}, fp={metric[4]}')\n\n        plt.tight_layout()\n        plt.show()\n        zz=0","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-22T19:24:52.099975Z","iopub.execute_input":"2024-11-22T19:24:52.100260Z","iopub.status.idle":"2024-11-22T19:24:53.700481Z","shell.execute_reply.started":"2024-11-22T19:24:52.100234Z","shell.execute_reply":"2024-11-22T19:24:53.699622Z"}},"outputs":[],"execution_count":null}]}