{"cells":[{"cell_type":"markdown","id":"cell-001","metadata":{},"source":"# A small reproducible MRI baseline for RSNA knee abnormalities\n\nA baseline is a simple reference model. This notebook uses magnetic resonance\nimaging (MRI) studies of the knee. One MRI study is one example for the model.\n\nA study can contain multiple image series. An acquisition is one MRI scan with\na specified setup. A series is a group of related images from one acquisition.\nEach image is a two-dimensional slice. The files use the Digital Imaging and\nCommunications in Medicine (DICOM) standard. ResNet-18 uses these slices to\nmake probabilities for 12 abnormalities.\n\nThis notebook explains each important part of the data path.\nPatient-level validation keeps each patient in only one data split. A fixed\nmethod selects the slices. The model averages valid features across slices and\nseries.\n\nInference is the use of a trained model to make predictions. The notebook\nreloads a saved checkpoint before test inference. A checkpoint is a file that\ncontains saved model weights. The last cell makes and validates a submission\nfile.\n\nTraining uses only labels supplied by the competition organizers. The model\ndoes not use radiology reports or external labels. A pseudo-label is a label\nthat another model makes. This model does not use pseudo-labels or leaderboard\nfeedback.\n\nThis reference run is small.\nIt trains on 54 patients and evaluates on 4 different patients. The run shows\nthat the pipeline can complete all steps.\nA score from four patients cannot give a stable estimate of model quality.\n\n"},{"cell_type":"markdown","id":"cell-002","metadata":{},"source":"## Load the tested code\n\nThe notebook uses short examples to explain the main steps. The `rsna_knee`\npackage does the complete run. A Python wheel is a file that contains a Python\npackage. An attached Kaggle Dataset contains this wheel.\n\nThe notebook does not need internet access. A hash is a value that identifies\nthe exact contents of a file. The reproducibility section shows the wheel hash\nand other source information.\n\n"},{"cell_type":"code","execution_count":null,"id":"cell-003","metadata":{},"outputs":[],"source":"import base64\nimport hashlib\nimport json\nimport subprocess\nimport sys\nimport time\nfrom pathlib import Path\n\nnotebook_started = time.perf_counter()\n\nWHEEL_DATASET_ID = \"plthon/rsna-knee-gold-baseline-package\"\nWHEEL_FILENAME = \"rsna_knee-0.1.0-py3-none-any.whl\"\nWHEEL_SHA256 = \"0529af6837f4029a30077c741fbbd390edec3e0752bfe12ac290135bf9a2efa7\"\nWEIGHT_DATASET_ID = \"plthon/torchvision-resnet18-imagenet1k-v1\"\nKERNEL_ID = \"plthon/rsna-knee-a-reproducible-mri-baseline\"\nGPU_APPROVED = json.loads(\"true\")\n\nwheel_roots = (\n    Path(\"/kaggle/input/datasets\") / WHEEL_DATASET_ID,\n    Path(\"/kaggle/input\") / \"rsna-knee-gold-baseline-package\",\n)\nwheel_candidates = {\n    path.resolve()\n    for root in wheel_roots\n    if root.exists()\n    for path in root.rglob(WHEEL_FILENAME)\n    if path.is_file() and hashlib.sha256(path.read_bytes()).hexdigest() == WHEEL_SHA256\n}\nif len(wheel_candidates) != 1:\n    raise FileNotFoundError(\n        f\"Expected one hash-verified {WHEEL_FILENAME}, found {len(wheel_candidates)}\"\n    )\nwheel_path = next(iter(wheel_candidates))\nsubprocess.run(\n    [\n        sys.executable,\n        \"-m\",\n        \"pip\",\n        \"install\",\n        \"--no-index\",\n        \"--no-deps\",\n        str(wheel_path),\n    ],\n    check=True,\n)\n\n"},{"cell_type":"code","execution_count":null,"id":"cell-004","metadata":{},"outputs":[],"source":"import matplotlib.pyplot as plt  # type: ignore[import-untyped]\nimport numpy as np\nimport pandas as pd  # type: ignore[import-untyped]\nimport torch\nimport torch.nn.functional as F\nfrom IPython.display import display  # type: ignore[import-not-found]\nfrom pydicom import dcmread\n\nfrom rsna_knee import (\n    PreprocessingConfig,\n    decode_pixels,\n    preprocess_series_result,\n    uniform_indices,\n)\nfrom rsna_knee.publication import (\n    load_reference_baseline_summary,\n    reproduce_reference_baseline,\n)\nfrom rsna_knee.training import (\n    build_initialized_model,\n    masked_binary_cross_entropy,\n)\n\nreference = load_reference_baseline_summary()\ntargets = reference[\"targets\"]\nprint(f\"Loaded the reference run for {len(targets)} prediction targets\")\n\n"},{"cell_type":"markdown","id":"cell-005","metadata":{},"source":"## What the model predicts\n\nThe competition requires one probability for each target in each study. The\nsample submission defines the output order. The code reads this order. It does\nnot use a manual list.\n\n"},{"cell_type":"code","execution_count":null,"id":"cell-006","metadata":{},"outputs":[],"source":"display(\n    pd.DataFrame(\n        {\n            \"Output position\": range(1, len(targets) + 1),\n            \"Target\": targets,\n        }\n    )\n)\n\n"},{"cell_type":"markdown","id":"cell-007","metadata":{},"source":"## What one MRI study contains\n\n`StudyInstanceUID` identifies each study. Each series has one directory of\nDICOM slices. The series table gives the anatomical plane. It also shows if an\nacquisition is fluid sensitive or uses fat suppression.\n\nThis cell uses a deterministic method to select one study with all 12 labels.\nA deterministic method gives the same result for the same input. The method\nfirst selects studies that have axial, coronal, and sagittal series. It then\nuses the first study ID in the sorted list. The code does not read report text.\n\n"},{"cell_type":"code","execution_count":null,"id":"cell-008","metadata":{},"outputs":[],"source":"competition_root = Path(\"/kaggle/input/competitions/rsna-knee-abnormality-detection\")\nrequired_training_files = (\"train.csv\", \"train_series.csv\", \"train_series\")\nif not all((competition_root / name).exists() for name in required_training_files):\n    raise FileNotFoundError(\n        f\"Competition training data not found at {competition_root}\"\n    )\ntrain_labels = pd.read_csv(\n    competition_root / \"train.csv\",\n    usecols=[\"StudyInstanceUID\", *targets],\n    dtype={\"StudyInstanceUID\": str},\n)\nseries_table = pd.read_csv(\n    competition_root / \"train_series.csv\",\n    dtype={\"StudyInstanceUID\": str, \"SeriesInstanceUID\": str},\n)\n\nfully_labeled = set(train_labels.dropna(subset=targets)[\"StudyInstanceUID\"].astype(str))\nseries_summary = (\n    series_table[series_table[\"StudyInstanceUID\"].isin(fully_labeled)]\n    .groupby(\"StudyInstanceUID\")\n    .agg(\n        series_count=(\"SeriesInstanceUID\", \"size\"),\n        plane_count=(\"Anatomical_Plane\", \"nunique\"),\n    )\n)\nthree_plane_studies = series_summary[series_summary[\"plane_count\"] >= 3]\ntwo_plane_studies = series_summary[series_summary[\"plane_count\"] >= 2]\nstudy_candidates = (\n    three_plane_studies\n    if not three_plane_studies.empty\n    else two_plane_studies\n    if not two_plane_studies.empty\n    else series_summary\n)\nstudy_uid = min(study_candidates.index)\nstudy_rows = (\n    series_table[series_table[\"StudyInstanceUID\"] == study_uid]\n    .sort_values([\"Anatomical_Plane\", \"SeriesInstanceUID\"])\n    .reset_index(drop=True)\n)\n\nseries_files: dict[str, list[Path]] = {}\nseries_inventory = []\npatient_ids: set[str] = set()\nfor row in study_rows.itertuples(index=False):\n    series_uid = str(row.SeriesInstanceUID)\n    series_dir = competition_root / \"train_series\" / study_uid / series_uid\n    paths = sorted(\n        path\n        for path in series_dir.iterdir()\n        if path.is_file() and path.suffix.casefold() in {\".dcm\", \".dicom\"}\n    )\n    if not paths:\n        raise FileNotFoundError(f\"No DICOM slices found in {series_dir}\")\n    header = dcmread(paths[0], stop_before_pixels=True)\n    patient_id = str(getattr(header, \"PatientID\", \"\"))\n    if patient_id:\n        patient_ids.add(patient_id)\n    series_files[series_uid] = paths\n    series_inventory.append(\n        {\n            \"SeriesInstanceUID\": series_uid,\n            \"Plane\": str(row.Anatomical_Plane),\n            \"Fluid sensitive\": str(row.Fluid_Sensitive),\n            \"Fat suppression\": str(row.Fat_Suppression),\n            \"Slices\": len(paths),\n            \"Matrix\": f\"{getattr(header, 'Rows', '?')} x {getattr(header, 'Columns', '?')}\",\n            \"Photometric\": str(getattr(header, \"PhotometricInterpretation\", \"\")),\n        }\n    )\n\nif len(patient_ids) != 1:\n    raise ValueError(\"Expected one PatientID across the representative study\")\npatient_id = next(iter(patient_ids))\nseries_inventory_frame = pd.DataFrame(series_inventory)\n\nprint(f\"PatientID: {patient_id}\")\nprint(f\"StudyInstanceUID: {study_uid}\")\ndisplay(series_inventory_frame)\n\n"},{"cell_type":"markdown","id":"cell-009","metadata":{},"source":"The next output shows a short directory tree. Each file name is a DICOM SOP\nInstance UID. A unique identifier (UID) is a value that identifies one item.\nThe output shows only the first and last file name in each series.\n\n"},{"cell_type":"code","execution_count":null,"id":"cell-010","metadata":{},"outputs":[],"source":"print(f\"train_series/{study_uid}/\")\nfor record in series_inventory:\n    series_uid = record[\"SeriesInstanceUID\"]\n    paths = series_files[series_uid]\n    print(f\"  {series_uid}/\")\n    print(f\"    {paths[0].name} ... {paths[-1].name}  ({len(paths)} DICOM files)\")\n\n"},{"cell_type":"markdown","id":"cell-011","metadata":{},"source":"## Order the slices by anatomical position\n\nFile names identify slices. They do not give anatomical positions. The code\ncalculates each slice coordinate from `ImagePositionPatient` and\n`ImageOrientationPatient`.\n\nA normal vector gives the direction perpendicular to the image plane. The\npackage verifies that all normal vectors have the same direction. It also\nverifies that each coordinate is unique.\n\nIf the geometry is not valid, the package uses unique `InstanceNumber` values.\nIf these values are not valid, the package uses file names. The code below is\na short example of this logic. `rsna_knee.indexing` contains the complete\nchecks and error records.\n\n\n"},{"cell_type":"code","execution_count":null,"id":"cell-012","metadata":{},"outputs":[],"source":"def slice_geometry(dataset):\n    try:\n        orientation = np.asarray(dataset.ImageOrientationPatient, dtype=np.float64)\n        position = np.asarray(dataset.ImagePositionPatient, dtype=np.float64)\n    except AttributeError:\n        return None\n    if orientation.shape != (6,) or position.shape != (3,):\n        return None\n    normal = np.cross(orientation[:3], orientation[3:])\n    norm = float(np.linalg.norm(normal))\n    if norm < 1e-6 or not np.isfinite([*normal, *position]).all():\n        return None\n    unit_normal = normal / norm\n    return unit_normal, float(np.dot(unit_normal, position))\n\n\ndef order_series(paths: list[Path]) -> tuple[list[Path], str]:\n    headers = [\n        (\n            path,\n            dcmread(\n                path,\n                stop_before_pixels=True,\n                specific_tags=[\n                    \"ImageOrientationPatient\",\n                    \"ImagePositionPatient\",\n                    \"InstanceNumber\",\n                ],\n            ),\n        )\n        for path in paths\n    ]\n    geometry = [slice_geometry(dataset) for _, dataset in headers]\n    if all(item is not None for item in geometry):\n        normals = [item[0] for item in geometry if item is not None]\n        coordinates = [item[1] for item in geometry if item is not None]\n        geometry_is_consistent = all(\n            float(np.dot(normals[0], normal)) >= 0.999 for normal in normals[1:]\n        ) and len({round(value, 6) for value in coordinates}) == len(coordinates)\n        if geometry_is_consistent:\n            return (\n                [\n                    path\n                    for _, path in sorted(\n                        zip(coordinates, paths, strict=True), key=lambda item: item[0]\n                    )\n                ],\n                \"ImagePositionPatient geometry\",\n            )\n\n    instances = [getattr(dataset, \"InstanceNumber\", None) for _, dataset in headers]\n    if all(value is not None for value in instances) and len(set(instances)) == len(\n        instances\n    ):\n        return (\n            [\n                path\n                for _, path in sorted(\n                    zip(instances, paths, strict=True), key=lambda item: int(item[0])\n                )\n            ],\n            \"InstanceNumber fallback\",\n        )\n    return sorted(paths, key=lambda path: path.name), \"filename fallback\"\n\n\nsagittal = series_inventory_frame[\n    series_inventory_frame[\"Plane\"].str.casefold() == \"sagittal\"\n]\nvisual_record = (\n    (sagittal if not sagittal.empty else series_inventory_frame)\n    .sort_values(\"Slices\", ascending=False)\n    .iloc[0]\n)\nvisual_series_uid = str(visual_record[\"SeriesInstanceUID\"])\nordered_paths, ordering_method = order_series(series_files[visual_series_uid])\nprint(f\"Series: {visual_series_uid}\")\nprint(f\"Plane: {visual_record['Plane']}\")\nprint(f\"Ordering used: {ordering_method}\")\nprint(f\"Ordered slices: {len(ordered_paths)}\")\n\n"},{"cell_type":"markdown","id":"cell-013","metadata":{},"source":"`InstanceNumber` can increase while the spatial coordinate decreases. This\ndoes not indicate an error. DICOM geometry locates slices in the patient\ncoordinate system. Instance numbers and file names do not always follow the\nanatomical direction. Thus, the package uses valid geometry as the primary\norder. It uses `InstanceNumber` only when the geometry is not valid.\n\n"},{"cell_type":"markdown","id":"cell-014","metadata":{},"source":"## Select eight slices from the volume\n\nThe baseline selects integer indices at equal intervals. It includes the first\nand last slice. NumPy uses truncation to make the intermediate positions into\nintegers. Truncation removes the decimal part. It does not round the value.\n\nIf a series has eight or fewer usable slices, the code selects each slice one\ntime. Unused model slots contain zero. A mask records which slots contain valid\nslices. A false mask value marks an unused slot or a decode error. The model\ndoes not include these slots in the mean.\n\n\n"},{"cell_type":"code","execution_count":null,"id":"cell-015","metadata":{},"outputs":[],"source":"# This example is simplified. It matches rsna_knee.preprocessing.uniform_indices.\ndef teaching_slice_indices(n_slices: int, n_selected: int = 8) -> tuple[int, ...]:\n    if n_slices < 0 or n_selected <= 0:\n        raise ValueError(\"n_slices must be non-negative and n_selected positive\")\n    if n_slices <= n_selected:\n        return tuple(range(n_slices))\n    return tuple(np.linspace(0, n_slices - 1, n_selected, dtype=int).tolist())\n\n\nfor example_length in (0, 5, 8, 9, 31):\n    assert teaching_slice_indices(example_length) == uniform_indices(example_length, 8)\n\nselected_indices = teaching_slice_indices(len(ordered_paths))\ndisplay(\n    pd.DataFrame(\n        {\n            \"Model slot\": range(1, len(selected_indices) + 1),\n            \"Ordered slice index\": selected_indices,\n            \"Position in series\": [\n                f\"{index + 1}/{len(ordered_paths)}\" for index in selected_indices\n            ],\n        }\n    )\n)\n\n"},{"cell_type":"markdown","id":"cell-016","metadata":{},"source":"## View the selected MRI slices\n\nThe preprocessing step uses these slices from the selected series. The slices\ncover the ordered volume. They are not eight adjacent images.\n\n"},{"cell_type":"code","execution_count":null,"id":"cell-017","metadata":{},"outputs":[],"source":"figure, axes = plt.subplots(2, 4, figsize=(16, 8))\nfor slot, (axis, source_index) in enumerate(zip(axes.flat, selected_indices), start=1):\n    path = ordered_paths[source_index]\n    header = dcmread(path, stop_before_pixels=True, specific_tags=[\"InstanceNumber\"])\n    axis.imshow(decode_pixels(path), cmap=\"gray\")\n    axis.set_title(\n        f\"slot {slot}: {source_index + 1}/{len(ordered_paths)}\\n\"\n        f\"InstanceNumber {getattr(header, 'InstanceNumber', '?')}\"\n    )\n    axis.axis(\"off\")\nfor axis in list(axes.flat)[len(selected_indices) :]:\n    axis.axis(\"off\")\nfigure.suptitle(\n    f\"{visual_record['Plane']} series: eight slices across the volume\", fontsize=16\n)\nfigure.tight_layout()\nplt.show()\n\n"},{"cell_type":"markdown","id":"cell-018","metadata":{},"source":"### Compare the anatomical planes\n\nThe code shows one middle slice from each available plane. It uses the same\nstudy and does not scan another study.\n\n"},{"cell_type":"code","execution_count":null,"id":"cell-019","metadata":{},"outputs":[],"source":"orientation_examples = []\nfor plane in (\"Axial\", \"Coronal\", \"Sagittal\"):\n    matching_series = series_inventory_frame[\n        series_inventory_frame[\"Plane\"].str.casefold() == plane.casefold()\n    ]\n    if matching_series.empty:\n        continue\n    record = matching_series.sort_values(\"Slices\", ascending=False).iloc[0]\n    plane_paths, _ = order_series(series_files[str(record[\"SeriesInstanceUID\"])])\n    middle_path = plane_paths[(len(plane_paths) - 1) // 2]\n    orientation_examples.append((plane, middle_path))\n\nfigure, axes = plt.subplots(\n    1,\n    len(orientation_examples),\n    figsize=(4 * len(orientation_examples), 4),\n    squeeze=False,\n)\nfor axis, (plane, path) in zip(axes.flat, orientation_examples, strict=True):\n    axis.imshow(decode_pixels(path), cmap=\"gray\")\n    axis.set_title(plane)\n    axis.axis(\"off\")\nfigure.tight_layout()\nplt.show()\n\n"},{"cell_type":"markdown","id":"cell-020","metadata":{},"source":"## Make a model input from one DICOM slice\n\nPixel decoding converts stored DICOM values to image intensities. A lookup\ntable (LUT) defines how to convert one value to another. The code applies the\nDICOM modality LUT when the file contains one.\n\nIf there is no LUT, the code uses `RescaleSlope` and `RescaleIntercept`. The\ncode inverts a `MONOCHROME1` image. Thus, larger values stay visually brighter.\n\nThe code calculates the 1st and 99th percentiles from all selected slices that\nit can decode. Each slice uses the same limits for its series. The code clips\neach value to the range `[0, 1]`.\n\nAspect ratio is the relation between image width and image height. The code\nkeeps this ratio when it changes the image size. It adds zero padding to make a\n`224 x 224` image. ResNet-18 requires three image channels. Thus, the code\nmakes three identical channels from the grayscale image.\n\nImageNet is an image dataset used to initialize ResNet-18. ImageNet\nnormalization uses fixed channel means and standard deviations. The code\napplies this normalization before ResNet-18 processes the image.\n\n"},{"cell_type":"code","execution_count":null,"id":"cell-021","metadata":{},"outputs":[],"source":"preprocessing_config = PreprocessingConfig(\n    slice_count=int(reference[\"preprocessing\"][\"slice_count\"]),\n    lower_percentile=float(reference[\"preprocessing\"][\"lower_percentile\"]),\n    upper_percentile=float(reference[\"preprocessing\"][\"upper_percentile\"]),\n    image_size=int(reference[\"preprocessing\"][\"image_size\"]),\n)\nprocessed_series = preprocess_series_result(ordered_paths, preprocessing_config)\nvolume_middle = (len(ordered_paths) - 1) / 2\nvalid_slot = min(\n    (slot for slot, is_valid in enumerate(processed_series.mask) if is_valid),\n    key=lambda slot: abs(processed_series.source_indices[slot] - volume_middle),\n)\nsource_index = processed_series.source_indices[valid_slot]\nexample_path = ordered_paths[source_index]\nexample_dataset = dcmread(example_path)\nstored_pixels = np.asarray(example_dataset.pixel_array)\ndecoded_pixels = decode_pixels(example_path)\nlower = float(processed_series.lower_value)\nupper = float(processed_series.upper_value)\nscaled_pixels = (\n    np.zeros_like(decoded_pixels, dtype=np.float32)\n    if lower == upper\n    else np.clip((decoded_pixels - lower) / (upper - lower), 0, 1).astype(np.float32)\n)\nresized_pixels = processed_series.pixels[valid_slot]\n\nfigure, axes = plt.subplots(1, 4, figsize=(18, 4))\nstages = (\n    (stored_pixels, f\"Stored DICOM pixels\\n{stored_pixels.dtype}\"),\n    (decoded_pixels, \"Corrected pixel values\"),\n    (scaled_pixels, f\"Scaled to percentile limits\\n{lower:.1f} to {upper:.1f}\"),\n    (\n        resized_pixels,\n        f\"Model input\\n{resized_pixels.shape[0]} x {resized_pixels.shape[1]}\",\n    ),\n)\nfor axis, (image, title) in zip(axes, stages, strict=True):\n    axis.imshow(image, cmap=\"gray\")\n    axis.set_title(title)\n    axis.axis(\"off\")\nfigure.tight_layout()\nplt.show()\n\nvalid_planes = torch.from_numpy(processed_series.pixels[processed_series.mask])\nthree_channel = valid_planes.unsqueeze(1).repeat(1, 3, 1, 1)\nimagenet_mean = torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1)\nimagenet_std = torch.tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1)\nnormalized_planes = (three_channel - imagenet_mean) / imagenet_std\nprint(f\"grayscale batch: {tuple(valid_planes.shape)}\")\nprint(f\"ResNet-18 input: {tuple(normalized_planes.shape)}\")\n\n"},{"cell_type":"markdown","id":"cell-022","metadata":{},"source":"## Keep each patient in one data split\n\nStudies from the same patient are related. Data leakage occurs when related\npatient data occurs in both data splits. Leakage can make a validation score\ntoo high.\n\nThe split uses `PatientID`. Each patient occurs in only one data split. All 58\nselected studies have known competition labels for all 12 targets. A holdout\nis the data split used for validation. The loss can ignore an unknown label.\nThis subset does not contain unknown labels.\n\n"},{"cell_type":"code","execution_count":null,"id":"cell-023","metadata":{},"outputs":[],"source":"split = reference[\"split\"]\npartitions = split[\"partitions\"]\ndisplay(\n    pd.DataFrame(\n        [\n            {\n                \"Split\": name.capitalize(),\n                \"Studies\": partitions[name][\"study_count\"],\n                \"Patients\": partitions[name][\"patient_count\"],\n            }\n            for name in (\"train\", \"holdout\")\n        ]\n    )\n)\nprint(f\"Patients shared by train and holdout: {split['patient_overlap']}\")\n\n"},{"cell_type":"code","execution_count":null,"id":"cell-024","metadata":{},"outputs":[],"source":"class_balance = []\nfor target in targets:\n    train_counts = partitions[\"train\"][\"per_target\"][target]\n    holdout_counts = partitions[\"holdout\"][\"per_target\"][target]\n    class_balance.append(\n        {\n            \"Target\": target,\n            \"Train positive\": train_counts[\"positive\"],\n            \"Train negative\": train_counts[\"negative\"],\n            \"Holdout positive\": holdout_counts[\"positive\"],\n            \"Holdout negative\": holdout_counts[\"negative\"],\n        }\n    )\ndisplay(pd.DataFrame(class_balance))\n\n"},{"cell_type":"markdown","id":"cell-025","metadata":{},"source":"## Make one study prediction from DICOM slices\n\nAn encoder converts an input image to a feature vector. ResNet-18 processes\neach valid slice with the same encoder and weights. A feature vector is also\ncalled an embedding. It is a numeric description that the model learns.\n\nThe model averages only valid embeddings. This operation is a masked mean.\nEmpty padding slots do not change the result. Thus, a short series does not get\nextra weight from empty slots.\n\n```text\nMRI study\n  -> one or more image series\n     -> 8 slices selected by a fixed method in each series\n        -> decode DICOM pixels, scale values, and resize to 224 x 224\n        -> process each slice with the shared ResNet-18 encoder\n     -> calculate the mean of valid slice embeddings\n  -> calculate the mean of valid series embeddings\n  -> make 12 logits and sigmoid probabilities\n```\n\nThe fixed selection method covers the full ordered series. It also makes memory\nuse and run time more predictable.\n\n"},{"cell_type":"code","execution_count":null,"id":"cell-026","metadata":{},"outputs":[],"source":"model_record = reference[\"model\"]\npreprocessing = reference[\"preprocessing\"]\ntraining = reference[\"training\"]\ndisplay(\n    pd.DataFrame(\n        [\n            {\"Setting\": \"Backbone\", \"Value\": \"ResNet-18\"},\n            {\n                \"Setting\": \"Initialization\",\n                \"Value\": model_record[\"initialization_identity\"],\n            },\n            {\n                \"Setting\": \"Image size\",\n                \"Value\": f\"{model_record['image_size']} x {model_record['image_size']}\",\n            },\n            {\"Setting\": \"Slices per series\", \"Value\": preprocessing[\"slice_count\"]},\n            {\n                \"Setting\": \"Slice pooling\",\n                \"Value\": model_record[\"slice_to_series_pooling\"],\n            },\n            {\n                \"Setting\": \"Series pooling\",\n                \"Value\": model_record[\"series_to_study_pooling\"],\n            },\n            {\"Setting\": \"Optimizer\", \"Value\": \"AdamW\"},\n            {\"Setting\": \"Epochs\", \"Value\": training[\"epochs\"]},\n            {\"Setting\": \"Learning rate\", \"Value\": training[\"learning_rate\"]},\n            {\"Setting\": \"Weight decay\", \"Value\": training[\"weight_decay\"]},\n            {\"Setting\": \"Mixed precision\", \"Value\": training[\"mixed_precision\"]},\n        ]\n    )\n)\n\n"},{"cell_type":"markdown","id":"cell-027","metadata":{},"source":"## Track the tensors through the model\n\nA tensor is a multidimensional array. The code uses tensors for images, masks,\nfeatures, and outputs. The same ResNet-18 encoder processes every valid slice.\nThe example removes the classification layer. Thus, the encoder makes one\nfeature vector with 512 values for each slice.\n\nThe slice mask has shape `[study, series, slice]`. A false value marks padding\nor a decode error. The series mask has shape `[study, series]`. Its value is\nfalse when a series has no valid slices.\n\nThe code applies each mask before it calculates a mean. Thus, padded zeros do\nnot change the study feature vector.\n\nThe next cell is a short CPU example. It uses a maximum of three series from\nthe selected study. The complete run uses every indexed series. The example\nloads the attached ImageNet weights. The package verifies the exact weight\nfile. The example does not train the 12-output head. It shows only the tensor\npath and the tensor shapes.\n\n"},{"cell_type":"code","execution_count":null,"id":"cell-028","metadata":{},"outputs":[],"source":"# This example is simplified. The tested forward pass is in rsna_knee.training.\nteaching_series_uids = [visual_series_uid] + [\n    uid\n    for uid in series_inventory_frame[\"SeriesInstanceUID\"].astype(str)\n    if uid != visual_series_uid\n]\nteaching_series_uids = teaching_series_uids[:3]\nteaching_results = [\n    preprocess_series_result(\n        order_series(series_files[series_uid])[0], preprocessing_config\n    )\n    for series_uid in teaching_series_uids\n]\nstudy_pixels = torch.from_numpy(\n    np.stack([result.pixels for result in teaching_results])\n).unsqueeze(0)\nslice_mask = torch.from_numpy(\n    np.stack([result.mask for result in teaching_results])\n).unsqueeze(0)\nseries_mask = slice_mask.any(dim=2)\n\nweight_roots = (\n    Path(\"/kaggle/input/datasets\") / WEIGHT_DATASET_ID,\n    Path(\"/kaggle/input\") / WEIGHT_DATASET_ID.partition(\"/\")[2],\n)\nweight_candidates = {\n    path.resolve()\n    for root in weight_roots\n    if root.exists()\n    for path in root.rglob(\"resnet18-f37072fd.pth\")\n    if path.is_file()\n}\nif len(weight_candidates) != 1:\n    raise FileNotFoundError(\n        f\"Expected one offline ResNet-18 weight file, found {len(weight_candidates)}\"\n    )\n\ntorch.manual_seed(int(training[\"seed\"]))\nteaching_model, teaching_initialization = build_initialized_model(\n    next(iter(weight_candidates)), targets\n)\nteaching_model.eval()\n\nvalid = slice_mask & series_mask.unsqueeze(-1)\nflat_valid = valid.reshape(-1)\nflat_pixels = study_pixels.reshape(-1, 224, 224)\nvalid_images = flat_pixels[flat_valid].unsqueeze(1).repeat(1, 3, 1, 1)\nvalid_images = (\n    valid_images - teaching_model.imagenet_mean\n) / teaching_model.imagenet_std\n\nwith torch.inference_mode():\n    valid_slice_features = teaching_model.backbone(valid_images)\n    flat_features = valid_slice_features.new_zeros(\n        (flat_pixels.shape[0], teaching_model.feature_count)\n    )\n    flat_features[flat_valid] = valid_slice_features\n    slice_features = flat_features.view(\n        *study_pixels.shape[:3], teaching_model.feature_count\n    )\n\n    slice_weights = valid.unsqueeze(-1).to(slice_features.dtype)\n    series_features = (slice_features * slice_weights).sum(dim=2) / (\n        slice_weights.sum(dim=2).clamp_min(1)\n    )\n\n    effective_series = valid.any(dim=2) & series_mask\n    series_weights = effective_series.unsqueeze(-1).to(series_features.dtype)\n    study_features = (series_features * series_weights).sum(dim=1) / (\n        series_weights.sum(dim=1).clamp_min(1)\n    )\n    logits = teaching_model.head(study_features)\n\ndisplay(\n    pd.DataFrame(\n        [\n            {\"Tensor\": \"study_pixels\", \"Shape\": tuple(study_pixels.shape)},\n            {\"Tensor\": \"slice_mask\", \"Shape\": tuple(slice_mask.shape)},\n            {\"Tensor\": \"valid_images\", \"Shape\": tuple(valid_images.shape)},\n            {\n                \"Tensor\": \"slice_features\",\n                \"Shape\": tuple(slice_features.shape),\n            },\n            {\n                \"Tensor\": \"series_features\",\n                \"Shape\": tuple(series_features.shape),\n            },\n            {\n                \"Tensor\": \"study_features\",\n                \"Shape\": tuple(study_features.shape),\n            },\n            {\"Tensor\": \"logits\", \"Shape\": tuple(logits.shape)},\n        ]\n    )\n)\nprint(f\"Verified weight SHA-256: {teaching_initialization.sha256}\")\n\n"},{"cell_type":"markdown","id":"cell-029","metadata":{},"source":"A logit is the model output before conversion to a probability. The output\nhead makes one logit for each target. A sigmoid function converts each logit to\na probability.\n\nThis is a multi-label problem. Multi-label means that one study can have more\nthan one positive target.\n\n"},{"cell_type":"markdown","id":"cell-030","metadata":{},"source":"## Calculate the multi-label loss\n\nBinary cross-entropy (BCE) is a loss for one binary target. Each study has 12\nbinary targets. Thus, the loss calculation starts with 12 separate BCE values.\n\nThe known-label mask selects the values for the mean. The reference group has\nall 12 labels for every selected study. The training code can also process a\nrow that has unknown labels.\n\n"},{"cell_type":"code","execution_count":null,"id":"cell-031","metadata":{},"outputs":[],"source":"# This example uses synthetic labels. They make the mask operation easy to see.\nexample_logits = torch.linspace(-1.5, 1.5, len(targets)).view(1, -1)\nexample_logits.requires_grad_()\nexample_targets = torch.tensor(\n    [[float(index % 3 == 0) for index in range(len(targets))]]\n)\nknown_mask = torch.ones_like(example_targets, dtype=torch.bool)\nknown_mask[:, -2:] = False\n\nelementwise_bce = F.binary_cross_entropy_with_logits(\n    example_logits, example_targets, reduction=\"none\"\n)\nteaching_loss = elementwise_bce[known_mask].mean()\nproduction_loss = masked_binary_cross_entropy(\n    example_logits, example_targets, known_mask\n)\nassert torch.allclose(teaching_loss, production_loss)\nteaching_loss.backward()\n\ndisplay(\n    pd.DataFrame(\n        {\n            \"Target\": targets,\n            \"Logit\": example_logits.detach().squeeze(0),\n            \"Binary label\": example_targets.squeeze(0),\n            \"Known\": known_mask.squeeze(0),\n            \"Element BCE\": elementwise_bce.detach().squeeze(0),\n        }\n    )\n)\nprint(f\"Mean BCE over known targets: {teaching_loss.item():.4f}\")\nprint(\n    \"Gradient for the two unknown targets:\",\n    example_logits.grad[0, -2:].tolist(),\n)\n\n"},{"cell_type":"markdown","id":"cell-032","metadata":{},"source":"## Reference run results\n\nThe receiver operating characteristic area under the curve (ROC AUC) measures\nhow the model ranks positive and negative examples. The competition metric is\nthe mean of the ROC AUC for the 12 targets.\n\nThe table gives the holdout class counts. Each target has data from only four\npatients. With four patients, only a few positive and negative pairs determine\nthe AUC. An AUC of 0 does not prove that the model always fails. An AUC of 1\ndoes not prove that the model is perfect.\n\n"},{"cell_type":"code","execution_count":null,"id":"cell-033","metadata":{},"outputs":[],"source":"results = reference[\"results\"]\nresult_rows = []\nfor target in targets:\n    target_result = results[\"per_target\"][target]\n    result_rows.append(\n        {\n            \"Target\": target,\n            \"Holdout positive\": target_result[\"positive\"],\n            \"Holdout negative\": target_result[\"negative\"],\n            \"ROC AUC\": target_result[\"roc_auc\"],\n        }\n    )\ndisplay(pd.DataFrame(result_rows))\nprint(f\"Observed 12-target macro ROC AUC: {results['macro_auc']:.4f}\")\n\n"},{"cell_type":"markdown","id":"cell-034","metadata":{},"source":"## Why four validation patients are too few\n\nThe run uses bootstrap resampling to examine uncertainty. For each resample,\nthe method selects four patients with replacement. Thus, it can select the\nsame patient more than one time. The run makes 10,000 bootstrap resamples.\n\nROC AUC is not defined when a resample has only one class for a target. The\nmacro AUC needs a valid AUC for all 12 targets. Thus, one undefined target\nmakes the macro AUC undefined for that resample.\n\n"},{"cell_type":"code","execution_count":null,"id":"cell-035","metadata":{},"outputs":[],"source":"limitations = reference[\"limitations\"]\ndefined = limitations[\"defined_macro_replicate_count\"]\nundefined = limitations[\"undefined_macro_replicate_count\"]\ntotal = limitations[\"bootstrap_replicate_count\"]\ndisplay(\n    pd.DataFrame(\n        [\n            {\n                \"Quantity\": \"Holdout patients\",\n                \"Value\": limitations[\"holdout_patient_count\"],\n            },\n            {\"Quantity\": \"Observed macro ROC AUC\", \"Value\": results[\"macro_auc\"]},\n            {\"Quantity\": \"Bootstrap resamples\", \"Value\": total},\n            {\n                \"Quantity\": \"Resamples with a defined macro AUC\",\n                \"Value\": f\"{defined} ({defined / total:.2%})\",\n            },\n            {\n                \"Quantity\": \"Resamples with an undefined macro AUC\",\n                \"Value\": f\"{undefined} ({undefined / total:.2%})\",\n            },\n            {\"Quantity\": \"Evidence for model selection\", \"Value\": \"Inadequate\"},\n        ]\n    )\n)\n\n"},{"cell_type":"markdown","id":"cell-036","metadata":{},"source":"In this run, 9,065 of 10,000 resamples (90.65%) had no macro AUC. The other 935\nresamples gave the same value. This does not show that the score is precise.\nThe small holdout caused this result. The calculation also excludes all\nresamples that have an undefined macro AUC.\n\nThe pipeline completed training, evaluation, checkpoint reload, and submission\nvalidation. The value 0.6111 describes only these four holdout patients. It\ndoes not estimate expected leaderboard or clinical performance. Thus, the\nobserved macro AUC is not reliable enough for model selection.\n\n"},{"cell_type":"markdown","id":"cell-037","metadata":{},"source":"## Run time\n\nThe recorded 5.1 minutes includes ten training epochs and their holdout\nevaluations on a Kaggle Tesla T4. It does not include the complete notebook\nrun. The previous complete notebook run took about 147 minutes.\n\nA DICOM header index records file paths and DICOM attributes. During the\ncomplete run, the package rebuilt and verified this index. It also reloaded\nthe checkpoint and made test predictions. Then, it validated the submission.\nIndex building and DICOM decoding used most of the complete run time.\n\n"},{"cell_type":"code","execution_count":null,"id":"cell-038","metadata":{},"outputs":[],"source":"runtime = reference[\"runtime\"]\ncheckpoint = reference[\"checkpoint\"]\ndisplay(\n    pd.DataFrame(\n        [\n            {\"Measurement\": \"GPU\", \"Value\": runtime[\"gpu\"]},\n            {\n                \"Measurement\": \"Training and holdout evaluation only\",\n                \"Value\": f\"{runtime['complete_run_seconds'] / 60:.1f} minutes\",\n            },\n            {\n                \"Measurement\": \"Previous complete notebook run\",\n                \"Value\": \"about 147 minutes\",\n            },\n            {\n                \"Measurement\": \"Training DICOM decode time\",\n                \"Value\": f\"{runtime['training_decode_seconds']:.1f} seconds\",\n            },\n            {\n                \"Measurement\": \"Training forward and backward time\",\n                \"Value\": f\"{runtime['training_forward_backward_seconds']:.1f} seconds\",\n            },\n            {\n                \"Measurement\": \"Peak CUDA memory\",\n                \"Value\": f\"{runtime['peak_cuda_memory_bytes'] / (1024**3):.2f} GiB\",\n            },\n            {\n                \"Measurement\": \"Valid training slices each second\",\n                \"Value\": f\"{runtime['training_slices_per_second']:.1f}\",\n            },\n            {\n                \"Measurement\": \"Training decode errors\",\n                \"Value\": len(runtime[\"training_decode_errors\"]),\n            },\n            {\n                \"Measurement\": \"Best checkpoint epoch\",\n                \"Value\": checkpoint[\"best\"][\"epoch\"],\n            },\n        ]\n    )\n)\n\n"},{"cell_type":"markdown","id":"cell-039","metadata":{},"source":"## Reproducibility details\n\nReproducibility means that another run can use the same source and settings.\nThe teaching cells use short examples. The package does the complete run with\nthe exact source and configuration bytes.\n\nThe notebook does not use internet access. The package verifies the hashes for\nthe wheel, ImageNet weights, and configuration. A split fingerprint identifies\nthe exact patient split. The package also verifies the saved checkpoint after\nit reloads the checkpoint.\n\nA dirty working tree has uncommitted source changes. The table reports if the\nworking tree was dirty when the build process made the notebook.\n\n"},{"cell_type":"code","execution_count":null,"id":"cell-040","metadata":{},"outputs":[],"source":"config_bytes = base64.b64decode(\"W3ByZXByb2Nlc3NpbmddCnNsaWNlX2NvdW50ID0gOApsb3dlcl9wZXJjZW50aWxlID0gMQp1cHBlcl9wZXJjZW50aWxlID0gOTkKaW1hZ2Vfc2l6ZSA9IDIyNAoKW3RyYWluaW5nXQplcG9jaHMgPSAxMApiYXRjaF9zaXplID0gMQpsZWFybmluZ19yYXRlID0gMC4wMDAxCndlaWdodF9kZWNheSA9IDAuMDAwMQpzZWVkID0gMjAyNjA4MTIKbWl4ZWRfcHJlY2lzaW9uID0gdHJ1ZQphcmNoaXRlY3R1cmUgPSAidG9yY2h2aXNpb24ubW9kZWxzLnJlc25ldDE4Igp3ZWlnaHRfaWRlbnRpdHkgPSAiUmVzTmV0MThfV2VpZ2h0cy5JTUFHRU5FVDFLX1YxIgoKW2Jvb3RzdHJhcF0KbWV0aG9kX3ZlcnNpb24gPSAicGF0aWVudC1pZC1ub25wYXJhbWV0cmljLWJvb3RzdHJhcC12MSIKc2VlZCA9IDIwMjYwODEyCnJlcGxpY2F0ZV9jb3VudCA9IDEwMDAwCmNvbmZpZGVuY2VfbGV2ZWwgPSAwLjk1CgpbdmFsaWRhdGlvbl0KcmVsb2FkX2Fic29sdXRlX3RvbGVyYW5jZSA9IDAuMDAwMDAxCg==\")\nconfig_sha256 = hashlib.sha256(config_bytes).hexdigest()\nif config_sha256 != \"089ec95b4d160607f219a5f81e06af027c7c03329f32051e67ee36128107694c\":\n    raise RuntimeError(\"Baseline configuration hash mismatch\")\n\ndisplay(\n    pd.DataFrame(\n        [\n            {\"Item\": \"Git commit\", \"Value\": \"03ef9f4416b3a85061fad22df047fc36fca31e75\"},\n            {\n                \"Item\": \"Working tree was dirty at build\",\n                \"Value\": json.loads(\"false\"),\n            },\n            {\"Item\": \"Project wheel SHA-256\", \"Value\": WHEEL_SHA256},\n            {\"Item\": \"Configuration SHA-256\", \"Value\": config_sha256},\n            {\n                \"Item\": \"Jupytext source SHA-256\",\n                \"Value\": \"604383db446fcf02cf59c206518e4cc186f6e88dc6d34914de5bb3eb4247a98a\",\n            },\n            {\n                \"Item\": \"ImageNet weight SHA-256\",\n                \"Value\": model_record[\"initialization_artifact_sha256\"],\n            },\n            {\"Item\": \"Patient split fingerprint\", \"Value\": split[\"fingerprint\"]},\n            {\n                \"Item\": \"Recorded best checkpoint SHA-256\",\n                \"Value\": checkpoint[\"best\"][\"sha256\"],\n            },\n            {\n                \"Item\": \"Checkpoint reload verification passed\",\n                \"Value\": checkpoint[\"reload_predictions_within_tolerance\"],\n            },\n            {\"Item\": \"Kaggle kernel\", \"Value\": KERNEL_ID},\n            {\"Item\": \"Internet enabled\", \"Value\": False},\n        ]\n    )\n)\n\n"},{"cell_type":"markdown","id":"cell-041","metadata":{},"source":"## Run the complete pipeline and validate test predictions\n\nThe final cell repeats the complete run with the attached competition data.\nThe notebook contains the fixed configuration. The package checks the\nconfiguration before use.\n\nThe package builds the index and preprocesses the DICOM slices. It trains the\nmodel and evaluates the holdout data. Then, it reloads the checkpoint. It makes\ntest predictions and validates the submission.\n\nThe temporary index contains row identifiers. A run artifact is a file that a\nrun makes. Checkpoints, predictions, and the submission are run artifacts. The\ncell writes these files under `/kaggle/working`. The package removes the files\nbefore the notebook shows the summary.\n\n"},{"cell_type":"code","execution_count":null,"id":"cell-042","metadata":{},"outputs":[],"source":"reproduction = reproduce_reference_baseline(\n    input_root=Path(\"/kaggle/input\"),\n    work_parent=Path(\"/kaggle/working\"),\n    config_bytes=config_bytes,\n    approved_gpu=GPU_APPROVED,\n    source_provenance={\n        \"git_commit\": \"03ef9f4416b3a85061fad22df047fc36fca31e75\",\n        \"git_dirty\": json.loads(\"false\"),\n        \"jupytext_source_sha256\": \"604383db446fcf02cf59c206518e4cc186f6e88dc6d34914de5bb3eb4247a98a\",\n        \"package_archive_sha256\": WHEEL_SHA256,\n        \"baseline_config_sha256\": \"089ec95b4d160607f219a5f81e06af027c7c03329f32051e67ee36128107694c\",\n        \"weight_dataset_id\": WEIGHT_DATASET_ID,\n    },\n    kernel_id=KERNEL_ID,\n)\nreproduction[\"notebook_elapsed_minutes\"] = round(\n    (time.perf_counter() - notebook_started) / 60, 1\n)\n\nlabels = {\n    \"training_outputs_verified\": \"Training output verification passed\",\n    \"patient_split_reproduced\": \"Patient split matched the reference\",\n    \"best_checkpoint_used_for_test_predictions\": \"Test predictions used the best checkpoint\",\n    \"submission_valid\": \"Submission schema valid\",\n    \"submission_rows\": \"Submission rows\",\n    \"test_predictions\": \"Test predictions\",\n    \"fallback_predictions\": \"Fallback predictions\",\n    \"decode_failures\": \"Inference decode failures\",\n    \"temporary_files_removed\": \"Temporary file removal passed\",\n    \"notebook_elapsed_minutes\": \"Complete notebook run time (minutes)\",\n}\ndisplay(\n    pd.DataFrame(\n        [{\"Check\": labels[key], \"Result\": value} for key, value in reproduction.items()]\n    )\n)\n\n"},{"cell_type":"markdown","id":"cell-043","metadata":{},"source":"## Limits of the reference run\n\nThe notebook shows that the study-level MRI pipeline can run from input to\nsubmission. The data split keeps patients separate. The pipeline can process\ndifferent numbers of slices and series. It reloads the saved checkpoint for\ninference. It also compares the submission rows and columns with the\ncompetition data.\n\nThis run does not show that ResNet-18 is the best architecture. It also does\nnot show how the model will perform on new data. More independent validation\npatients are necessary for these conclusions. This run does not change the\nsplit, tune the model, or add supervision.\n"}],"metadata":{"jupytext":{"formats":"py:percent","text_representation":{"extension":".py","format_name":"percent","format_version":"1.3","jupytext_version":"1.17"}},"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12"}},"nbformat":4,"nbformat_minor":5}