{"metadata":{"accelerator":"GPU","colab":{"gpuType":"T4","provenance":[]},"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":9577817,"sourceType":"datasetVersion","datasetId":5839157},{"sourceId":9577867,"sourceType":"datasetVersion","datasetId":5839195}],"dockerImageVersionId":30786,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## Load Library","metadata":{}},{"cell_type":"code","source":"#%pip install torch\n#%pip install torchvision\n#%pip install transformers\n\nimport csv\nfrom datasets import Dataset, load_from_disk\nfrom datetime import datetime\nfrom PIL import Image as PILImage\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport os\nimport pandas as pd\nimport pydicom\nimport random\nimport shutil\nimport torch\nfrom torchvision import transforms\nfrom transformers import (\n    AutoImageProcessor,\n    ResNetForImageClassification,\n)\nimport warnings\nwarnings.filterwarnings(\"ignore\")","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:27:52.585041Z","iopub.execute_input":"2024-10-09T20:27:52.585447Z","iopub.status.idle":"2024-10-09T20:28:13.167885Z","shell.execute_reply.started":"2024-10-09T20:27:52.585403Z","shell.execute_reply":"2024-10-09T20:28:13.166970Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"current_date_yymmdd = datetime.now().strftime(\"%y%m%d%H%M\")","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:13.170067Z","iopub.execute_input":"2024-10-09T20:28:13.170898Z","iopub.status.idle":"2024-10-09T20:28:13.176179Z","shell.execute_reply.started":"2024-10-09T20:28:13.170848Z","shell.execute_reply":"2024-10-09T20:28:13.175066Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Directory where the test images are stored\nbase_dir = \"../input/rsna-2024-lumbar-spine-test-dataset/test_images\"\n\n# Output CSV file path\noutput_csv = \"test_data_file.csv\"","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:13.177874Z","iopub.execute_input":"2024-10-09T20:28:13.178265Z","iopub.status.idle":"2024-10-09T20:28:13.231578Z","shell.execute_reply.started":"2024-10-09T20:28:13.178219Z","shell.execute_reply":"2024-10-09T20:28:13.230449Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Open the CSV file for writing\nwith open(output_csv, mode='w', newline='') as file:\n    writer = csv.writer(file)\n    \n    # Write the header\n    writer.writerow(['study_id', 'series_id', 'instance_number', 'image_path'])\n\n    # Walk through the directory\n    for study_id in os.listdir(base_dir):\n        study_dir = os.path.join(base_dir, study_id)\n        if os.path.isdir(study_dir):\n            for series_id in os.listdir(study_dir):\n                series_dir = os.path.join(study_dir, series_id)\n                if os.path.isdir(series_dir):\n                    for instance_file in os.listdir(series_dir):\n                        if instance_file.endswith('.dcm'):\n                            # Extract instance number from the file name\n                            instance_number = instance_file.split('.')[0]\n                            # Build the full path to the file with forward slashes\n                            file_path = os.path.join(series_dir, instance_file).replace(os.sep, \"/\")\n                            # Write the row to the CSV\n                            writer.writerow([study_id, series_id, instance_number, file_path])\n\nprint(f\"CSV file created: {output_csv}\")","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:13.234253Z","iopub.execute_input":"2024-10-09T20:28:13.234792Z","iopub.status.idle":"2024-10-09T20:28:13.306873Z","shell.execute_reply.started":"2024-10-09T20:28:13.234739Z","shell.execute_reply":"2024-10-09T20:28:13.305705Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_test_data = pd.read_csv(f\"/kaggle/working/test_data_file.csv\")\ndf_test_data.head()","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:13.308121Z","iopub.execute_input":"2024-10-09T20:28:13.308477Z","iopub.status.idle":"2024-10-09T20:28:13.335116Z","shell.execute_reply.started":"2024-10-09T20:28:13.308440Z","shell.execute_reply":"2024-10-09T20:28:13.334115Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Convert dicom image to numpy array and create the dataset","metadata":{}},{"cell_type":"markdown","source":"The data originally stored in `dicom` format that cannot direcly accessed in Python. Here, we utilize `pydicom` library that made to convert the `dicom` data to `numpy` array.","metadata":{}},{"cell_type":"code","source":"def read_dicom_image(file_path, target_shape=(128, 128)):\n    dicom = pydicom.dcmread(file_path)\n    # Convert the DICOM pixel data to a NumPy array\n    image = dicom.pixel_array\n    # Normalize pixel values (if necessary)\n    image = (image / np.max(image) * 255).astype(np.uint8)\n    # Convert NumPy array to PIL Image\n    pil_image = PILImage.fromarray(image)\n    # Resize image to the target shape\n    resized_image = pil_image.resize(target_shape)\n    return resized_image","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:13.336477Z","iopub.execute_input":"2024-10-09T20:28:13.336806Z","iopub.status.idle":"2024-10-09T20:28:13.345408Z","shell.execute_reply.started":"2024-10-09T20:28:13.336771Z","shell.execute_reply":"2024-10-09T20:28:13.344471Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Huggingface has a `datasets` library that enables the user to create a custom dataset stored in `json` format easily.","metadata":{}},{"cell_type":"code","source":"image_path = df_test_data['image_path'].values\nstudy_id = df_test_data['study_id'].values\nseries_id = df_test_data['series_id'].values\ninstance_number = df_test_data['instance_number'].values","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:13.346518Z","iopub.execute_input":"2024-10-09T20:28:13.346837Z","iopub.status.idle":"2024-10-09T20:28:13.355353Z","shell.execute_reply.started":"2024-10-09T20:28:13.346804Z","shell.execute_reply":"2024-10-09T20:28:13.354476Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dataset = Dataset.from_dict({\n    \"image_path\": image_path,\n    \"study_id\": study_id,\n    \"series_id\": series_id,\n    \"instance_number\": instance_number\n})","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:13.356638Z","iopub.execute_input":"2024-10-09T20:28:13.357520Z","iopub.status.idle":"2024-10-09T20:28:13.395858Z","shell.execute_reply.started":"2024-10-09T20:28:13.357471Z","shell.execute_reply":"2024-10-09T20:28:13.394974Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dataset[0]","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:13.397058Z","iopub.execute_input":"2024-10-09T20:28:13.397396Z","iopub.status.idle":"2024-10-09T20:28:13.405275Z","shell.execute_reply.started":"2024-10-09T20:28:13.397361Z","shell.execute_reply":"2024-10-09T20:28:13.404361Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Dicom to numpy conversion begin**\n\n`target_shape` argument used for adjusting the numpy array resolution","metadata":{}},{"cell_type":"code","source":"def converts(example):\n    example['image'] = read_dicom_image(example['image_path'], target_shape=(256, 256))\n    return example","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:13.409037Z","iopub.execute_input":"2024-10-09T20:28:13.409467Z","iopub.status.idle":"2024-10-09T20:28:13.414253Z","shell.execute_reply.started":"2024-10-09T20:28:13.409430Z","shell.execute_reply":"2024-10-09T20:28:13.413349Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dataset = test_dataset.map(converts)","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:13.415466Z","iopub.execute_input":"2024-10-09T20:28:13.415763Z","iopub.status.idle":"2024-10-09T20:28:16.677808Z","shell.execute_reply.started":"2024-10-09T20:28:13.415723Z","shell.execute_reply":"2024-10-09T20:28:16.676703Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dataset","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:16.679235Z","iopub.execute_input":"2024-10-09T20:28:16.679695Z","iopub.status.idle":"2024-10-09T20:28:16.687916Z","shell.execute_reply.started":"2024-10-09T20:28:16.679647Z","shell.execute_reply":"2024-10-09T20:28:16.686879Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Displaying random image from the dataset","metadata":{}},{"cell_type":"code","source":"def display_images(dataset, num_rows=2, num_columns=5, figsize=(12, 10), max_title_length=30):\n    total_images = num_rows * num_columns\n\n    # Shuffle the dataset to get a random selection of images\n    indices = list(range(len(dataset)))\n    random.shuffle(indices)\n\n    fig, axes = plt.subplots(num_rows, num_columns, figsize=figsize)\n\n    for i, idx in enumerate(indices):\n        if i >= total_images:\n            break\n        example = dataset[idx]\n\n        row = i // num_columns\n        col = i % num_columns\n\n        image = example[\"image\"]\n        series_id = example[\"series_id\"]\n\n        # Display image\n        axes[row, col].imshow(image)\n        axes[row, col].axis('off')\n\n        axes[row, col].set_title(series_id, wrap=True, fontsize='small')\n    # Adjust spacing and layout\n    plt.tight_layout()\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:16.689491Z","iopub.execute_input":"2024-10-09T20:28:16.690040Z","iopub.status.idle":"2024-10-09T20:28:16.699469Z","shell.execute_reply.started":"2024-10-09T20:28:16.689991Z","shell.execute_reply":"2024-10-09T20:28:16.698372Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"display_images(test_dataset)","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:16.700796Z","iopub.execute_input":"2024-10-09T20:28:16.701235Z","iopub.status.idle":"2024-10-09T20:28:17.937609Z","shell.execute_reply.started":"2024-10-09T20:28:16.701189Z","shell.execute_reply":"2024-10-09T20:28:17.936461Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Save the dataset to local directory","metadata":{"id":"89d4d971-946e-4210-a676-5d48d08d970e"}},{"cell_type":"code","source":"# change to your desired target path\ntarget_path = \"test_dataset_256_condition_level_encoded_tmp_\"+current_date_yymmdd+\".hf\"\n\n# Check if the target path exists and is a directory, then remove it\nif os.path.exists(target_path):\n    if os.path.isdir(target_path):\n        shutil.rmtree(target_path)  # Removes the entire directory\n        print(f\"Existing directory at {target_path} has been deleted.\")\n    else:\n        os.remove(target_path)  # If it's a file, remove it\n        print(f\"Existing file at {target_path} has been deleted.\")\n\n# Save the dataset to the disk\ntest_dataset.save_to_disk(target_path)\nprint(f\"New dataset saved to {target_path}.\")","metadata":{"id":"d5b9e620-fc35-45f1-8ebb-3e481d65e626","outputId":"d911f3ab-5625-468f-82a2-dc54fc4b88d3","execution":{"iopub.status.busy":"2024-10-09T20:28:17.938961Z","iopub.execute_input":"2024-10-09T20:28:17.939429Z","iopub.status.idle":"2024-10-09T20:28:17.990904Z","shell.execute_reply.started":"2024-10-09T20:28:17.939391Z","shell.execute_reply":"2024-10-09T20:28:17.989892Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# change dataset_path to the path of your own .hf dataset\ndataset_path = target_path\ntest_dataset_256_load = load_from_disk(dataset_path)","metadata":{"id":"bc9b9cab-d886-4e07-84c5-b03698558d03","execution":{"iopub.status.busy":"2024-10-09T20:28:17.992441Z","iopub.execute_input":"2024-10-09T20:28:17.993156Z","iopub.status.idle":"2024-10-09T20:28:18.004005Z","shell.execute_reply.started":"2024-10-09T20:28:17.993104Z","shell.execute_reply":"2024-10-09T20:28:18.002901Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dataset_256_load[0]","metadata":{"id":"8cf892d8-1e7a-4de9-988d-c37e7b736883","outputId":"08cc7fca-fc39-4358-9fe6-00a0bd4b76ff","execution":{"iopub.status.busy":"2024-10-09T20:28:18.005470Z","iopub.execute_input":"2024-10-09T20:28:18.005895Z","iopub.status.idle":"2024-10-09T20:28:18.015693Z","shell.execute_reply.started":"2024-10-09T20:28:18.005847Z","shell.execute_reply":"2024-10-09T20:28:18.014550Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dataset_256_load[0]['image']","metadata":{"id":"f88ea50a-29b2-4561-8c54-e3cd3d804b59","outputId":"dfdbc972-5b53-40b8-a041-97d54c13bd74","execution":{"iopub.status.busy":"2024-10-09T20:28:18.016991Z","iopub.execute_input":"2024-10-09T20:28:18.017291Z","iopub.status.idle":"2024-10-09T20:28:18.035098Z","shell.execute_reply.started":"2024-10-09T20:28:18.017258Z","shell.execute_reply":"2024-10-09T20:28:18.034149Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Prediction","metadata":{}},{"cell_type":"code","source":"current_model = \"resnet\"\n\nmodel_resnet_path = \"../input/models/resnet-adw8bit-01-multi-targets\"\n\nmodel_dict = {\n    \"resnet\": {\n        \"model\": ResNetForImageClassification.from_pretrained(model_resnet_path, num_labels=75),\n        \"processor\": AutoImageProcessor.from_pretrained(model_resnet_path),\n    },\n}\n\nprocessor = model_dict[current_model][\"processor\"]\nmodel = model_dict[current_model][\"model\"]\nif current_model != \"resnet\":\n    size = processor.size[\"height\"]\nelse:\n    size = processor.size[\"shortest_edge\"]\n\nimage_mean, image_std = processor.image_mean, processor.image_std\nnormalize = transforms.Normalize(mean=image_mean, std=image_std)\n\nimage_transformer = transforms.Compose(\n    [transforms.Resize((size, size)), transforms.ToTensor(), normalize]\n)\n\ndef preprocess_function(examples):\n    examples['pixel_values'] = [image_transformer(image.convert(\"RGB\")) for image in examples['image']]\n    return examples\n\ntest_dataset_256_load.set_transform(preprocess_function)    \n\n#pixel_value = image_transformer(image.convert(\"RGB\"))\n#to_pil = transforms.ToPILImage()\n#pil_image = to_pil(pixel_value)\n\n#inputs = processor(pil_image, return_tensors=\"pt\").to()","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:18.036257Z","iopub.execute_input":"2024-10-09T20:28:18.036603Z","iopub.status.idle":"2024-10-09T20:28:18.803874Z","shell.execute_reply.started":"2024-10-09T20:28:18.036567Z","shell.execute_reply":"2024-10-09T20:28:18.803012Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dataset_256_load","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:18.805015Z","iopub.execute_input":"2024-10-09T20:28:18.805328Z","iopub.status.idle":"2024-10-09T20:28:18.811569Z","shell.execute_reply.started":"2024-10-09T20:28:18.805275Z","shell.execute_reply":"2024-10-09T20:28:18.810535Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Set the model to evaluation mode\nmodel.eval()","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:18.812872Z","iopub.execute_input":"2024-10-09T20:28:18.813259Z","iopub.status.idle":"2024-10-09T20:28:18.827337Z","shell.execute_reply.started":"2024-10-09T20:28:18.813213Z","shell.execute_reply":"2024-10-09T20:28:18.826367Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import Trainer, TrainingArguments","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:18.828708Z","iopub.execute_input":"2024-10-09T20:28:18.829397Z","iopub.status.idle":"2024-10-09T20:28:20.575834Z","shell.execute_reply.started":"2024-10-09T20:28:18.829349Z","shell.execute_reply":"2024-10-09T20:28:20.574863Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"target_columns = [\"spinal_canal_stenosis_l1_l2\",\n                  \"spinal_canal_stenosis_l2_l3\",\n                    \"spinal_canal_stenosis_l3_l4\",\n                    \"spinal_canal_stenosis_l4_l5\",\n                    \"spinal_canal_stenosis_l5_s1\",\n                    \"left_neural_foraminal_narrowing_l1_l2\",\n                    \"left_neural_foraminal_narrowing_l2_l3\",\n                    \"left_neural_foraminal_narrowing_l3_l4\",\n                    \"left_neural_foraminal_narrowing_l4_l5\",\n                    \"left_neural_foraminal_narrowing_l5_s1\",\n                    \"right_neural_foraminal_narrowing_l1_l2\",\n                    \"right_neural_foraminal_narrowing_l2_l3\",\n                    \"right_neural_foraminal_narrowing_l3_l4\",\n                    \"right_neural_foraminal_narrowing_l4_l5\",\n                    \"right_neural_foraminal_narrowing_l5_s1\",\n                    \"left_subarticular_stenosis_l1_l2\",\n                    \"left_subarticular_stenosis_l2_l3\",\n                    \"left_subarticular_stenosis_l3_l4\",\n                    \"left_subarticular_stenosis_l4_l5\",\n                    \"left_subarticular_stenosis_l5_s1\",\n                    \"right_subarticular_stenosis_l1_l2\",\n                    \"right_subarticular_stenosis_l2_l3\",\n                    \"right_subarticular_stenosis_l3_l4\",\n                    \"right_subarticular_stenosis_l4_l5\",\n                    \"right_subarticular_stenosis_l5_s1\"]","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:20.577328Z","iopub.execute_input":"2024-10-09T20:28:20.578288Z","iopub.status.idle":"2024-10-09T20:28:20.584691Z","shell.execute_reply.started":"2024-10-09T20:28:20.578237Z","shell.execute_reply":"2024-10-09T20:28:20.583558Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def collate_fn(examples):\n    pixel_values = torch.stack([example[\"pixel_values\"] for example in examples])\n    \n    # Handle 25 labels for multi-target learning\n    # Build the labels tensor while keeping the missing values (-1) intact\n    labels = torch.tensor([[example.get(target, -1) for target in target_columns] for example in examples])\n    \n    # labels should have shape (batch_size, 25) with 25 labels for each sample\n    return {\"pixel_values\": pixel_values, \"labels\": labels}","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:20.586127Z","iopub.execute_input":"2024-10-09T20:28:20.586746Z","iopub.status.idle":"2024-10-09T20:28:20.605618Z","shell.execute_reply.started":"2024-10-09T20:28:20.586697Z","shell.execute_reply":"2024-10-09T20:28:20.604642Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"training_args = TrainingArguments(\n    output_dir=\"./result_tmp/\" + current_date_yymmdd,\n    report_to=[],\n    #report_to=\"tensorboard\",\n    evaluation_strategy='epoch',\n    save_strategy='epoch',\n    learning_rate=1e-3,\n    per_device_train_batch_size=32,\n    per_device_eval_batch_size=32,\n    #gradient_accumulation_steps=4,   # Accumulate gradients over 4 steps allow to save memory\n    num_train_epochs=5,\n    optim='adamw_bnb_8bit',\n    weight_decay=0.001,\n    warmup_ratio=0.05,\n    #warmup_steps=50,\n    remove_unused_columns=False,\n    load_best_model_at_end=True,\n    save_total_limit=1,\n    logging_steps=1,\n    #fp16=True,  # Enable mixed precision training can greatly reduce memory usage by using half-precision floating-point numbers.\n    run_name=os.getenv('TEMP_RUN_NAME')\n)","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:20.606851Z","iopub.execute_input":"2024-10-09T20:28:20.607164Z","iopub.status.idle":"2024-10-09T20:28:20.619611Z","shell.execute_reply.started":"2024-10-09T20:28:20.607131Z","shell.execute_reply":"2024-10-09T20:28:20.618700Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nimport torch.nn.functional as F\n\ndef masked_loss(predictions, targets):\n    # predictions shape: (batch_size, 25, 3) - logit predictions for each target\n    # targets shape: (batch_size, 25) - multi-target labels\n    mask = targets != -1  # create a mask for non-missing labels\n    \n    # Apply the mask to the loss calculation\n    loss = F.cross_entropy(predictions[mask], targets[mask], reduction='mean')\n    return loss\n\nclass CustomTrainer(Trainer):\n    def compute_loss(self, model, inputs, return_outputs=False):\n        labels = inputs.pop(\"labels\")\n        outputs = model(**inputs)\n        logits = outputs.get(\"logits\")\n\n        # Reshape to (batch_size, 25, 3) and apply masked loss\n        logits = logits.view(-1, 25, 3)\n        loss = masked_loss(logits, labels)\n    \n        return (loss, outputs) if return_outputs else loss\n# Assuming you already have a Trainer instance\n\n\ntrainer = CustomTrainer(model=model, args=training_args, data_collator=collate_fn)","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:20.620873Z","iopub.execute_input":"2024-10-09T20:28:20.621365Z","iopub.status.idle":"2024-10-09T20:28:20.653124Z","shell.execute_reply.started":"2024-10-09T20:28:20.621301Z","shell.execute_reply":"2024-10-09T20:28:20.652072Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#print(trainer.args)","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:20.654519Z","iopub.execute_input":"2024-10-09T20:28:20.654932Z","iopub.status.idle":"2024-10-09T20:28:20.659582Z","shell.execute_reply.started":"2024-10-09T20:28:20.654883Z","shell.execute_reply":"2024-10-09T20:28:20.658574Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#test_dataset_256_load[0]","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:20.665414Z","iopub.execute_input":"2024-10-09T20:28:20.665834Z","iopub.status.idle":"2024-10-09T20:28:20.670124Z","shell.execute_reply.started":"2024-10-09T20:28:20.665798Z","shell.execute_reply":"2024-10-09T20:28:20.669154Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"outputs_simul = trainer.predict(test_dataset_256_load)","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:20.671673Z","iopub.execute_input":"2024-10-09T20:28:20.672072Z","iopub.status.idle":"2024-10-09T20:28:35.226059Z","shell.execute_reply.started":"2024-10-09T20:28:20.672026Z","shell.execute_reply":"2024-10-09T20:28:35.225024Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"y_pred_simul = outputs_simul.predictions.reshape(-1, 25, 3).argmax(2)\nprint(np.unique(y_pred_simul))","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:35.227260Z","iopub.execute_input":"2024-10-09T20:28:35.227618Z","iopub.status.idle":"2024-10-09T20:28:35.233813Z","shell.execute_reply.started":"2024-10-09T20:28:35.227582Z","shell.execute_reply":"2024-10-09T20:28:35.232806Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from collections import defaultdict\n\n#For submission test set : comment\n#probs = torch.softmax(torch.tensor(outputs.predictions), dim=-1)  # Get probabilities\n#For submission test set : uncomment\nprobs = torch.softmax(torch.tensor(outputs_simul.predictions), dim=-1)  # Get probabilities\n\n# Dictionary to accumulate probabilities and counts for averaging\naveraged_probs = defaultdict(lambda: [0.0, 0.0, 0.0])  # To store total sums of probabilities\ncount_dict = defaultdict(int)  # To store counts of occurrences\n\n# Iterate through the test set\nfor idx, example in enumerate(test_dataset_256_load):\n    study_id = example['study_id']\n\n    # Get predicted probabilities for this sample (shape: (batch_size, num_labels))\n    sample_probs = probs[idx]\n\n    # Assuming each target has 3 probabilities (Normal/Mild, Moderate, Severe)\n    for i, target in enumerate(target_columns):\n        if 2 == -1:\n        #if example[target] == -1:  #case to remove for dataset_test_simul or add the target columns as empty in dataset_test_simul\n            # Skip this target if it has a missing value (-1)\n            continue\n\n        # Calculate the correct slice indices for the 3 severity probabilities\n        start_idx = i * 3\n        end_idx = start_idx + 3\n\n        # Extract the logits for this particular target\n        target_logits = sample_probs[start_idx:end_idx]\n\n        # Apply softmax to normalize the probabilities\n        target_probs = F.softmax(torch.tensor(target_logits), dim=0)\n\n        # Update the accumulated probabilities\n        averaged_probs[f\"{study_id}_{target}\"][0] += float(target_probs[0])  # Normal/Mild\n        averaged_probs[f\"{study_id}_{target}\"][1] += float(target_probs[1])  # Moderate\n        averaged_probs[f\"{study_id}_{target}\"][2] += float(target_probs[2])  # Severe\n\n        # Increment the count for this row_id\n        count_dict[f\"{study_id}_{target}\"] += 1\n\n# Open the output CSV file for writing\nwith open('submission.csv', mode='w', newline='') as file:\n    writer = csv.writer(file)\n\n    # Write the header\n    writer.writerow(['row_id', 'normal_mild', 'moderate', 'severe'])\n\n    # Calculate averages and write to the CSV\n    for row_id, sum_probs in averaged_probs.items():\n        count = count_dict[row_id]  # Get the count for this row_id\n\n        # Only calculate the average if the count is greater than zero\n        if count > 0:\n            avg_normal_mild = sum_probs[0] / count\n            avg_moderate = sum_probs[1] / count\n            avg_severe = sum_probs[2] / count\n\n        # Write the row to the CSV\n        writer.writerow([row_id, avg_normal_mild, avg_moderate, avg_severe])\n\nprint(\"Submission file created as 'submission.csv' at \" + current_date_yymmdd + \".\")","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:35.235592Z","iopub.execute_input":"2024-10-09T20:28:35.236215Z","iopub.status.idle":"2024-10-09T20:28:35.640614Z","shell.execute_reply.started":"2024-10-09T20:28:35.236164Z","shell.execute_reply":"2024-10-09T20:28:35.639531Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Define the path to the submission file\nsubmission_file = '/kaggle/working/submission.csv'\n\n# Walk through the /kaggle/working directory\nfor item in os.listdir('/kaggle/working'):\n    item_path = os.path.join('/kaggle/working', item)\n    \n    # Skip the submission.csv file\n    if item_path != submission_file:\n        # Check if the item is a file or a directory\n        if os.path.isfile(item_path):\n            os.remove(item_path)  # Remove the file\n        elif os.path.isdir(item_path):\n            shutil.rmtree(item_path)  # Remove the directory\n\nprint(\"All files and folders (except submission.csv) have been removed.\")","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:35.641727Z","iopub.execute_input":"2024-10-09T20:28:35.642024Z","iopub.status.idle":"2024-10-09T20:28:35.649814Z","shell.execute_reply.started":"2024-10-09T20:28:35.641990Z","shell.execute_reply":"2024-10-09T20:28:35.648756Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pd.read_csv(submission_file)","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:35.651168Z","iopub.execute_input":"2024-10-09T20:28:35.651590Z","iopub.status.idle":"2024-10-09T20:28:35.673128Z","shell.execute_reply.started":"2024-10-09T20:28:35.651542Z","shell.execute_reply":"2024-10-09T20:28:35.672037Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_file","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:35.674463Z","iopub.execute_input":"2024-10-09T20:28:35.674847Z","iopub.status.idle":"2024-10-09T20:28:35.682052Z","shell.execute_reply.started":"2024-10-09T20:28:35.674799Z","shell.execute_reply":"2024-10-09T20:28:35.680816Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pred_values = pd.read_csv(\"/kaggle/working/submission.csv\")\n\n# Use pred_values[\"row_id\"] as the index and select your target columns\npreds_df = (\n    pd.DataFrame(\n        pred_values[[\"row_id\", \"normal_mild\", \"moderate\", \"severe\"]]\n    )\n    .set_index(\"row_id\")  # Set row_id as the index\n    .fillna(1 / 3)  # Handle missing values by filling with 1/3\n    .sort_values(\"row_id\", ascending=True)  # Sort by row_id\n    .reset_index()  # Reset the index so that row_id is a column again\n)\nprint(preds_df.head())\n","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:35.683717Z","iopub.execute_input":"2024-10-09T20:28:35.684468Z","iopub.status.idle":"2024-10-09T20:28:35.706584Z","shell.execute_reply.started":"2024-10-09T20:28:35.684420Z","shell.execute_reply":"2024-10-09T20:28:35.705576Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ss_df = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/sample_submission.csv')\nss_df.tail()","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:35.707709Z","iopub.execute_input":"2024-10-09T20:28:35.708008Z","iopub.status.idle":"2024-10-09T20:28:35.724374Z","shell.execute_reply.started":"2024-10-09T20:28:35.707976Z","shell.execute_reply":"2024-10-09T20:28:35.723257Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Find row_ids present in ss_df but not in preds_df\nmissing_in_preds_df = set(ss_df.row_id.to_list()) - set(preds_df.row_id.to_list())","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:35.725555Z","iopub.execute_input":"2024-10-09T20:28:35.725875Z","iopub.status.idle":"2024-10-09T20:28:35.731018Z","shell.execute_reply.started":"2024-10-09T20:28:35.725841Z","shell.execute_reply":"2024-10-09T20:28:35.729975Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Find row_ids present in preds_df but not in ss_df\nmissing_in_ss_df = set(preds_df.row_id.to_list()) - set(ss_df.row_id.to_list())","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:35.732514Z","iopub.execute_input":"2024-10-09T20:28:35.732898Z","iopub.status.idle":"2024-10-09T20:28:35.745877Z","shell.execute_reply.started":"2024-10-09T20:28:35.732856Z","shell.execute_reply":"2024-10-09T20:28:35.744984Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Display the differences\nprint(\"Missing in preds_df:\", missing_in_preds_df)\nprint(\"Missing in ss_df:\", missing_in_ss_df)","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:35.747520Z","iopub.execute_input":"2024-10-09T20:28:35.747914Z","iopub.status.idle":"2024-10-09T20:28:35.756249Z","shell.execute_reply.started":"2024-10-09T20:28:35.747869Z","shell.execute_reply":"2024-10-09T20:28:35.755294Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"set(ss_df.row_id.to_list())==set(preds_df.row_id.to_list()) ","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:35.757425Z","iopub.execute_input":"2024-10-09T20:28:35.757822Z","iopub.status.idle":"2024-10-09T20:28:35.767840Z","shell.execute_reply.started":"2024-10-09T20:28:35.757787Z","shell.execute_reply":"2024-10-09T20:28:35.766814Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Confirm that all row_ids are unique\nassert preds_df['row_id'].is_unique, \"row_id values are not unique!\"","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:35.768948Z","iopub.execute_input":"2024-10-09T20:28:35.769251Z","iopub.status.idle":"2024-10-09T20:28:35.782166Z","shell.execute_reply.started":"2024-10-09T20:28:35.769216Z","shell.execute_reply":"2024-10-09T20:28:35.781109Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#preds_df[['normal_mild', 'moderate', 'severe']] = preds_df[['normal_mild', 'moderate', 'severe']].round(6)\n\n# Round the float columns 'normal_mild', 'moderate', and 'severe' to 6 decimals\npreds_df['normal_mild'] = preds_df['normal_mild'].round(6)\npreds_df['moderate'] = preds_df['moderate'].round(6)\npreds_df['severe'] = preds_df['severe'].round(6)\n\n# Verify if the rounding was applied\nprint(preds_df.head())","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:35.783637Z","iopub.execute_input":"2024-10-09T20:28:35.783984Z","iopub.status.idle":"2024-10-09T20:28:35.797866Z","shell.execute_reply.started":"2024-10-09T20:28:35.783948Z","shell.execute_reply":"2024-10-09T20:28:35.796706Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission = pd.DataFrame(ss_df['row_id']).merge(preds_df,left_on='row_id',right_on='row_id',how='left', validate=\"1:1\").fillna(1./3)\nv = submission[['normal_mild','moderate','severe']].values\nv = v/v.sum(1).reshape(-1,1)\nsubmission[['normal_mild','moderate','severe']] = v\nsubmission","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:35.799477Z","iopub.execute_input":"2024-10-09T20:28:35.800042Z","iopub.status.idle":"2024-10-09T20:28:35.827446Z","shell.execute_reply.started":"2024-10-09T20:28:35.799990Z","shell.execute_reply":"2024-10-09T20:28:35.826375Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(submission.dtypes)","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:35.828740Z","iopub.execute_input":"2024-10-09T20:28:35.829216Z","iopub.status.idle":"2024-10-09T20:28:35.839148Z","shell.execute_reply.started":"2024-10-09T20:28:35.829168Z","shell.execute_reply":"2024-10-09T20:28:35.838124Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission.to_csv(\"/kaggle/working/submission.csv\", index=False, float_format=\"%.6f\")","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:35.840712Z","iopub.execute_input":"2024-10-09T20:28:35.841366Z","iopub.status.idle":"2024-10-09T20:28:35.853060Z","shell.execute_reply.started":"2024-10-09T20:28:35.841300Z","shell.execute_reply":"2024-10-09T20:28:35.852028Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Load the CSV file\n##df = pd.read_csv(\"/kaggle/working/submission.csv\").reset_index(drop=True)\n\n# Sort the dataframe by 'row_id'\n##df_sorted = df.sort_values(by='row_id')\n\n# Save the sorted dataframe back to the CSV\n##df_sorted.to_csv(\"/kaggle/working/submission.csv\", index=False)\n\n##print(\"CSV sorted by 'row_id' and saved.\")","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:35.854398Z","iopub.execute_input":"2024-10-09T20:28:35.854782Z","iopub.status.idle":"2024-10-09T20:28:35.859603Z","shell.execute_reply.started":"2024-10-09T20:28:35.854743Z","shell.execute_reply":"2024-10-09T20:28:35.858469Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from subprocess import check_output\nprint(check_output([\"ls\", \"/kaggle/working\"]).decode(\"utf8\"))","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:35.861025Z","iopub.execute_input":"2024-10-09T20:28:35.861351Z","iopub.status.idle":"2024-10-09T20:28:35.873202Z","shell.execute_reply.started":"2024-10-09T20:28:35.861297Z","shell.execute_reply":"2024-10-09T20:28:35.872157Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pd.read_csv(submission_file)","metadata":{"execution":{"iopub.status.busy":"2024-10-09T20:28:35.874475Z","iopub.execute_input":"2024-10-09T20:28:35.875116Z","iopub.status.idle":"2024-10-09T20:28:35.890492Z","shell.execute_reply.started":"2024-10-09T20:28:35.875078Z","shell.execute_reply":"2024-10-09T20:28:35.889539Z"},"trusted":true},"execution_count":null,"outputs":[]}]}