{"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":"none","dataSources":[{"sourceId":84969,"databundleVersionId":10033515,"sourceType":"competition"}],"dockerImageVersionId":30822,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install zarr\n# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np  # linear algebra\nimport pandas as pd  # data processing, CSV file I/O (e.g. pd.read_csv)\nimport zarr\nimport os\nimport matplotlib.pyplot as plt\nimport joblib\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport pickle\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2024-12-22T09:00:16.309304Z","iopub.execute_input":"2024-12-22T09:00:16.309754Z","iopub.status.idle":"2024-12-22T09:00:20.837312Z","shell.execute_reply.started":"2024-12-22T09:00:16.309718Z","shell.execute_reply":"2024-12-22T09:00:20.836093Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 2: Set paths\ntrain_path = \"/kaggle/input/czii-cryo-et-object-identification/train\"\ntest_path = \"/kaggle/input/czii-cryo-et-object-identification/test\"\nsubmission_path = \"/kaggle/input/czii-cryo-et-object-identification/sample_submission.csv\"\nmodel_path = \"/kaggle/working/trained_model.pkl\"  # Path for the pickled model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T09:00:20.839326Z","iopub.execute_input":"2024-12-22T09:00:20.839763Z","iopub.status.idle":"2024-12-22T09:00:20.844566Z","shell.execute_reply.started":"2024-12-22T09:00:20.839729Z","shell.execute_reply":"2024-12-22T09:00:20.843301Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 3: Load data function\ndef load_data(path):\n    \"\"\"Load zarr files from the specified path.\"\"\"\n    return zarr.open(path, mode='r')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T09:00:20.846299Z","iopub.execute_input":"2024-12-22T09:00:20.846577Z","iopub.status.idle":"2024-12-22T09:00:20.861614Z","shell.execute_reply.started":"2024-12-22T09:00:20.846548Z","shell.execute_reply":"2024-12-22T09:00:20.860627Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 4: Define Dataset Class\nclass CryoETDataset(Dataset):\n    def __init__(self, path):\n        self.path = path\n        self.data = self.load_data()\n\n    def load_data(self):\n        \"\"\"Load zarr files from the specified path.\"\"\"\n        return zarr.open(self.path, mode='r')\n\n    def __len__(self):\n        return len(self.data)\n\n    def __getitem__(self, idx):\n        tomogram = self.data[idx]\n        tomogram_normalized = tomogram / np.max(tomogram)  # Normalize\n        tomogram_tensor = torch.tensor(tomogram_normalized, dtype=torch.float32)\n        return tomogram_tensor","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T09:00:20.863139Z","iopub.execute_input":"2024-12-22T09:00:20.863541Z","iopub.status.idle":"2024-12-22T09:00:20.872441Z","shell.execute_reply.started":"2024-12-22T09:00:20.863480Z","shell.execute_reply":"2024-12-22T09:00:20.871425Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 5: Define the Model\nclass UNet(nn.Module):\n    def __init__(self):\n        super(UNet, self).__init__()\n        # Define the layers of the U-Net model\n        self.encoder = nn.Sequential(\n            nn.Conv3d(1, 64, kernel_size=3, padding=1),\n            nn.ReLU(inplace=True),\n            nn.MaxPool3d(kernel_size=2, stride=2)\n        )\n        self.decoder = nn.Sequential(\n            nn.ConvTranspose3d(64, 1, kernel_size=2, stride=2),\n            nn.ReLU(inplace=True)\n        )\n\n    def forward(self, x):\n        x = self.encoder(x)\n        x = self.decoder(x)\n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T09:00:20.873582Z","iopub.execute_input":"2024-12-22T09:00:20.873879Z","iopub.status.idle":"2024-12-22T09:00:20.893418Z","shell.execute_reply.started":"2024-12-22T09:00:20.873847Z","shell.execute_reply":"2024-12-22T09:00:20.892365Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 6: Training Function\ndef train_model(model, dataloader, criterion, optimizer, num_epochs=10):\n    model.train()\n    for epoch in range(num_epochs):\n        for inputs in dataloader:\n            optimizer.zero_grad()\n            outputs = model(inputs.unsqueeze(1))  # Add channel dimension\n            loss = criterion(outputs, inputs.unsqueeze(1))  # Assuming reconstruction task\n            loss.backward()\n            optimizer.step()\n        print(f'Epoch [{epoch+1}/{num_epochs}], Loss: {loss.item():.4f}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T09:00:20.894563Z","iopub.execute_input":"2024-12-22T09:00:20.894974Z","iopub.status.idle":"2024-12-22T09:00:20.908161Z","shell.execute_reply.started":"2024-12-22T09:00:20.894926Z","shell.execute_reply":"2024-12-22T09:00:20.907107Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 7: Evaluation Function\ndef evaluate_model(model, dataloader, criterion):\n    model.eval()\n    total_loss = 0\n    with torch.no_grad():\n        for inputs in dataloader:\n            outputs = model(inputs.unsqueeze(1))\n            loss = criterion(outputs, inputs.unsqueeze(1))\n            total_loss += loss.item()\n    average_loss = total_loss / len(dataloader)\n    print(f'Validation Loss: {average_loss:.4f}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T09:00:20.909239Z","iopub.execute_input":"2024-12-22T09:00:20.909626Z","iopub.status.idle":"2024-12-22T09:00:20.923750Z","shell.execute_reply.started":"2024-12-22T09:00:20.909596Z","shell.execute_reply":"2024-12-22T09:00:20.922681Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 8: Save the Model\ndef save_model(model, filename):\n    \"\"\"Save the trained model.\"\"\"\n    torch.save(model.state_dict(), filename)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T09:00:20.926520Z","iopub.execute_input":"2024-12-22T09:00:20.926836Z","iopub.status.idle":"2024-12-22T09:00:20.946852Z","shell.execute_reply.started":"2024-12-22T09:00:20.926808Z","shell.execute_reply":"2024-12-22T09:00:20.945896Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 9: Visualization Function\ndef visualize_predictions(predictions):\n    \"\"\"Visualize the predicted tomogram slices.\"\"\"\n    plt.figure(figsize=(10, 10))\n    # Example: Plotting the first tomogram slice\n    plt.imshow(predictions[0].cpu().numpy(), cmap='gray')  # Ensure to move tensor to CPU for numpy conversion\n    plt.title('Predicted Tomogram Slice')\n    plt.axis('off')\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T09:00:20.948580Z","iopub.execute_input":"2024-12-22T09:00:20.948872Z","iopub.status.idle":"2024-12-22T09:00:20.964128Z","shell.execute_reply.started":"2024-12-22T09:00:20.948845Z","shell.execute_reply":"2024-12-22T09:00:20.963082Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 10: Implementing a simple model evaluation metric - F-beta score\ndef f_beta_score(y_true, y_pred, beta=4):\n    \"\"\"Calculate F-beta score based on true and predicted values.\"\"\"\n    precision = np.sum(y_pred * y_true) / (np.sum(y_pred) + 1e-6)\n    recall = np.sum(y_pred * y_true) / (np.sum(y_true) + 1e-6)\n    f_beta = (1 + beta**2) * (precision * recall) / (beta**2 * precision + recall + 1e-6)\n    return f_beta","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T09:00:20.965271Z","iopub.execute_input":"2024-12-22T09:00:20.965626Z","iopub.status.idle":"2024-12-22T09:00:20.979626Z","shell.execute_reply.started":"2024-12-22T09:00:20.965583Z","shell.execute_reply":"2024-12-22T09:00:20.978577Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 11: Main Execution\nif __name__ == \"__main__\":\n    # Construct the training data path\n    train_data_path = os.path.join(train_path, \"static\", \"ExperimentRuns\")\n    \n    # Check if the path exists\n    if not os.path.exists(train_data_path):\n        raise FileNotFoundError(f\"Training data path does not exist: {train_data_path}\")\n    else:\n        print(f\"Training data path exists: {train_data_path}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T09:00:20.980890Z","iopub.execute_input":"2024-12-22T09:00:20.981262Z","iopub.status.idle":"2024-12-22T09:00:20.998309Z","shell.execute_reply.started":"2024-12-22T09:00:20.981223Z","shell.execute_reply":"2024-12-22T09:00:20.997219Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 12: Load and preprocess training data\ntry:\n    train_dataset = CryoETDataset(train_data_path)\n    print(f\"Loaded dataset with {len(train_dataset)} items.\")\n    train_loader = DataLoader(train_dataset, batch_size=4, shuffle=True)\nexcept Exception as e:\n    print(f\"Error loading training data: {e}\")\n    train_loader = None","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T09:00:20.999606Z","iopub.execute_input":"2024-12-22T09:00:20.999951Z","iopub.status.idle":"2024-12-22T09:00:21.016032Z","shell.execute_reply.started":"2024-12-22T09:00:20.999921Z","shell.execute_reply":"2024-12-22T09:00:21.015034Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 13: Initialize model, criterion, and optimizer\nmodel = UNet()\ncriterion = nn.MSELoss()\noptimizer = optim.Adam(model.parameters(), lr=0.001)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T09:00:21.017174Z","iopub.execute_input":"2024-12-22T09:00:21.017627Z","iopub.status.idle":"2024-12-22T09:00:21.036987Z","shell.execute_reply.started":"2024-12-22T09:00:21.017533Z","shell.execute_reply":"2024-12-22T09:00:21.036011Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 14: Train the model\nif train_loader is not None:\n    try:\n        train_model(model, train_loader, criterion, optimizer, num_epochs=10)\n    except Exception as e:\n        print(f\"Error during training: {e}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T09:00:21.038341Z","iopub.execute_input":"2024-12-22T09:00:21.038754Z","iopub.status.idle":"2024-12-22T09:00:21.050438Z","shell.execute_reply.started":"2024-12-22T09:00:21.038709Z","shell.execute_reply":"2024-12-22T09:00:21.049286Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 15: Construct the test data path\ntest_data_path = os.path.join(test_path, \"static\", \"ExperimentRuns\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T09:00:21.051600Z","iopub.execute_input":"2024-12-22T09:00:21.051969Z","iopub.status.idle":"2024-12-22T09:00:21.065767Z","shell.execute_reply.started":"2024-12-22T09:00:21.051941Z","shell.execute_reply":"2024-12-22T09:00:21.064619Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 16: Load test dataset and make predictions\ntry:\n    test_dataset = CryoETDataset(test_data_path)\n    print(f\"Loaded test dataset with {len(test_dataset)} items.\")\n    test_loader = DataLoader(test_dataset, batch_size=4, shuffle=False)\nexcept Exception as e:\n    print(f\"Error loading test data: {e}\")\n    test_loader = None  # Ensure test_loader is defined","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T09:00:21.066879Z","iopub.execute_input":"2024-12-22T09:00:21.067187Z","iopub.status.idle":"2024-12-22T09:00:21.081645Z","shell.execute_reply.started":"2024-12-22T09:00:21.067159Z","shell.execute_reply":"2024-12-22T09:00:21.080671Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 17: Make predictions on the test set\ndef make_predictions(model, dataloader):\n    model.eval()\n    predictions = []\n    with torch.no_grad():\n        for inputs in dataloader:\n            outputs = model(inputs.unsqueeze(1))\n            predictions.append(outputs)\n    return torch.cat(predictions)\n\nif test_loader is not None:\n    try:\n        predictions = make_predictions(model, test_loader)\n    except Exception as e:\n        print(f\"Error during predictions: {e}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T09:00:21.082644Z","iopub.execute_input":"2024-12-22T09:00:21.082938Z","iopub.status.idle":"2024-12-22T09:00:21.091831Z","shell.execute_reply.started":"2024-12-22T09:00:21.082911Z","shell.execute_reply":"2024-12-22T09:00:21.090987Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 18: Evaluate the model\nif train_loader is not None:\n    evaluate_model(model, train_loader, criterion)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T09:00:21.092919Z","iopub.execute_input":"2024-12-22T09:00:21.093300Z","iopub.status.idle":"2024-12-22T09:00:21.108601Z","shell.execute_reply.started":"2024-12-22T09:00:21.093273Z","shell.execute_reply":"2024-12-22T09:00:21.107693Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 19: Save the Model\nsave_model(model, 'trained_model.pth')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T09:00:21.109683Z","iopub.execute_input":"2024-12-22T09:00:21.110022Z","iopub.status.idle":"2024-12-22T09:00:21.126607Z","shell.execute_reply.started":"2024-12-22T09:00:21.109996Z","shell.execute_reply":"2024-12-22T09:00:21.125461Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 20: Prepare Submission File\nif 'predictions' in locals():\n    submission_data = predictions.numpy()\n    np.save(submission_path.replace('.csv', '.npy'), submission_data)\n    print(\"Model training and predictions completed. Model saved and predictions saved as .npy file.\")\nelse:\n    print(\"Predictions not available for submission.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T09:00:21.127725Z","iopub.execute_input":"2024-12-22T09:00:21.128072Z","iopub.status.idle":"2024-12-22T09:00:21.141908Z","shell.execute_reply.started":"2024-12-22T09:00:21.128036Z","shell.execute_reply":"2024-12-22T09:00:21.141009Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 21: Visualize Predictions\nif 'predictions' in locals():\n    visualize_predictions(predictions)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T09:00:21.142850Z","iopub.execute_input":"2024-12-22T09:00:21.143147Z","iopub.status.idle":"2024-12-22T09:00:21.161061Z","shell.execute_reply.started":"2024-12-22T09:00:21.143123Z","shell.execute_reply":"2024-12-22T09:00:21.159991Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 22: Example usage of the F-beta score\ntry:\n    ground_truth = load_data(\"path_to_ground_truth\")  # Load ground truth data\n    score = f_beta_score(ground_truth, predictions.cpu().numpy())  # Ensure predictions are on CPU for numpy conversion\n    print(f'F-beta Score: {score}')\nexcept Exception as e:\n    print(f\"Error calculating F-beta score: {e}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T09:00:21.161969Z","iopub.execute_input":"2024-12-22T09:00:21.162291Z","iopub.status.idle":"2024-12-22T09:00:21.176132Z","shell.execute_reply.started":"2024-12-22T09:00:21.162265Z","shell.execute_reply":"2024-12-22T09:00:21.175095Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 23: Save the model if needed\nsave_model(model, 'trained_model.pkl')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T09:00:21.179180Z","iopub.execute_input":"2024-12-22T09:00:21.179472Z","iopub.status.idle":"2024-12-22T09:00:21.194323Z","shell.execute_reply.started":"2024-12-22T09:00:21.179444Z","shell.execute_reply":"2024-12-22T09:00:21.193389Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 24: Load the model from .pth file\nloaded_model = UNet()  # Replace with your model class\nloaded_model.load_state_dict(torch.load('trained_model.pth'))\nloaded_model.eval()  # Set model to evaluation mode","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T09:00:21.195599Z","iopub.execute_input":"2024-12-22T09:00:21.195946Z","iopub.status.idle":"2024-12-22T09:00:21.213866Z","shell.execute_reply.started":"2024-12-22T09:00:21.195919Z","shell.execute_reply":"2024-12-22T09:00:21.212748Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 25: Load a pickled object from .pkl file\ntry:\n    # If the object is a PyTorch model, use torch.load instead of pickle.load\n    loaded_model = UNet()  # Initialize the model class\n    loaded_model.load_state_dict(torch.load(model_path))  # Load the model state\n    loaded_model.eval()  # Set model to evaluation mode\n    print(\"Model loaded successfully.\")\nexcept FileNotFoundError as e:\n    print(f\"Error loading model file: {e}\")\nexcept Exception as e:\n    print(f\"Error loading model: {e}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T09:00:21.215111Z","iopub.execute_input":"2024-12-22T09:00:21.215458Z","iopub.status.idle":"2024-12-22T09:00:21.225559Z","shell.execute_reply.started":"2024-12-22T09:00:21.215422Z","shell.execute_reply":"2024-12-22T09:00:21.224541Z"}},"outputs":[],"execution_count":null}]}