{"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":"gpu","dataSources":[{"sourceId":84969,"databundleVersionId":10033515,"sourceType":"competition"},{"sourceId":9846452,"sourceType":"datasetVersion","datasetId":6040935},{"sourceId":36518,"sourceType":"modelInstanceVersion","modelInstanceId":30754,"modelId":43936},{"sourceId":67132,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":55968,"modelId":77111}],"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#!pip download numcodecs\n\ntry:\n    import zarr\nexcept: \n    !cp -r '/kaggle/input/hengck-czii-cryo-et-01/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-09T00:59:44.484362Z","iopub.execute_input":"2024-11-09T00:59:44.485113Z","iopub.status.idle":"2024-11-09T01:00:44.007266Z","shell.execute_reply.started":"2024-11-09T00:59:44.485065Z","shell.execute_reply":"2024-11-09T01:00:44.006162Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#https://github.com/seung-lab/connected-components-3d/\n#pip install connected-components-3d\n#pip install zarr\n\n\nfrom datetime import datetime\nprint('LOGGING TIME OF START:', datetime.now().strftime('%Y-%m-%d_%H-%M-%S'))\n\nimport sys\nsys.path.append('/kaggle/input/hengck-czii-cryo-et-01')\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\nprint('IMPORT OK!!!')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2024-11-09T01:00:44.009342Z","iopub.execute_input":"2024-11-09T01:00:44.009678Z","iopub.status.idle":"2024-11-09T01:00:50.574462Z","shell.execute_reply.started":"2024-11-09T01:00:44.009627Z","shell.execute_reply":"2024-11-09T01:00:50.573488Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DATA_KAGGLE_DIR = '/kaggle/input/czii-cryo-et-object-identification'\n\nMODE='submit'\n\nif MODE=='local':\n    valid_dir =f'{DATA_KAGGLE_DIR}/train'\nif MODE=='submit':\n    valid_dir =f'{DATA_KAGGLE_DIR}/test'\n\nvalid_id = glob.glob(f'{valid_dir}/static/ExperimentRuns/*')\nvalid_id = [f.split('/')[-1] for f in valid_id]\nprint('valid_id:',len(valid_id), valid_id)\n\ncfg = dotdict(\n    checkpoint='/kaggle/input/model_60.pth/pytorch/ep/1/model_ep_60.pth',\n    threshold=0.1,\n)\n\nprint('MODE:', MODE)\nprint('SETTING OK!!!')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-09T01:00:50.575743Z","iopub.execute_input":"2024-11-09T01:00:50.576533Z","iopub.status.idle":"2024-11-09T01:00:50.587612Z","shell.execute_reply.started":"2024-11-09T01:00:50.576485Z","shell.execute_reply":"2024-11-09T01:00:50.586680Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"net = Net(pretrained=False)\nstate_dict = torch.load(cfg.checkpoint, map_location=lambda storage, loc: storage)\nprint(net.load_state_dict(state_dict, strict=False))\n\n#print(net.arch)\nprint('MODEL OK!!!')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-09T01:00:50.589325Z","iopub.execute_input":"2024-11-09T01:00:50.589678Z","iopub.status.idle":"2024-11-09T01:00:52.499792Z","shell.execute_reply.started":"2024-11-09T01:00:50.589632Z","shell.execute_reply":"2024-11-09T01:00:52.498811Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def probability_to_location(probability,cfg):\n\n    location={}\n    for p in PARTICLE:\n        p = dotdict(p)\n        l = p.label\n\n        cc, P = cc3d.connected_components(probability[l]>cfg.threshold, return_N=True)\n        stats = cc3d.statistics(cc)\n        zyx=stats['centroids'][1:]*10\n        xyz =np.ascontiguousarray(zyx[:,::-1])\n        location[p.name]=xyz\n        '''\n            j=1\n            z,y,x = np.where(cc==j)\n            z=z.mean()\n            y=y.mean()\n            x=x.mean()\n            print([x,y,z])\n        '''\n    return location\n\ndef location_to_df(location):\n    location_df = []\n    for p in PARTICLE:\n        p = dotdict(p)\n        xyz = location[p.name]\n        if len(xyz)>0:\n            df = pd.DataFrame(data=xyz, columns=['x','y','z'])\n            #df.loc[:,'particle_type']= p.name\n            df.insert(loc=0, column='particle_type', value=p.name)\n            location_df.append(df)\n    location_df = pd.concat(location_df)\n    return location_df\n#start here !!!!\ndef run_submit():\n    \n    net.cuda()\n    net.eval()\n    net.output_type = ['infer']\n\n    submit_df = []\n    start_timer = timer()\n    for i,id in enumerate(valid_id):\n        print(i, id, '---------------')\n        volume = read_one_data(id, static_dir=f'{valid_dir}/static/ExperimentRuns')\n        D, H, W = volume.shape\n        print(D, H, W)\n\n        probability = np.zeros((7, D, H, W), dtype=np.float32)\n        count = np.zeros((7, D, H, W), dtype=np.float32)\n        pad_volume = np.pad(volume, [[0, 0], [0, 640 - H], [0, 640 - W]], mode='constant', constant_values=0)\n        \n        num_slice=64\n        zz = list(range(0, D - num_slice, num_slice//2)) + [D - num_slice]\n        for z in zz:\n            print('\\r',f'z:{z}', end='',flush=True)\n            image = pad_volume[z:z + num_slice]\n            batch = dotdict(\n                image=torch.from_numpy(image).unsqueeze(0),\n            )\n            with torch.amp.autocast('cuda', enabled=True):\n                with torch.no_grad():\n                    output = net(batch)\n            prob = output['particle'][0].cpu().numpy()\n            probability[:, z:z + num_slice] = prob[:, :, :H, :W]\n            count[:, z:z + num_slice] += 1\n        probability = probability / (count + 0.0001)\n        location = probability_to_location(probability, cfg)\n        df = location_to_df(location)\n        df.insert(loc=0, column='experiment', value=id)\n        submit_df.append(df)\n        print('')\n        print(time_to_str(timer() - start_timer, 'sec'))\n\n    print('\\ndone!')\n    total_time = timer() - start_timer\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()\n    print('submit_df', submit_df.shape)\n    print(submit_df)\n    submit_df.to_csv('submission.csv', index=False)\n\n#!ls\nprint('SUBMIT OK!!!')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-09T01:00:52.501912Z","iopub.execute_input":"2024-11-09T01:00:52.502245Z","iopub.status.idle":"2024-11-09T01:02:07.388454Z","shell.execute_reply.started":"2024-11-09T01:00:52.502210Z","shell.execute_reply":"2024-11-09T01:02:07.387250Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def do_one_eval(truth, predict, radius):\n    P=len(predict)\n    T=len(truth)\n\n    distance = predict.reshape(P,1,3)-truth.reshape(1,T,3)\n    distance = distance**2\n    distance = distance.sum(axis=2)\n    distance = np.sqrt(distance)\n    p_index, t_index = linear_sum_assignment(distance)\n\n    threshold = radius * 0.5\n    hit = []\n    fp = []\n    miss = set(range(T))\n    for pi, ti in zip(p_index, t_index):\n        if distance[pi, ti] <= threshold:\n            hit.append([pi, ti])\n            miss.remove(ti)\n        else:\n            fp.append(pi)\n    miss = list(miss)\n    hit = (np.array(hit).T).tolist()\n    if hit==[]:\n        hit=[[],[]]\n\n    #todo: compute precision,recall, F-beta score\n    return hit, fp, miss\n    \n#debug\nif 1:\n    if MODE=='local':\n    #if 1:\n        import matplotlib.pyplot as plt\n        from mpl_toolkits.mplot3d import Axes3D\n\n        fig = plt.figure(figsize=(18, 8))\n\n        id = valid_id[1]\n        truth = read_one_truth(id,overlay_dir=f'{valid_dir}/overlay/ExperimentRuns')\n\n        submit_df=pd.read_csv(\n            #'submission.csv'\n            '/kaggle/input/czii-cryo-et-object-identification/sample_submission.csv'\n        )\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 = 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\n            ax = fig.add_subplot(2, 3, p.label, projection='3d')\n            if hit[0]:\n                pt = xyz_predict[hit[0]]\n                ax.scatter(pt[:, 0], pt[:, 1], pt[:, 2], alpha=0.5, color='r')\n                pt = xyz_truth[hit[1]]\n                ax.scatter(pt[:,0], pt[:,1], pt[:,2], s=80, facecolors='none', edgecolors='r')\n            if fp:\n                pt = xyz_predict[fp]\n                ax.scatter(pt[:, 0], pt[:, 1], pt[:, 2], s=80, alpha=1, color='k')\n            if miss:\n                pt = xyz_truth[miss]\n                ax.scatter(pt[:, 0], pt[:, 1], pt[:, 2], s=160, alpha=1, facecolors='none', edgecolors='k')\n\n            ax.set_title(f'{p.name} ({p.difficulty})')\n\n        plt.tight_layout()\n        plt.show()\n        \n        #---\n\n        \n        zz=0","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-09T01:02:07.389754Z","iopub.execute_input":"2024-11-09T01:02:07.390082Z","iopub.status.idle":"2024-11-09T01:02:07.407107Z","shell.execute_reply.started":"2024-11-09T01:02:07.390047Z","shell.execute_reply":"2024-11-09T01:02:07.406077Z"}},"outputs":[],"execution_count":null}]}