{"metadata":{"colab":{"provenance":[],"machine_shape":"hm","gpuType":"A100"},"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"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"},"accelerator":"GPU","kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":71698,"databundleVersionId":7906362,"sourceType":"competition"}],"dockerImageVersionId":30698,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# BugNIST ML Competitions\n\nThe notebook was initially executed on Colab with GPU support.\n\n* **Training:** The models are trained on individually scanned bugs, which are easier to annotate automatically.\n* **Testing:** The models are tested on mixtures of bugs, where the context shifts but the appearance of the objects remains the same.\n","metadata":{"id":"u2xr15HHO7yE"}},{"cell_type":"markdown","source":"## Set up (SKIP IF RUN LOCALLY)","metadata":{"id":"5g3y3taURwX5"}},{"cell_type":"code","source":"!which python\n!python --version\n!pip install virtualenv\n!virtualenv myenv\n!wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh\n!chmod +x Miniconda3-latest-Linux-x86_64.sh\n!./Miniconda3-latest-Linux-x86_64.sh -b -f -p /usr/local\n!conda install -q -y --prefix /usr/local python=3.8 ujson","metadata":{"id":"Z0szcMEtBEre"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nimport os\nsys.path.append('/usr/local/lib/python3.8/site-packages/')\nos.environ['CONDA_PREFIX'] = '/usr/local/envs/myenv'","metadata":{"id":"MlZIuEIRB3eR"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install -q \"monai[all]\"\n!pip install -q \"monai-weekly[pillow, tqdm]\"\n!pip install wandb\n!pip install opencv-python","metadata":{"id":"GE1xqnkXTqHM","collapsed":true,"jupyter":{"outputs_hidden":true}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!nvcc --version\n!nvidia-smi","metadata":{"id":"SWhWbhpKG9nu","outputId":"d77e2a81-a203-4c5a-8fa1-eed463741a04"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Import private repo","metadata":{"id":"0Aoc1GwET5HN"}},{"cell_type":"markdown","source":"```python\n!wget -q https://raw.githubusercontent.com/tsunrise/colab-github/main/colab_github.py\nimport colab_github\ncolab_github.github_auth(persistent_key=True)\n\nrepo = \"x/bugnist\"\naddr = f\"git@github.com:{repo}.git\"\n!git clone $addr\n```","metadata":{}},{"cell_type":"code","source":"!mv BugNIST_DATA/test bugnist/data/BugNIST_DATA/\n!mv BugNIST_DATA/train bugnist/data/BugNIST_DATA/\n%cd bugnist","metadata":{"id":"r0RnHwoFVDxz","outputId":"4c58f91b-01a2-417e-873c-d0330cf44f8e"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data","metadata":{"id":"C6ID4gQbSVrY"}},{"cell_type":"markdown","source":"### Data Visualizations\nThis is used for final report purpose.","metadata":{"id":"-pHnTr2iW6DM"}},{"cell_type":"code","source":"import os\nimport glob\nimport matplotlib.pyplot as plt\n\ndata_dir = \"/kaggle/input/bugnist2024fgvc/BugNIST_DATA/train\"\n\nlabels = []\nlabels_count = []\n\nfor dirname in os.listdir(data_dir):\n      num_files = len(os.listdir(os.path.join(data_dir, dirname)))\n      labels.append(dirname)\n      labels_count.append(num_files)","metadata":{"id":"J-PEAS5hXYz-","execution":{"iopub.status.busy":"2024-05-15T03:53:24.882567Z","iopub.execute_input":"2024-05-15T03:53:24.883031Z","iopub.status.idle":"2024-05-15T03:53:24.899875Z","shell.execute_reply.started":"2024-05-15T03:53:24.882999Z","shell.execute_reply":"2024-05-15T03:53:24.898968Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(12, 5))\nplt.bar(range(len(labels)), labels_count, tick_label=labels, width=0.8)\n\nfor i, count in enumerate(labels_count):\n    plt.text(i, count + 0.5, str(count), ha='center', va='bottom')\n\nplt.title(\"Class distribution of training dataset\")\nplt.xlabel(\"Classes\")\nplt.ylabel(\"Count\")\nplt.xticks(rotation=45, ha=\"right\")\nplt.tight_layout()\nplt.show()","metadata":{"id":"_mfys7M8W9wN","outputId":"f11b706e-a3a9-4151-fb9f-bf628b97ae71","execution":{"iopub.status.busy":"2024-05-15T03:53:26.797036Z","iopub.execute_input":"2024-05-15T03:53:26.797997Z","iopub.status.idle":"2024-05-15T03:53:27.214585Z","shell.execute_reply.started":"2024-05-15T03:53:26.797961Z","shell.execute_reply":"2024-05-15T03:53:27.213602Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from skimage.io import imread\nfrom mpl_toolkits.mplot3d import Axes3D\nimport plotly.graph_objects as go\nfrom skimage import measure, io\nimport numpy as np\nimport random\n\ndef surfacePlot(ax, img_path, label):\n    # plot the center point of each bug\n    img = imread(img_path)\n    verts, faces, _, _ = measure.marching_cubes(img, level=42)\n\n    x, y, z = verts.T\n    ax.plot_trisurf(x, y, faces, z, color='blue', edgecolor='none', alpha=0.3)\n    ax.set_xlabel('')\n    ax.set_ylabel('')\n    ax.set_zlabel('')\n    ax.grid(False)\n    ax.set_xticks([])\n    ax.set_yticks([])\n    ax.set_zticks([])\n    ax.xaxis.pane.set_edgecolor('w')\n    ax.yaxis.pane.set_edgecolor('w')\n    ax.zaxis.pane.set_edgecolor('w')\n    ax.grid(False)\n    ax.set_axis_off()\n\n    cenx, ceny, cenz = np.mean(x), np.mean(y), np.mean(z)\n    center_point = (cenx, ceny, cenz)\n    ax.scatter(cenx, ceny, cenz, color='red', s=150)\n\n    ax.text2D(0.05, 0.95, label, transform=ax.transAxes, fontsize=12, verticalalignment='top', fontweight='bold')\n    ax.patch.set_edgecolor('black')\n    ax.patch.set_linewidth(1.5)\n\n\nfig, axs = plt.subplots(2, 6, figsize=(18, 6), subplot_kw={'projection': '3d'})\n\nrandom.seed(17)\nimg_paths = []\nfor label in labels:\n    files = os.listdir(f'/kaggle/input/bugnist2024fgvc/BugNIST_DATA/train/{label}/')\n    img_name = random.choice(files)\n    img_path = f'/kaggle/input/bugnist2024fgvc/BugNIST_DATA/train/{label}/{img_name}'\n    img_paths.append(img_path)\n\nfor ax, img_path, label in zip(axs.flat, img_paths, labels):\n    surfacePlot(ax, img_path, label)\n\nplt.tight_layout()\nplt.show()","metadata":{"id":"wQ69WXFCbUIH","outputId":"8d304d67-f6d9-45e7-b06d-18031714dc7a","execution":{"iopub.status.busy":"2024-05-15T03:53:54.506225Z","iopub.execute_input":"2024-05-15T03:53:54.507312Z","iopub.status.idle":"2024-05-15T03:54:07.661599Z","shell.execute_reply.started":"2024-05-15T03:53:54.507275Z","shell.execute_reply":"2024-05-15T03:54:07.660601Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig, axes = plt.subplots(2, 6, figsize=(12, 8))\naxes = axes.flatten()\n\nfor i, ax in enumerate(axes):\n    if i < len(img_paths):\n        img = imread(img_paths[i])\n        ax.imshow(img.max(axis=1), cmap='gray')\n        ax.axis('off')\n\nplt.tight_layout()\nplt.show()","metadata":{"id":"ESO8q5ynR3mi","outputId":"b1f67af4-0b79-43a1-bae5-a0c126232ae0","execution":{"iopub.status.busy":"2024-05-15T03:54:10.396149Z","iopub.execute_input":"2024-05-15T03:54:10.396924Z","iopub.status.idle":"2024-05-15T03:54:11.169705Z","shell.execute_reply.started":"2024-05-15T03:54:10.396891Z","shell.execute_reply":"2024-05-15T03:54:11.168769Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## img_path\nimport monai.transforms as T\nfrom dataset import CopyImaged, TiffReader\nfrom monai.transforms import MapTransform\nfrom monai.data import ImageReader\nfrom monai.data import DataLoader, Dataset, CacheDataset\nimport torch\n\n\nimg = imread(img_paths[0])\nshape = torch.tensor(img.shape)\npadding = (0,\n    int((shape[0] - shape[1]) // 2),\n    int((shape[0] - shape[2]) // 2),)\n\nkeys = [\"image\", \"mask\"]\n\ntransforms = T.Compose([\n      T.LoadImaged(\n          keys=\"image\",\n          reader=TiffReader,\n          image_only=True,\n      ),\n      T.Resized(keys=\"image\", spatial_size=shape),\n      T.ScaleIntensityd(keys=\"image\"),\n      T.BorderPadd(keys=\"image\", spatial_border=padding),\n      CopyImaged(key_to_copy=\"image\", new_key=\"mask\"),\n      T.GaussianSmoothd(keys=\"mask\", sigma=2),\n      T.AsDiscreted(\n          keys=\"mask\",\n          threshold=0.25,\n          dtype=torch.long,\n      ),\n      T.KeepLargestConnectedComponentd(keys=\"mask\", applied_labels=[0]),\n      T.EnsureTyped(\n          keys=[\"image\", \"mask\", \"label\"],\n          track_meta=False,\n      ),\n      T.RandAffined(\n          keys=[\"image\", \"mask\"],\n          prob=0.95,\n          rotate_range=(np.pi / 2,) * 3,\n          translate_range=shape // torch.tensor([2, 1, 1]),\n          padding_mode=\"zeros\",\n      ),\n      T.RandAxisFlipd(keys=keys, prob=0.5),\n      T.RandScaleIntensityd(keys=\"image\", factors=0.25, prob=0.5),\n      T.RandZoomd(keys=keys, prob=0.5),\n      T.SqueezeDimd(keys=\"mask\"),\n      T.CastToTyped(keys=\"mask\", dtype=torch.long),])\n\n\ndata = Dataset(\n    [{ 'image': f, 'label': l }\n        for f, l in zip(img_paths, [path.split('/')[-2] for path in img_paths])],\n    transform=transforms,\n)\n\ndataloader = DataLoader(\n    data,\n    shuffle=True,\n    num_workers=0\n)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image_batch = []\nmask_batch = []\nlabel_batch = []\n\nfor batch in dataloader:\n    image_batch.append(batch['image'][0][0].numpy())\n    mask_batch.append(batch['mask'][0].numpy())\n    label_batch.append(batch['label'][0])","metadata":{"id":"PzH6XcdfdI3r"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 3D\narray = image_batch[0]\nverts, faces, _, _ = measure.marching_cubes(array, level=0.5)\n\nmesh = go.Mesh3d(\n    x=verts[:, 0], y=verts[:, 1], z=verts[:, 2],\n    i=faces[:, 0], j=faces[:, 1], k=faces[:, 2],\n    opacity=0.5,\n    color='blue')\n\nx_data = mesh.x\ny_data = mesh.y\nz_data = mesh.z\n\ncenx, ceny, cenz = np.mean(x_data), np.mean(y_data), np.mean(z_data)\ncenter_point = (cenx, ceny, cenz)\n\ncenter_scatter = go.Scatter3d(\n    x=[cenx], y=[ceny], z=[cenz],\n    mode='markers',\n    marker=dict(size=5, color='red'),\n    name='Center Point'\n)\n\nfig = go.Figure(data=[mesh, center_scatter])\nfig.update_layout(\n    title='3D Surface Extraction',\n    scene=dict(\n        xaxis=dict(title='X Axis'),\n        yaxis=dict(title='Y Axis'),\n        zaxis=dict(title='Z Axis')\n    )\n)\n\nfig.show()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def surfacePlot3d(img, mask, label):\n    mask = mask.astype(bool)\n    masked_img = np.where(mask, img, 0)\n\n    verts, faces, _, _ = measure.marching_cubes(masked_img, level=0.5)\n\n    mesh = go.Mesh3d(\n        x=verts[:, 0], y=verts[:, 1], z=verts[:, 2],\n        i=faces[:, 0], j=faces[:, 1], k=faces[:, 2],\n        opacity=0.5,\n        color='blue'\n    )\n    x_data = mesh.x\n    y_data = mesh.y\n    z_data = mesh.z\n\n    cenx, ceny, cenz = np.mean(x_data), np.mean(y_data), np.mean(z_data)\n    center_point = (cenx, ceny, cenz)\n\n    center_scatter = go.Scatter3d(\n        x=[cenx], y=[ceny], z=[cenz],\n        mode='markers',\n        marker=dict(size=5, color='red'),\n        name='Center Point'\n    )\n\n    fig = go.Figure(data=[mesh, center_scatter])\n    fig.update_layout(\n        title=label,\n        scene=dict(\n            xaxis=dict(title='X Axis'),\n            yaxis=dict(title='Y Axis'),\n            zaxis=dict(title='Z Axis')\n        )\n    )\n\n    fig.show()\n\n#surfacePlot3d(image_batch[0], mask_batch[0], label_batch[0])","metadata":{"id":"XCqRsWWRlveA","outputId":"b2fd200a-3a9f-426e-8216-0420face6ecb","execution":{"iopub.status.busy":"2024-05-15T03:54:49.275392Z","iopub.execute_input":"2024-05-15T03:54:49.275977Z","iopub.status.idle":"2024-05-15T03:54:49.285526Z","shell.execute_reply.started":"2024-05-15T03:54:49.275948Z","shell.execute_reply":"2024-05-15T03:54:49.284550Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# bounding box\ndef surfacePlot1(ax, img, label):\n    verts, faces, _, _ = measure.marching_cubes(img, level=0.5)\n    x, y, z = verts.T\n    ax.plot_trisurf(x, y, faces, z, color='blue', edgecolor='none', alpha=0.3)\n\n    ax.set_xlabel('')\n    ax.set_ylabel('')\n    ax.set_zlabel('')\n    ax.grid(False)\n    ax.set_xticks([])\n    ax.set_yticks([])\n    ax.set_zticks([])\n    ax.xaxis.pane.set_edgecolor('w')\n    ax.yaxis.pane.set_edgecolor('w')\n    ax.zaxis.pane.set_edgecolor('w')\n    ax.set_axis_off()\n\n    cenx, ceny, cenz = np.mean(x), np.mean(y), np.mean(z)\n    center_point = (cenx, ceny, cenz)\n    ax.scatter(cenx, ceny, cenz, color='red', s=150)\n\n    # Calculate the bounding box\n    min_x, min_y, min_z = np.min(x), np.min(y), np.min(z)\n    max_x, max_y, max_z = np.max(x), np.max(y), np.max(z)\n\n    # Create the bounding box vertices\n    box_verts = np.array([\n        [min_x, min_y, min_z],\n        [min_x, min_y, max_z],\n        [min_x, max_y, min_z],\n        [min_x, max_y, max_z],\n        [max_x, min_y, min_z],\n        [max_x, min_y, max_z],\n        [max_x, max_y, min_z],\n        [max_x, max_y, max_z]\n    ])\n\n    # Define the edges of the bounding box\n    box_edges = [\n        [box_verts[0], box_verts[1]],\n        [box_verts[0], box_verts[2]],\n        [box_verts[0], box_verts[4]],\n        [box_verts[1], box_verts[3]],\n        [box_verts[1], box_verts[5]],\n        [box_verts[2], box_verts[3]],\n        [box_verts[2], box_verts[6]],\n        [box_verts[3], box_verts[7]],\n        [box_verts[4], box_verts[5]],\n        [box_verts[4], box_verts[6]],\n        [box_verts[5], box_verts[7]],\n        [box_verts[6], box_verts[7]]\n    ]\n\n    for edge in box_edges:\n        ax.plot3D(*zip(*edge), color=\"black\")\n\n    ax.text2D(0.05, 0.95, label, transform=ax.transAxes, fontsize=12, verticalalignment='top', fontweight='bold')\n    ax.patch.set_edgecolor('black')\n    ax.patch.set_linewidth(1.5)\n\n\"\"\"fig, axs = plt.subplots(2, 6, figsize=(18, 6), subplot_kw={'projection': '3d'})\nfor ax, img, label in zip(axs.flat, image_batch, label_batch):\n    surfacePlot1(ax, img, label)\n\nplt.tight_layout()\nplt.show()\"\"\"","metadata":{"id":"GwWtPMydU7Ek","outputId":"4d88762e-b0b3-4e92-80e3-dd489072b452","execution":{"iopub.status.busy":"2024-05-15T03:54:52.856360Z","iopub.execute_input":"2024-05-15T03:54:52.857224Z","iopub.status.idle":"2024-05-15T03:54:52.968897Z","shell.execute_reply.started":"2024-05-15T03:54:52.857191Z","shell.execute_reply":"2024-05-15T03:54:52.967838Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# mix - 3D\ndef surfacePlot2(ax, img_path):\n    img = imread(img_path)\n    verts, faces, _, _ = measure.marching_cubes(img, level=90)\n\n    x, y, z = verts.T\n    ax.plot_trisurf(x, y, faces, z, color='blue', edgecolor='none', alpha=0.3)\n    ax.set_xlabel('')\n    ax.set_ylabel('')\n    ax.set_zlabel('')\n    ax.grid(False)\n    ax.set_xticks([])\n    ax.set_yticks([])\n    ax.set_zticks([])\n    ax.xaxis.pane.set_edgecolor('w')\n    ax.yaxis.pane.set_edgecolor('w')\n    ax.zaxis.pane.set_edgecolor('w')\n    ax.grid(False)\n    ax.set_axis_off()\n    ax.patch.set_edgecolor('black')\n    ax.patch.set_linewidth(1.5)\n\nfig, axs = plt.subplots(1, 3, figsize=(16, 6), subplot_kw={'projection': '3d'})\n\nrandom.seed(7)\n\nimg_paths = []\nfor i in range(3):\n    files = os.listdir(f'/kaggle/input/bugnist2024fgvc/BugNIST_DATA/test/')\n    img_name = random.choice(files)\n    img_path = f'/kaggle/input/bugnist2024fgvc/BugNIST_DATA/test/{img_name}'\n    img_paths.append(img_path)\nfor ax, img_path in zip(axs.flat, img_paths):\n    surfacePlot2(ax, img_path)\n\nplt.tight_layout()\nplt.show()","metadata":{"id":"iRzucx0_LMNQ","outputId":"94d99589-9955-4253-ee82-aa90eba6e036","execution":{"iopub.status.busy":"2024-05-15T03:55:49.153244Z","iopub.execute_input":"2024-05-15T03:55:49.153601Z","iopub.status.idle":"2024-05-15T03:56:14.869042Z","shell.execute_reply.started":"2024-05-15T03:55:49.153573Z","shell.execute_reply":"2024-05-15T03:56:14.868040Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 2d - mix\nfig, axes = plt.subplots(1, 3, figsize=(10, 5))\n\nfor i, ax in enumerate(axes):\n    img = imread(img_paths[i])\n    ax.imshow(img.max(axis=1), cmap='gray')\n    ax.set_axis_off()\nplt.tight_layout()\nplt.show()","metadata":{"id":"RNthfLrtRn_H","outputId":"2a4d94a2-ef7c-4233-c4aa-cd4b8bbf2522","execution":{"iopub.status.busy":"2024-05-15T03:56:14.871063Z","iopub.execute_input":"2024-05-15T03:56:14.871802Z","iopub.status.idle":"2024-05-15T03:56:15.055054Z","shell.execute_reply.started":"2024-05-15T03:56:14.871744Z","shell.execute_reply":"2024-05-15T03:56:15.053945Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 3D - mix\nimg = imread(img_paths[0])\nverts, faces, _, _ = measure.marching_cubes(img, level=90)\n\nmesh = go.Mesh3d(\n    x=verts[:, 0], y=verts[:, 1], z=verts[:, 2],\n    i=faces[:, 0], j=faces[:, 1], k=faces[:, 2],\n    opacity=0.5,\n    color='blue')\n\nx_data = mesh.x\ny_data = mesh.y\nz_data = mesh.z\n\nfig = go.Figure(data=[mesh])\nfig.update_layout(\n    title='3D Surface Extraction',\n    scene=dict(\n        xaxis=dict(title='X Axis'),\n        yaxis=dict(title='Y Axis'),\n        zaxis=dict(title='Z Axis')\n    )\n)\n\nfig.show()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Classifications Experiments\n - ResNet\n - DenseNet\n - Vision Transformers from `MONAI`","metadata":{"id":"O7fm-NIaSmXt"}},{"cell_type":"code","source":"%cd bugnist","metadata":{"id":"ls3n1v5HMyVA","outputId":"e9691b52-7d15-4791-e2d7-0bf9fe5eff93"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import dataset\nfrom sklearn.model_selection import train_test_split\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import DataLoader, Subset\nfrom torchvision import models, transforms\nfrom tqdm import tqdm\n\ndata_dir = \"data/BugNIST_DATA\"\ndset = dataset.BugNist(root=data_dir, split=\"train\")\n\ntrain_indices, test_indices = train_test_split(\n    range(len(dset)),\n    test_size=0.2,\n    stratify=dset.all_labels,\n    random_state=17\n)\n\ntrain_subset = Subset(dset, train_indices)\ntest_subset = Subset(dset, test_indices)","metadata":{"id":"di5Ztu7CSWXM"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### ResNet\n\ncontinue..","metadata":{"id":"ht2fv0absqgA"}},{"cell_type":"markdown","source":"### DenseNet","metadata":{"id":"SKpVYRBbtDBK"}},{"cell_type":"code","source":"import logging\nimport os\nimport sys\nimport shutil\nimport tempfile\n\nimport matplotlib.pyplot as plt\nimport torch\nfrom torch.utils.tensorboard import SummaryWriter\nimport monai\nfrom monai.apps import download_and_extract\nfrom monai.config import print_config\nfrom monai.data import DataLoader, ImageDataset\n\npin_memory = torch.cuda.is_available()\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nlogging.basicConfig(stream=sys.stdout, level=logging.INFO)\nprint_config()","metadata":{"id":"ZjIiaqwuFtu_","outputId":"ed236149-e590-4430-c273-dc087187dcd6"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_loader = DataLoader(train_subset, batch_size=4, shuffle=True)\ntest_loader = DataLoader(test_subset, batch_size=4, shuffle=False)","metadata":{"id":"PFGUyJfmFyt8"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = monai.networks.nets.DenseNet121(spatial_dims=3, in_channels=1, out_channels=2).to(device)\nloss_function = torch.nn.CrossEntropyLoss()\noptimizer = torch.optim.Adam(model.parameters(), 1e-4)","metadata":{"id":"CivbFbU2CycQ"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_interval = 2\nbest_metric = -1\nbest_metric_epoch = -1\nepoch_loss_values = []\nmetric_values = []\nwriter = SummaryWriter()\nmax_epochs = 100","metadata":{"id":"pIfnZWQvowzD"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.environ['CUDA_LAUNCH_BLOCKING'] = \"1\"","metadata":{"id":"_LrpnXh1LW0U"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for epoch in range(max_epochs):\n    print(\"-\" * 10)\n    print(f\"epoch {epoch + 1}/{max_epochs}\")\n    model.train()\n    epoch_loss = 0\n    step = 0\n\n    for batch_data in train_loader:\n        step += 1\n        inputs, labels = batch_data[\"image\"].to(device), batch_data[\"label\"].to(device)\n        optimizer.zero_grad()\n        outputs = model(inputs)\n        loss = loss_function(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        epoch_loss += loss.item()\n        epoch_len = len(train_loader) // train_loader.batch_size\n        print(f\"{step}/{epoch_len}, train_loss: {loss.item():.4f}\")\n        writer.add_scalar(\"train_loss\", loss.item(), epoch_len * epoch + step)\n\n    epoch_loss /= step\n    epoch_loss_values.append(epoch_loss)\n    print(f\"epoch {epoch + 1} average loss: {epoch_loss:.4f}\")\n\n    if (epoch + 1) % val_interval == 0:\n        model.eval()\n\n        num_correct = 0.0\n        metric_count = 0\n        for val_data in val_loader:\n            val_images, val_labels = val_data[\"image\"].to(device), val_data[\"label\"].to(device)\n            with torch.no_grad():\n                val_outputs = model(val_images)\n                value = torch.eq(val_outputs.argmax(dim=1), val_labels.argmax(dim=1))\n                metric_count += len(value)\n                num_correct += value.sum().item()\n\n        metric = num_correct / metric_count\n        metric_values.append(metric)\n\n        if metric > best_metric:\n            best_metric = metric\n            best_metric_epoch = epoch + 1\n            torch.save(model.state_dict(), \"best_metric_model_classification3d_array.pth\")\n            print(\"saved new best metric model\")\n\n        print(f\"Current epoch: {epoch+1} current accuracy: {metric:.4f} \")\n        print(f\"Best accuracy: {best_metric:.4f} at epoch {best_metric_epoch}\")\n        writer.add_scalar(\"val_accuracy\", metric, epoch + 1)\n\nprint(f\"Training completed, best_metric: {best_metric:.4f} at epoch: {best_metric_epoch}\")","metadata":{"id":"B2hZx93nDMNP"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Detections Experiments\n- U-Net\n- NNDetetions","metadata":{"id":"5JkEihY9StNv"}},{"cell_type":"code","source":"# 3D U-Net\n!python train.py --data_dir data/BugNIST_DATA --batch_size 16 --epochs 500 --lr 0.001","metadata":{"id":"AwGvZ4GNJGhd"},"execution_count":null,"outputs":[]}]}