{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","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"}],"dockerImageVersionId":30839,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install tensorflow-probability\n!pip install imageio\n!pip install git+https://github.com/tensorflow/docs\n!pip install zarr\n!pip install fsspec\n!pip install pydantic\n!pip install trimesh\n!pip install \"copick[all]\"\n!pip install copick git+https://github.com/copick/copick-utils.git git+https://github.com/copick/DeepFindET.git","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-02-13T02:54:17.883876Z","iopub.execute_input":"2025-02-13T02:54:17.884234Z","iopub.status.idle":"2025-02-13T02:55:05.686636Z","shell.execute_reply.started":"2025-02-13T02:54:17.884203Z","shell.execute_reply":"2025-02-13T02:55:05.685538Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"config_blob = \"\"\"{\n    \"name\": \"czii_cryoet_mlchallenge_2024\",\n    \"description\": \"2024 CZII CryoET ML Challenge training data.\",\n    \"version\": \"1.0.0\",\n\n    \"pickable_objects\": [\n        {\n            \"name\": \"apo-ferritin\",\n            \"is_particle\": true,\n            \"pdb_id\": \"4V1W\",\n            \"label\": 1,\n            \"color\": [  0, 117, 220, 128],\n            \"radius\": 60,\n            \"map_threshold\": 0.0418\n        },\n        {\n            \"name\": \"beta-amylase\",\n            \"is_particle\": true,\n            \"pdb_id\": \"1FA2\",\n            \"label\": 2,\n            \"color\": [153,  63,   0, 128],\n            \"radius\": 65,\n            \"map_threshold\": 0.035\n        },\n        {\n            \"name\": \"beta-galactosidase\",\n            \"is_particle\": true,\n            \"pdb_id\": \"6X1Q\",\n            \"label\": 3,\n            \"color\": [ 76,   0,  92, 128],\n            \"radius\": 90,\n            \"map_threshold\": 0.0578\n        },\n        {\n            \"name\": \"ribosome\",\n            \"is_particle\": true,\n            \"pdb_id\": \"6EK0\",\n            \"label\": 4,\n            \"color\": [  0,  92,  49, 128],\n            \"radius\": 150,\n            \"map_threshold\": 0.0374\n        },\n        {\n            \"name\": \"thyroglobulin\",\n            \"is_particle\": true,\n            \"pdb_id\": \"6SCJ\",\n            \"label\": 5,\n            \"color\": [ 43, 206,  72, 128],\n            \"radius\": 130,\n            \"map_threshold\": 0.0278\n        },\n        {\n            \"name\": \"virus-like-particle\",\n            \"is_particle\": true,\n            \"label\": 6,\n            \"color\": [255, 204, 153, 128],\n            \"radius\": 135,\n            \"map_threshold\": 0.201\n        },\n        {\n            \"name\": \"membrane\",\n            \"is_particle\": false,\n            \"label\": 8,\n            \"color\": [100, 100, 100, 128]\n        },\n        {\n            \"name\": \"background\",\n            \"is_particle\": false,\n            \"label\": 9,\n            \"color\": [10, 150, 200, 128]\n        }\n    ],\n\n    \"overlay_root\": \"/kaggle/working/overlay\",\n\n    \"overlay_fs_args\": {\n        \"auto_mkdir\": true\n    },\n\n    \"static_root\": \"/kaggle/input/czii-cryo-et-object-identification/train/static\"\n}\"\"\"\n\ncopick_config_path = \"/kaggle/working/copick.config\"\noutput_overlay = \"/kaggle/working/overlay\"\n\nwith open(copick_config_path, \"w\") as f:\n    f.write(config_blob)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-13T02:55:05.687989Z","iopub.execute_input":"2025-02-13T02:55:05.688280Z","iopub.status.idle":"2025-02-13T02:55:05.693640Z","shell.execute_reply.started":"2025-02-13T02:55:05.688256Z","shell.execute_reply":"2025-02-13T02:55:05.692827Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Setup new overlay directory\nimport os\nimport shutil\n\n# Define source and destination directories\nsource_dir = '/kaggle/input/czii-cryo-et-object-identification/train/overlay'\ndestination_dir = '/kaggle/working/overlay'\n\n# Walk through the source directory\nfor root, dirs, files in os.walk(source_dir):\n    # Create corresponding subdirectories in the destination\n    relative_path = os.path.relpath(root, source_dir)\n    target_dir = os.path.join(destination_dir, relative_path)\n    os.makedirs(target_dir, exist_ok=True)\n    \n    # Copy and rename each file\n    for file in files:\n        if file.startswith(\"curation_0_\"):\n            new_filename = file\n        else:\n            new_filename = f\"curation_0_{file}\"\n            \n        \n        # Define full paths for the source and destination files\n        source_file = os.path.join(root, file)\n        destination_file = os.path.join(target_dir, new_filename)\n        \n        # Copy the file with the new name\n        shutil.copy2(source_file, destination_file)\n        print(f\"Copied {source_file} to {destination_file}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-13T02:55:05.695433Z","iopub.execute_input":"2025-02-13T02:55:05.695708Z","iopub.status.idle":"2025-02-13T02:55:05.796741Z","shell.execute_reply.started":"2025-02-13T02:55:05.695686Z","shell.execute_reply":"2025-02-13T02:55:05.795864Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from deepfindET.entry_points import step1\nfrom deepfindET.utils import copick_tools\nimport matplotlib.pyplot as plt\nimport copick\n\n%matplotlib inline\n\n################## Input Parameters #################\n\n# Config File\nconfig = '/kaggle/working/copick.config'\n\n# Query Tomogram\nvoxel_size = 10 \ntomogram_algorithm = 'denoised'\n\n# Output Name for the Segmentation Targets\nout_name = 'remotetargets'\nout_user_id = 'deepfindET'\nout_session_id = '0'\n\n# Read Copick Directory\ncopickRoot = copick.from_file(config)\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-13T02:55:05.798005Z","iopub.execute_input":"2025-02-13T02:55:05.798297Z","iopub.status.idle":"2025-02-13T02:55:05.805592Z","shell.execute_reply.started":"2025-02-13T02:55:05.798274Z","shell.execute_reply":"2025-02-13T02:55:05.804747Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"[(obj.name, None, None, (obj.radius / voxel_size)) for obj in copickRoot.pickable_objects if obj.is_particle]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-13T02:55:05.806478Z","iopub.execute_input":"2025-02-13T02:55:05.806825Z","iopub.status.idle":"2025-02-13T02:55:05.825559Z","shell.execute_reply.started":"2025-02-13T02:55:05.806789Z","shell.execute_reply":"2025-02-13T02:55:05.824657Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Query Train Protein Coordiantes and any Associated Segmentations\ntrain_targets = {}\n\n# Define protein targets with their respective radii\n# We can Provide two forms of inputs, either \n# ('protein-name',radius) or ('protein-name', 'user-id', 'session-id', 'radius')\ntargets = [(obj.name, None, None, (obj.radius / voxel_size)) for obj in copickRoot.pickable_objects if obj.is_particle]\n\n# Set run_ids to None, indicating that targets will be generated for the entire CoPick project by default.\n# If specific Run-IDs were provided, this variable would contain a list of those IDs.\nrun_ids = None","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-13T02:55:05.826331Z","iopub.execute_input":"2025-02-13T02:55:05.826592Z","iopub.status.idle":"2025-02-13T02:55:05.838909Z","shell.execute_reply.started":"2025-02-13T02:55:05.826569Z","shell.execute_reply":"2025-02-13T02:55:05.838198Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Generate train target information\nfor t in targets:\n    obj_name, user_id, session_id, radius = t\n    info = {\n        \"label\": copickRoot.get_object(obj_name).label,\n        \"user_id\": user_id,\n        \"session_id\": session_id,\n        \"radius\": radius,\n        \"is_particle_target\": True,\n    }\n    train_targets[obj_name] = info\n\n\n# Define segmentation target (e.g., membrane)\nseg_targets = [('membrane', None, None)]\n\n# Generate segmentation target information\nfor s in seg_targets:\n    obj_name, user_id, session_id = s\n    info = {\n        \"label\": copickRoot.get_object(obj_name).label,\n        \"user_id\": user_id,\n        \"session_id\": session_id,\n        \"radius\": None,       \n        \"is_particle_target\": False,                 \n    }\n    train_targets[obj_name] = info\n\n# Call the create_train_targets function from step1 to generate the training targets for the 3D U-Net model.\n# The function will use the parameters defined in the previous cells and the following inputs:\nstep1.create_train_targets(\n    config,              # The configuration file path specifying various settings and parameters for the project.\n    train_targets,       # A dictionary containing the target information for each protein or object to be segmented.\n    run_ids,             # The list of Run-IDs for which to generate targets. None means targets for the entire project.\n    voxel_size,          # The voxel size to be used in the tomogram data.\n    tomogram_algorithm,  # The reconstruction algorithm used for the tomograms, e.g., 'wbp' (weighted back projection).\n    out_name,            # The output name for the generated segmentation targets.\n    out_user_id,         # The user ID under which the output targets will be saved.\n    out_session_id,      # The session ID associated with the output, typically used for tracking purposes.\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-13T02:55:05.839634Z","iopub.execute_input":"2025-02-13T02:55:05.839876Z","iopub.status.idle":"2025-02-13T02:55:18.390506Z","shell.execute_reply.started":"2025-02-13T02:55:05.839855Z","shell.execute_reply":"2025-02-13T02:55:18.389816Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Option 1: Query All RunIDs\n# Retrieve all available Run-IDs from the CoPick project. This generates a list of Run-IDs by iterating over all runs in copickRoot.\nrun_ids = [run.name for run in copickRoot.runs]\n\n# Option 2: Manually Specify Specific Run\n# Define a specific Run-ID manually. This is useful for extracting volumes for a specific run.\nrunID = 'TS_6_4'\n\n# Retrieve the specific run object from CoPick using the manually specified Run-ID.\ncopick_run = copickRoot.get_run(runID)\n\n# Extract the segmentation target associated with the specified run.\n# The function get_copick_segmentation retrieves the segmentation data (e.g., target volume) based on the run object,\n# segmentation name, user ID, and session ID.\ntrain_target = copick_tools.get_copick_segmentation(\n    copick_run,                 # The run object obtained from CoPick for the specific Run-ID.\n    segmentationName='remotetargets',  # The name of the segmentation target to retrieve.\n    userID='deepfindET',        # The user ID under which the segmentation data is saved.\n    sessionID='0'               # The session ID associated with the segmentation data.\n)\n\n# Retrieve the tomogram associated with the specified Run-ID from the CoPick project.\n# The function get_copick_tomogram extracts the tomogram data, using the voxel size, algorithm, and Run-ID.\ntrain_tomogram = copick_tools.get_copick_tomogram(\n    copickRoot,                 # The root object for the CoPick project, containing all runs and associated data.\n    voxelSize=voxel_size,       # The voxel size to be used for retrieving the tomogram.\n    tomoAlgorithm='wbp',        # The reconstruction algorithm used for the tomogram, e.g., 'wbp' (weighted back projection).\n    tomoID=runID                # The specific Run-ID for which the tomogram is being retrieved.\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-13T02:55:18.393837Z","iopub.execute_input":"2025-02-13T02:55:18.394076Z","iopub.status.idle":"2025-02-13T02:55:18.414642Z","shell.execute_reply.started":"2025-02-13T02:55:18.394055Z","shell.execute_reply":"2025-02-13T02:55:18.414063Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from deepfindET.entry_points import step1\nfrom deepfindET.utils import copick_tools\nimport matplotlib.pyplot as plt\nimport copick\n\n%matplotlib inline\n\n################## Input Parameters #################\n\n# Config File\nconfig = '/kaggle/working/copick.config'\n\n# Query Tomogram\nvoxel_size = 10 \ntomogram_algorithm = 'denoised'\n\n# Output Name for the Segmentation Targets\nout_name = 'remotetargets'\nout_user_id = 'deepfindET'\nout_session_id = '0'\n\n# Read Copick Directory\ncopickRoot = copick.from_file(config)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-13T02:55:18.416203Z","iopub.execute_input":"2025-02-13T02:55:18.416412Z","iopub.status.idle":"2025-02-13T02:55:18.422591Z","shell.execute_reply.started":"2025-02-13T02:55:18.416393Z","shell.execute_reply":"2025-02-13T02:55:18.421720Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Plot the images\nplt.figure(figsize=(15, 5))\n\n# Original Image\nplt.subplot(1, 2, 1)\nplt.title('Tomogram')\nplt.imshow(train_tomogram[90,],cmap='gray')\nprint(train_tomogram[90,])\nplt.axis('off')\n\n# Original Image\nplt.subplot(1, 2, 2)\nplt.title('Train Target')\nplt.imshow(train_target[90,])\nplt.axis('off')\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-13T02:55:18.423515Z","iopub.execute_input":"2025-02-13T02:55:18.423857Z","iopub.status.idle":"2025-02-13T02:55:19.805510Z","shell.execute_reply.started":"2025-02-13T02:55:18.423824Z","shell.execute_reply":"2025-02-13T02:55:19.804661Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from IPython import display\n\nimport glob\nimport imageio\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport PIL\nimport tensorflow as tf\nimport tensorflow_probability as tfp\nimport time","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-13T02:55:19.806391Z","iopub.execute_input":"2025-02-13T02:55:19.806682Z","iopub.status.idle":"2025-02-13T02:55:19.811447Z","shell.execute_reply.started":"2025-02-13T02:55:19.806654Z","shell.execute_reply":"2025-02-13T02:55:19.810729Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import tensorflow as tf\nimport numpy as np\n\ntrain_runs = []\nfor runID in run_ids:\n    copick_run = copickRoot.get_run(runID)\n    train_target = copick_tools.get_copick_segmentation(\n        copick_run,                 # The run object obtained from CoPick for the specific Run-ID.\n        segmentationName='remotetargets',  # The name of the segmentation target to retrieve.\n        userID='deepfindET',        # The user ID under which the segmentation data is saved.\n        sessionID='0'               # The session ID associated with the segmentation data.\n    )\n    train_tomogram = copick_tools.get_copick_tomogram(\n        copickRoot,                 # The root object for the CoPick project, containing all runs and associated data.\n        voxelSize=voxel_size,       # The voxel size to be used for retrieving the tomogram.\n        tomoAlgorithm='wbp',        # The reconstruction algorithm used for the tomogram, e.g., 'wbp' (weighted back projection).\n        tomoID=runID                # The specific Run-ID for which the tomogram is being retrieved.\n    )\n    tf_tomogram = np.array(train_tomogram[:])\n    train_runs.append(tf_tomogram)\n    print(tf_tomogram.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-13T02:55:19.812316Z","iopub.execute_input":"2025-02-13T02:55:19.812590Z","iopub.status.idle":"2025-02-13T02:55:25.086516Z","shell.execute_reply.started":"2025-02-13T02:55:19.812563Z","shell.execute_reply":"2025-02-13T02:55:25.085460Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"new_train_runs = []\nfor run in train_runs:\n    for image in run:\n        new_train_runs.append(image)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-13T02:55:25.087439Z","iopub.execute_input":"2025-02-13T02:55:25.087688Z","iopub.status.idle":"2025-02-13T02:55:25.091911Z","shell.execute_reply.started":"2025-02-13T02:55:25.087665Z","shell.execute_reply":"2025-02-13T02:55:25.091104Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(np.array(new_train_runs).shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-13T02:55:25.092680Z","iopub.execute_input":"2025-02-13T02:55:25.092923Z","iopub.status.idle":"2025-02-13T02:55:25.717981Z","shell.execute_reply.started":"2025-02-13T02:55:25.092902Z","shell.execute_reply":"2025-02-13T02:55:25.716965Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_size = 1000\nbatch_size = 16\ntest_size = 288","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-13T02:55:25.718976Z","iopub.execute_input":"2025-02-13T02:55:25.719342Z","iopub.status.idle":"2025-02-13T02:55:25.723190Z","shell.execute_reply.started":"2025-02-13T02:55:25.719306Z","shell.execute_reply":"2025-02-13T02:55:25.722322Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_dataset = new_train_runs[0:train_size]\ntest_dataset = new_train_runs[train_size:]\n\nimport tensorflow as tf\nimport cv2\ndef normalize_image(image):\n    \"\"\"Ensure image is in the range [0, 1] regardless of its original scale.\"\"\"\n    # If the image data is already in the expected range [0, 1], no changes are needed.\n    image = (image+1)/2\n#     image_min = np.min(image)\n#     image_max = np.max(image)\n\n# # Normalize to [0, 1]\n#     image_normalized = (image - image_min) / (image_max - image_min)\n    return image\n\ntrain_dataset_normalized = [normalize_image(img[:256, :256]) for img in train_dataset]\ntest_dataset_normalized = [normalize_image(img[:256, :256]) for img in test_dataset]\ntrain_dataset_normalized = np.expand_dims(train_dataset_normalized, axis=1)\ntest_dataset_normalized = np.expand_dims(test_dataset_normalized, axis=1)\n\n\n# # Fill the new dimension with 0.5\n# train_dataset_normalized[0] = 0.5\n# test_dataset_normalized[0] = 0.5","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-13T03:31:46.470893Z","iopub.execute_input":"2025-02-13T03:31:46.471207Z","iopub.status.idle":"2025-02-13T03:31:47.010219Z","shell.execute_reply.started":"2025-02-13T03:31:46.471183Z","shell.execute_reply":"2025-02-13T03:31:47.009492Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(train_dataset_normalized.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-13T03:31:51.546084Z","iopub.execute_input":"2025-02-13T03:31:51.546388Z","iopub.status.idle":"2025-02-13T03:31:51.551264Z","shell.execute_reply.started":"2025-02-13T03:31:51.546363Z","shell.execute_reply":"2025-02-13T03:31:51.550134Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(train_dataset_normalized[0])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-13T03:31:53.982004Z","iopub.execute_input":"2025-02-13T03:31:53.982332Z","iopub.status.idle":"2025-02-13T03:31:53.987550Z","shell.execute_reply.started":"2025-02-13T03:31:53.982304Z","shell.execute_reply":"2025-02-13T03:31:53.986712Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.imshow(train_dataset_normalized[0][0], cmap='gray')  # 'gray' colormap for grayscale\nplt.axis('off')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-13T03:31:58.168478Z","iopub.execute_input":"2025-02-13T03:31:58.168845Z","iopub.status.idle":"2025-02-13T03:31:58.288380Z","shell.execute_reply.started":"2025-02-13T03:31:58.168812Z","shell.execute_reply":"2025-02-13T03:31:58.287568Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pip install torchgan==0.6.0 torchvision==0.14.1","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-13T02:55:26.016656Z","iopub.execute_input":"2025-02-13T02:55:26.016933Z","iopub.status.idle":"2025-02-13T02:55:26.978319Z","shell.execute_reply.started":"2025-02-13T02:55:26.016910Z","shell.execute_reply":"2025-02-13T02:55:26.977198Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nfrom torch.utils.data import Dataset, DataLoader\nimport numpy as np\nimport torch.nn.functional as F\n\n\nclass NumpyDataset(Dataset):\n    \"\"\"Custom Dataset for loading normalized NumPy arrays.\"\"\"\n    def __init__(self, dataset):\n        self.dataset = dataset\n        self.shape = dataset.shape\n\n    def __len__(self):\n        return len(self.dataset)\n\n    def __getitem__(self, idx):\n        # Return the image at index idx\n        image = self.dataset[idx]\n        # Convert to torch tensor\n        image_tensor = torch.tensor(image, dtype=torch.float32)\n        return image_tensor\n\n\n# Convert your normalized train and test datasets to Dataset objects\ntrain_dataset_pytorch = NumpyDataset(train_dataset_normalized)\ntest_dataset_pytorch = NumpyDataset(test_dataset_normalized)\nprint(train_dataset_pytorch.shape)\n# Create DataLoader for train and test datasets\ntrain_loader = DataLoader(train_dataset_pytorch, batch_size=64, shuffle=True)\ntest_loader = DataLoader(test_dataset_pytorch, batch_size=64, shuffle=False)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-13T03:32:11.470489Z","iopub.execute_input":"2025-02-13T03:32:11.470851Z","iopub.status.idle":"2025-02-13T03:32:11.477720Z","shell.execute_reply.started":"2025-02-13T03:32:11.470821Z","shell.execute_reply":"2025-02-13T03:32:11.476861Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torchgan.models import DCGANGenerator, DCGANDiscriminator\nfrom torchgan.losses import MinimaxGeneratorLoss, MinimaxDiscriminatorLoss\nfrom torchgan.trainer import Trainer\n\nlatent_dim = 128\nimage_size = 256\nimage_channels = 1\n\ngenerator = DCGANGenerator(\n    encoding_dims=latent_dim,\n    out_size=image_size,\n    out_channels=image_channels,\n    last_nonlinearity=torch.nn.Sigmoid()\n)\n\ndiscriminator = DCGANDiscriminator(\n    in_size=image_size,\n    in_channels=image_channels,\n    last_nonlinearity=torch.nn.Sigmoid()\n)\n\n# Define the optimizers for the generator and discriminator\noptimizer_generator = optim.Adam(generator.parameters(), lr=0.0002, betas=(0.5, 0.999))\noptimizer_discriminator = optim.Adam(discriminator.parameters(), lr=0.0002, betas=(0.5, 0.999))\n\n# Define the loss functions\nlosses = [MinimaxGeneratorLoss(), MinimaxDiscriminatorLoss()]\n\n# Define the model dictionary\nmodels = {\n    \"generator\": {\n        \"name\": DCGANGenerator,\n        \"args\": {\"out_channels\": image_channels, \"out_size\": image_size, \"step_channels\": 64},\n        \"optimizer\": {\"name\": optim.Adam, \"args\": {\"lr\": 0.0002, \"betas\": (0.5, 0.999)}}\n    },\n    \"discriminator\": {\n        \"name\": DCGANDiscriminator,\n        \"args\": {\"in_channels\": image_channels, \"in_size\": image_size, \"step_channels\": 64},\n        \"optimizer\": {\"name\": optim.Adam, \"args\": {\"lr\": 0.0002, \"betas\": (0.5, 0.999)}}\n    }\n}\n\n# Define the devices (for multi-GPU use, list all available GPUs, e.g., [0, 1])\ndevices = [0]  # Change to your device IDs (use [0] for CPU or a single GPU)\n\n# Define the Trainer with other necessary parameters\ntrainer = Trainer(\n    models,\n    losses,\n    sample_size=64,\n    epochs=1000,\n    ncritic=1,  # Train the generator and discriminator equally\n    retain_checkpoints=3,  # Retain the last 3 checkpoints\n    nrow=8,  # Arrange the generated images in 8 rows\n    \n)\n\n# Assuming you have a DataLoader `train_loader` for your dataset\ntrainer.train(train_loader)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-13T03:32:14.137824Z","iopub.execute_input":"2025-02-13T03:32:14.138227Z","execution_failed":"2025-02-13T06:40:27.457Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"NO NEED TO RUN AFTER THIS POINT","metadata":{}}]}