{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.16","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"tpu1vmV38","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"}],"dockerImageVersionId":30920,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import seaborn as sns\n\nimport matplotlib.pyplot as plt\nimport os\nimport time\nimport numpy as np\nimport glob\nimport json\nimport collections\nimport torch\nimport torch.nn as nn\n\n\nimport matplotlib.patches as patches\n\nfrom matplotlib import animation, rc\nimport pandas as pd\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-29T06:40:22.582868Z","iopub.execute_input":"2025-03-29T06:40:22.584557Z","iopub.status.idle":"2025-03-29T06:40:22.589619Z","shell.execute_reply.started":"2025-03-29T06:40:22.584519Z","shell.execute_reply":"2025-03-29T06:40:22.588588Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install pydicom\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-29T06:40:22.590635Z","iopub.execute_input":"2025-03-29T06:40:22.591181Z","iopub.status.idle":"2025-03-29T06:40:25.95341Z","shell.execute_reply.started":"2025-03-29T06:40:22.591155Z","shell.execute_reply":"2025-03-29T06:40:25.951838Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nimport pydicom as dicom # dicom\nimport pydicom\nfrom pydicom.pixel_data_handlers.util import apply_voi_lut","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-29T06:40:25.954618Z","iopub.execute_input":"2025-03-29T06:40:25.954906Z","iopub.status.idle":"2025-03-29T06:40:25.959517Z","shell.execute_reply.started":"2025-03-29T06:40:25.954877Z","shell.execute_reply":"2025-03-29T06:40:25.958553Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# read data\ntrain_path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/'\n\ntrain  = pd.read_csv(train_path + 'train.csv')\nlabel = pd.read_csv(train_path + 'train_label_coordinates.csv')\ntrain_desc  = pd.read_csv(train_path + 'train_series_descriptions.csv')\ntest_desc   = pd.read_csv(train_path + 'test_series_descriptions.csv')\nsub         = pd.read_csv(train_path + 'sample_submission.csv')\nlen(test_desc) #number of test_description.csv rows ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-29T06:40:25.960543Z","iopub.execute_input":"2025-03-29T06:40:25.960771Z","iopub.status.idle":"2025-03-29T06:40:26.139362Z","shell.execute_reply.started":"2025-03-29T06:40:25.960748Z","shell.execute_reply":"2025-03-29T06:40:26.138054Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Function to generate image paths based on directory structure\ndef generate_image_paths(df, data_dir):\n    image_paths = []\n    for study_id, series_id in zip(df['study_id'], df['series_id']):\n        study_dir = os.path.join(data_dir, str(study_id))\n        series_dir = os.path.join(study_dir, str(series_id))\n        images = os.listdir(series_dir)\n        image_paths.extend([os.path.join(series_dir, img) for img in images])\n    return image_paths\n\n# Generate image paths for train and test data\ntrain_image_paths = generate_image_paths(train_desc, f'{train_path}/train_images')\ntest_image_paths = generate_image_paths(test_desc, f'{train_path}/test_images')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-29T06:40:26.140284Z","iopub.execute_input":"2025-03-29T06:40:26.140518Z","iopub.status.idle":"2025-03-29T06:40:32.706571Z","shell.execute_reply.started":"2025-03-29T06:40:26.140494Z","shell.execute_reply":"2025-03-29T06:40:32.705183Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define function to reshape a single row of the DataFrame\ndef reshape_row(row):\n    data = {'study_id': [], 'condition': [], 'level': [], 'severity': []}\n    \n    for column, value in row.items():\n        if column not in ['study_id', 'series_id', 'instance_number', 'x', 'y', 'series_description']:\n            parts = column.split('_')\n            condition = ' '.join([word.capitalize() for word in parts[:-2]])\n            level = parts[-2].capitalize() + '/' + parts[-1].capitalize()\n            data['study_id'].append(row['study_id'])\n            data['condition'].append(condition)\n            data['level'].append(level)\n            data['severity'].append(value)\n    \n    return pd.DataFrame(data)\n\n# Reshape the DataFrame for all rows\nnew_train_df = pd.concat([reshape_row(row) for _, row in train.iterrows()], ignore_index=True)\n\n# Display the first few rows of the reshaped dataframe\nnew_train_df.head(5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-29T06:40:32.707573Z","iopub.execute_input":"2025-03-29T06:40:32.707853Z","iopub.status.idle":"2025-03-29T06:40:33.789065Z","shell.execute_reply.started":"2025-03-29T06:40:32.707821Z","shell.execute_reply":"2025-03-29T06:40:33.788078Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Merge the dataframes on the common columns\nmerged_df = pd.merge(new_train_df, label, on=['study_id', 'condition', 'level'], how='inner')\n# Merge the dataframes on the common column 'series_id'\nfinal_merged_df = pd.merge(merged_df, train_desc, on='series_id', how='inner')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-29T06:40:33.790226Z","iopub.execute_input":"2025-03-29T06:40:33.790455Z","iopub.status.idle":"2025-03-29T06:40:33.836648Z","shell.execute_reply.started":"2025-03-29T06:40:33.790433Z","shell.execute_reply":"2025-03-29T06:40:33.835813Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Merge the dataframes on the common column 'series_id'\nfinal_merged_df = pd.merge(merged_df, train_desc, on=['series_id','study_id'], how='inner')\n# Display the first few rows of the final merged dataframe\nfinal_merged_df.head(5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-29T06:40:33.837846Z","iopub.execute_input":"2025-03-29T06:40:33.838212Z","iopub.status.idle":"2025-03-29T06:40:33.861192Z","shell.execute_reply.started":"2025-03-29T06:40:33.838186Z","shell.execute_reply":"2025-03-29T06:40:33.860501Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\n# Create the row_id column\nfinal_merged_df['row_id'] = (\n    final_merged_df['study_id'].astype(str) + '_' +\n    final_merged_df['condition'].str.lower().str.replace(' ', '_') + '_' +\n    final_merged_df['level'].str.lower().str.replace('/', '_')\n)\n\n# Create the image_path column\nfinal_merged_df['image_path'] = (\n    f'{train_path}/train_images/' + \n    final_merged_df['study_id'].astype(str) + '/' +\n    final_merged_df['series_id'].astype(str) + '/' +\n    final_merged_df['instance_number'].astype(str) + '.dcm'\n)\n\n# Note: Check image path, since there's 1 instance id, for 1 image, but there's many more images other than the ones labelled in the instance ID. \n\n# Display the updated dataframe\nfinal_merged_df.head(5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-29T06:40:33.862562Z","iopub.execute_input":"2025-03-29T06:40:33.862815Z","iopub.status.idle":"2025-03-29T06:40:34.021328Z","shell.execute_reply.started":"2025-03-29T06:40:33.862774Z","shell.execute_reply":"2025-03-29T06:40:34.020086Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define the base path for test images\nbase_path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_images/'\n\n# Function to get image paths for a series\ndef get_image_paths(row):\n    series_path = os.path.join(base_path, str(row['study_id']), str(row['series_id']))\n    if os.path.exists(series_path):\n        return [os.path.join(series_path, f) for f in os.listdir(series_path) if os.path.isfile(os.path.join(series_path, f))]\n    return []\n\n# Mapping of series_description to conditions\ncondition_mapping = {\n    'Sagittal T1': {'left': 'left_neural_foraminal_narrowing', 'right': 'right_neural_foraminal_narrowing'},\n    'Axial T2': {'left': 'left_subarticular_stenosis', 'right': 'right_subarticular_stenosis'},\n    'Sagittal T2/STIR': 'spinal_canal_stenosis'\n}\n\n# Create a list to store the expanded rows\nexpanded_rows = []\n\n# Expand the dataframe by adding new rows for each file path\nfor index, row in test_desc.iterrows():\n    image_paths = get_image_paths(row)\n    conditions = condition_mapping.get(row['series_description'], {})\n    if isinstance(conditions, str):  # Single condition\n        conditions = {'left': conditions, 'right': conditions}\n    for side, condition in conditions.items():\n        for image_path in image_paths:\n            expanded_rows.append({\n                'study_id': row['study_id'],\n                'series_id': row['series_id'],\n                'series_description': row['series_description'],\n                'image_path': image_path,\n                'condition': condition,\n                'row_id': f\"{row['study_id']}_{condition}\"\n            })\n\n# Create a new dataframe from the expanded rows\nexpanded_test_desc = pd.DataFrame(expanded_rows)\n\n# Display the resulting dataframe\nexpanded_test_desc.head(5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-29T06:40:34.022356Z","iopub.execute_input":"2025-03-29T06:40:34.022594Z","iopub.status.idle":"2025-03-29T06:40:34.18286Z","shell.execute_reply.started":"2025-03-29T06:40:34.02257Z","shell.execute_reply":"2025-03-29T06:40:34.181422Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# change severity column labels\n#Normal/Mild': 'normal_mild', 'Moderate': 'moderate', 'Severe': 'severe'}\nfinal_merged_df['severity'] = final_merged_df['severity'].map({'Normal/Mild': 'normal_mild', 'Moderate': 'moderate', 'Severe': 'severe'})","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-29T06:40:34.183666Z","iopub.execute_input":"2025-03-29T06:40:34.183947Z","iopub.status.idle":"2025-03-29T06:40:34.193071Z","shell.execute_reply.started":"2025-03-29T06:40:34.183907Z","shell.execute_reply":"2025-03-29T06:40:34.191821Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_data = expanded_test_desc\ntrain_data = final_merged_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-29T06:40:34.194149Z","iopub.execute_input":"2025-03-29T06:40:34.194392Z","iopub.status.idle":"2025-03-29T06:40:34.209755Z","shell.execute_reply.started":"2025-03-29T06:40:34.194367Z","shell.execute_reply":"2025-03-29T06:40:34.208581Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_dicom(path):\n    dicom = pydicom.dcmread(path)\n    data = dicom.pixel_array\n    data = data - np.min(data)\n    if np.max(data) != 0:\n        data = data / np.max(data)\n    data = (data * 255).astype(np.uint8)\n    return data","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-29T06:40:34.210345Z","iopub.execute_input":"2025-03-29T06:40:34.210558Z","iopub.status.idle":"2025-03-29T06:40:34.231608Z","shell.execute_reply.started":"2025-03-29T06:40:34.210536Z","shell.execute_reply":"2025-03-29T06:40:34.230161Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import random\nimport matplotlib.pyplot as plt\n\n# Yeni sıfırlanmış indekslerle rastgele seçim yapalım\nfinal_merged_df_reset = final_merged_df.reset_index(drop=True)\n\n# Rastgele iki indeks seçelim\nselected_indices = random.sample(range(len(final_merged_df_reset)), 2)\n\nimages = []\nrow_ids = []\n\n# Seçilen indekslerle görselleri yükleyelim\nfor i in selected_indices:\n    image = load_dicom(final_merged_df_reset['image_path'][i])  # Yeni sıfırlanmış indeksi kullan\n    images.append(image)\n    row_ids.append(final_merged_df_reset['row_id'][i])  # Yeni sıfırlanmış indeksi kullan\n\n# Görselleri çizdirelim\nfig, ax = plt.subplots(1, 2, figsize=(8, 4))\nfor i in range(2):\n    ax[i].imshow(images[i], cmap='gray')\n    ax[i].set_title(f'Row ID: {row_ids[i]}', fontsize=8)\n    ax[i].axis('off')\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-29T06:40:34.233107Z","iopub.execute_input":"2025-03-29T06:40:34.233332Z","iopub.status.idle":"2025-03-29T06:40:34.73408Z","shell.execute_reply.started":"2025-03-29T06:40:34.233309Z","shell.execute_reply":"2025-03-29T06:40:34.73247Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_data = train_data.dropna()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-29T06:40:34.735066Z","iopub.execute_input":"2025-03-29T06:40:34.73533Z","iopub.status.idle":"2025-03-29T06:40:34.764327Z","shell.execute_reply.started":"2025-03-29T06:40:34.735298Z","shell.execute_reply":"2025-03-29T06:40:34.762967Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nfrom sklearn.model_selection import train_test_split\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as transforms\nimport torch\nimport torch.optim.lr_scheduler as lr_scheduler\nfrom tqdm import tqdm\n\n# Define a custom dataset class\nclass CustomDataset(Dataset):\n    def __init__(self, dataframe, transform=None):\n        self.dataframe = dataframe\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.dataframe)\n\n    def __getitem__(self, index):\n        image_path = self.dataframe['image_path'][index]\n        image = load_dicom(image_path)  # Define this function to load your DICOM images\n        label = self.dataframe['severity'][index]\n        \n        if self.transform:\n            image = self.transform(image)\n\n        return image, label\n\n\n# Function to create datasets and dataloaders for each series description\ndef create_datasets_and_loaders(df, series_description, transform, batch_size=8):\n    filtered_df = df[df['series_description'] == series_description]\n    \n    # %5'ini al frac değerini değiştirerek trainde verinin ne kadarını kullanacagını belirlersin\n    filtered_df = filtered_df.sample(frac=1.0, random_state=42)  \n    \n    train_df, val_df = train_test_split(filtered_df, test_size=0.2, random_state=42)\n    train_df = train_df.reset_index(drop=True)\n    val_df = val_df.reset_index(drop=True)\n\n    train_dataset = CustomDataset(train_df, transform)\n    val_dataset = CustomDataset(val_df, transform)\n\n    trainloader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)\n    valloader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False)\n    \n    return trainloader, valloader, len(train_df), len(val_df)\n\n\n# Define the transforms\ntransform = transforms.Compose([\n    transforms.Lambda(lambda x: (x * 255).astype(np.uint8)),  # Convert back to uint8 for PIL\n    transforms.ToPILImage(),\n    transforms.Resize((224, 224)),\n    transforms.Grayscale(num_output_channels=3),\n    transforms.ToTensor(),\n])\n\n# Create dataloaders for each series description\ndataloaders = {}\nlengths = {}\n\ntrainloader_t1, valloader_t1, len_train_t1, len_val_t1 = create_datasets_and_loaders(train_data, 'Sagittal T1', transform)\ntrainloader_t2, valloader_t2, len_train_t2, len_val_t2 = create_datasets_and_loaders(train_data, 'Axial T2', transform)\ntrainloader_t2stir, valloader_t2stir, len_train_t2stir, len_val_t2stir = create_datasets_and_loaders(train_data, 'Sagittal T2/STIR', transform)\n\ndataloaders['Sagittal T1'] = (trainloader_t1, valloader_t1)\ndataloaders['Axial T2'] = (trainloader_t2, valloader_t2)\ndataloaders['Sagittal T2/STIR'] = (trainloader_t2stir, valloader_t2stir)\n\nlengths['Sagittal T1'] = (len_train_t1, len_val_t1)\nlengths['Axial T2'] = (len_train_t2, len_val_t2)\nlengths['Sagittal T2/STIR'] = (len_train_t2stir, len_val_t2stir)\n\n# Dictionary mapping labels to indices\nlabel_map = {'Mild': 0, 'Moderate': 1, 'Severe': 2}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-29T06:40:34.765706Z","iopub.execute_input":"2025-03-29T06:40:34.765981Z","iopub.status.idle":"2025-03-29T06:40:40.967837Z","shell.execute_reply.started":"2025-03-29T06:40:34.765951Z","shell.execute_reply":"2025-03-29T06:40:40.965465Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\n# Function to visualize a batch of images\ndef visualize_batch(dataloader):\n    images, labels = next(iter(dataloader))\n    fig, axes = plt.subplots(1, len(images), figsize=(20, 5))\n    for i, (img, lbl) in enumerate(zip(images, labels)):\n        ax = axes[i]\n        img = img.permute(1, 2, 0)  # Convert to HWC for visualization\n        ax.imshow(img)\n        ax.set_title(f\"Label: {lbl}\")\n        ax.axis('off')\n    plt.show()\n\n# Visualize samples from each dataloader\nprint(\"Visualizing Sagittal T1 samples\")\nvisualize_batch(trainloader_t1)\nprint(\"Visualizing Axial T2 samples\")\nvisualize_batch(trainloader_t2)\nprint(\"Visualizing Sagittal T2/STIR samples\")\nvisualize_batch(trainloader_t2stir)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-29T06:40:40.968945Z","iopub.execute_input":"2025-03-29T06:40:40.969483Z","iopub.status.idle":"2025-03-29T06:40:42.81018Z","shell.execute_reply.started":"2025-03-29T06:40:40.96945Z","shell.execute_reply":"2025-03-29T06:40:42.80909Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torchvision.models as models\nfrom torchvision import transforms\nfrom torch.utils.data import DataLoader\nfrom sklearn.model_selection import train_test_split\nimport pandas as pd\nfrom tqdm import tqdm\nfrom torchvision.models import resnet50, ResNet50_Weights\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n\n\nclass CustomResNet50(nn.Module):\n    def __init__(self, num_classes=3, pretrained=True):\n        super(CustomResNet50, self).__init__()\n        weights = ResNet50_Weights.IMAGENET1K_V1 if pretrained else None\n        \n        self.model = resnet50(weights=weights)\n        self.model.fc = nn.Linear(self.model.fc.in_features, num_class\n    def forward(self, x):\n        return self.model(x)\n\n    def unfreeze_model(self):\n        \"\"\"Tüm katmanları çöz.\"\"\"\n        for param in self.model.parameters():\n            param.requires_grad = True\n\n    def unfreeze_specific_layers(self, layer_names=None):\n        \"\"\"\n        Belirli katmanları çözmek için kullanılabilir.\n        Eğer layer_names None ise, tüm katmanlar çözülür.\n        \"\"\"\n        for name, param in self.model.named_parameters():\n            if layer_names is None or any(layer in name for layer in layer_names):\n                param.requires_grad = True\n            else:\n                param.requires_grad = False\n\n# Cihaz seçimi\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# Modeli başlat\nsagittal_t1_model = CustomResNet50(num_classes=3).to(device)\naxial_t2_model = CustomResNet50(num_classes=3).to(device)\nsagittal_t2stir_model = CustomResNet50(num_classes=3).to(device)\n\n\"\"\"# Son fully connected katmanı çözme\nfor param in sagittal_t1_model.model.fc.parameters():\n    param.requires_grad = True\nfor param in axial_t2_model.model.fc.parameters():\n    param.requires_grad = True\nfor param in sagittal_t2stir_model.model.fc.parameters():\n    param.requires_grad = True\n\n# Başlangıç katmanlarını dondurma\nfor param in sagittal_t1_model.model.parameters():\n    if param is not sagittal_t1_model.model.fc.weight and param is not sagittal_t1_model.model.fc.bias:\n        param.requires_grad = False\n\nfor param in axial_t2_model.model.parameters():\n    if param is not axial_t2_model.model.fc.weight and param is not axial_t2_model.model.fc.bias:\n        param.requires_grad = False\n\nfor param in sagittal_t2stir_model.model.parameters():\n    if param is not sagittal_t2stir_model.model.fc.weight and param is not sagittal_t2stir_model.model.fc.bias:\n        param.requires_grad = False\n\n# Eğitim parametreleri\ncriterion = nn.CrossEntropyLoss()\"\"\"\n# Tüm katmanları çözmek için\nfor model in [sagittal_t1_model, axial_t2_model, sagittal_t2stir_model]:\n    model.unfreeze_model()  # Bütün katmanları çöz\n\n# Eğitim parametreleri 05.12.2024 saat 0423'de güncellendi.\nweights = torch.tensor([1.0, 2.0, 4.0])\ncriterion = nn.CrossEntropyLoss(weight=weights.to(device))\n\n\n# Using SGD optimizer with momentum for all models\noptimizer_sagittal_t1 = torch.optim.SGD(sagittal_t1_model.model.fc.parameters(), lr=0.001, momentum=0.9)\noptimizer_axial_t2 = torch.optim.SGD(axial_t2_model.model.fc.parameters(), lr=0.001, momentum=0.9)\noptimizer_sagittal_t2stir = torch.optim.SGD(sagittal_t2stir_model.model.fc.parameters(), lr=0.001, momentum=0.9)\n\n# Modelleri ve optimizörleri saklamak için dictionary\nmodels = {\n    'Sagittal T1': sagittal_t1_model,\n    'Axial T2': axial_t2_model,\n    'Sagittal T2/STIR': sagittal_t2stir_model,\n}\n\noptimizers = {\n    'Sagittal T1': optimizer_sagittal_t1,\n    'Axial T2': optimizer_axial_t2,\n    'Sagittal T2/STIR': optimizer_sagittal_t2stir,\n}\n\n\nfor model_name, model in models.items():\n    trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    print(f\"Trainable parameters for {model_name}: {trainable_params}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-29T06:44:52.293683Z","iopub.execute_input":"2025-03-29T06:44:52.294107Z","iopub.status.idle":"2025-03-29T06:44:53.654544Z","shell.execute_reply.started":"2025-03-29T06:44:52.294067Z","shell.execute_reply":"2025-03-29T06:44:53.653056Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"label_map = {'normal_mild': 0, 'moderate': 1, 'severe': 2}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-29T06:44:57.820437Z","iopub.execute_input":"2025-03-29T06:44:57.820781Z","iopub.status.idle":"2025-03-29T06:44:57.824524Z","shell.execute_reply.started":"2025-03-29T06:44:57.820752Z","shell.execute_reply":"2025-03-29T06:44:57.823674Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for images, labels in trainloader_t2:\n    labels = torch.tensor([label_map[label] for label in labels])\n    labels = labels.to(device)\n    print(labels)\n    break","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-29T06:45:02.340871Z","iopub.execute_input":"2025-03-29T06:45:02.341265Z","iopub.status.idle":"2025-03-29T06:45:02.501949Z","shell.execute_reply.started":"2025-03-29T06:45:02.341235Z","shell.execute_reply":"2025-03-29T06:45:02.500912Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch.optim.lr_scheduler as lr_scheduler\nfrom copy import deepcopy\n\ndef train_model(model, trainloader, valloader, len_train, len_val, optimizer, num_epochs=20, patience=3):\n    # Learning rate scheduler\n    scheduler = lr_scheduler.StepLR(optimizer, step_size=2, gamma=0.1)\n    \n    best_val_acc = 0.0\n    best_model_wts = deepcopy(model.state_dict())\n    counter = 0\n    \n    for epoch in range(num_epochs):\n        model.train()\n        train_loss = 0\n        correct_train = 0\n        \n        with tqdm(trainloader, unit=\"batch\") as tepoch:\n            for images, labels in tepoch:\n                images, labels = images.to(device), torch.tensor([label_map[label] for label in labels]).to(device)\n                optimizer.zero_grad()\n                outputs = model(images)\n                loss = criterion(outputs, labels)\n                loss.backward()\n                optimizer.step()\n                train_loss += loss.item()\n                \n                probabilities = torch.softmax(outputs, dim=1)\n                _, predicted = torch.max(probabilities, 1)\n                correct_train += (predicted == labels).sum().item()\n                \n                tepoch.set_postfix(epoch=epoch+1)\n        \n        scheduler.step()\n        \n        train_loss /= len(trainloader)\n        train_acc = 100 * correct_train / len_train\n        \n        model.eval()\n        val_loss, correct_val = 0, 0\n        with torch.no_grad():\n            with tqdm(valloader, unit=\"batch\") as vepoch:\n                for images, labels in vepoch:\n                    images, labels = images.to(device), torch.tensor([label_map[label] for label in labels]).to(device)\n                    outputs = model(images)\n                    loss = criterion(outputs, labels)\n                    val_loss += loss.item()\n                    \n                    probabilities = torch.softmax(outputs, dim=1)\n\n                    # Eğer batch size 1 ise, dim=0 kullanarak doğru boyutta işlem yapabilirsiniz\n                    if probabilities.dim() == 1:\n                        _, predicted = torch.max(probabilities, 0)  # batch size 1 ise dim=0\n                    else:\n                        _, predicted = torch.max(probabilities, 1)  # normal durumda dim=1\n                    correct_val += (predicted == labels).sum().item()\n                    \n                    vepoch.set_postfix(epoch=epoch+1)\n        \n        val_loss /= len(valloader)\n        val_acc = 100 * correct_val / len_val\n        \n        print(f\"Epoch {epoch+1}, Train Loss: {train_loss:.4f}, Train Acc: {train_acc:.2f}%, Val Loss: {val_loss:.4f}, Val Acc: {val_acc:.2f}%\")\n        \n        # Save the best model and check for early stopping\n        if val_acc > best_val_acc:\n            best_val_acc = val_acc\n            best_model_wts = deepcopy(model.state_dict())\n            counter = 0\n            torch.save(best_model_wts, f'best_model_{epoch+1}.pth')\n        else:\n            counter += 1\n        \n        # # Early stopping\n        # if counter >= patience:\n        #     print(f\"Early stopping triggered after {epoch+1} epochs\")\n        #     break\n    \n    # Load best model weights\n    model.load_state_dict(best_model_wts)\n    return model, best_val_acc","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-29T06:45:10.069334Z","iopub.execute_input":"2025-03-29T06:45:10.069731Z","iopub.status.idle":"2025-03-29T06:45:10.081623Z","shell.execute_reply.started":"2025-03-29T06:45:10.069688Z","shell.execute_reply":"2025-03-29T06:45:10.080763Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Training all models\nfor desc, model in models.items():\n    # if desc == 'Sagittal T1':\n    #     trainloader, valloader, len_train, len_val = trainloader_t1, valloader_t1, len_train_t1, len_val_t1\n    # if desc == 'Axial T2':\n    #     trainloader, valloader, len_train, len_val = trainloader_t2, valloader_t2, len_train_t2, len_val_t2\n    if desc == 'Sagittal T2/STIR':\n        trainloader, valloader, len_train, len_val = trainloader_t2stir, valloader_t2stir, len_train_t2stir, len_val_t2stir\n    \n        print(f\"Training model for {desc}\")\n        train_model(model, trainloader, valloader, len_train, len_val, optimizers[desc])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-29T06:45:19.069799Z","iopub.execute_input":"2025-03-29T06:45:19.070214Z"}},"outputs":[],"execution_count":null}]}