{"cells":[{"cell_type":"markdown","metadata":{},"source":"# RSNA Knee ConvNeXt V2 baseline\n\nFrozen ConvNeXt V2 encoder with a regularized multilabel head. The input is one\ncentral slice from each of the sagittal, coronal and axial series.\n"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"import subprocess, sys\nsubprocess.check_call([sys.executable, '-m', 'pip', 'install', '-q', 'timm'])\n"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"from pathlib import Path\nimport random, warnings\nimport numpy as np\nimport pandas as pd\nimport pydicom\nfrom PIL import Image\nimport torch\nfrom torch import nn\n\nSEED = 42\nrandom.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)\ndevice = 'cpu'\nlabels = ['ACL', 'MCL', 'Medial Meniscus', 'Lateral Meniscus', 'Medial OA', 'Lateral OA', 'PF OA', 'Effusion', 'Synovitis', \"Baker's\", 'Contusion', 'Fracture']\nroot = next(p for p in Path('/kaggle/input').glob('**/train.csv') if (p.parent / 'sample_submission.csv').exists())\ntrain = pd.read_csv(root); test = pd.read_csv(root.parent / 'test.csv')\ntrain_series = pd.read_csv(root.parent / 'train_series.csv'); test_series = pd.read_csv(root.parent / 'test_series.csv')\ngold = train.dropna(subset=labels).reset_index(drop=True)\nprint('gold studies:', len(gold), 'test studies:', len(test))\n"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"def choose_series(df):\n    x = df.copy()\n    x['_score'] = 2 * pd.to_numeric(x.get('Fluid_Sensitive', 0), errors='coerce').fillna(0) + pd.to_numeric(x.get('Fat_Suppression', 0), errors='coerce').fillna(0)\n    plane = x.get('Anatomical_Plane', '').fillna('').astype(str).str.lower()\n    x.loc[plane.str.contains('sagittal'), '_score'] += 2\n    x.loc[plane.str.contains('coronal'), '_score'] += 1\n    selected = {}\n    for study, group in x.groupby('StudyInstanceUID'):\n        chosen = []\n        for wanted in ['sagittal', 'coronal', 'axial']:\n            match = group[group['Anatomical_Plane'].fillna('').astype(str).str.lower().eq(wanted)]\n            if len(match): chosen.append(match.sort_values('_score', ascending=False).iloc[0]['SeriesInstanceUID'])\n        if len(chosen) < 3:\n            for sid in group.sort_values('_score', ascending=False)['SeriesInstanceUID']:\n                if sid not in chosen: chosen.append(sid)\n                if len(chosen) == 3: break\n        selected[study] = chosen[:3]\n    return selected\n\ndef ordered_dicoms(folder):\n    rows = []\n    for path in Path(folder).glob('*.dcm'):\n        try:\n            ds = pydicom.dcmread(str(path), stop_before_pixels=True, force=True)\n            pos = np.asarray(getattr(ds, 'ImagePositionPatient', []), dtype='float32')\n            ori = np.asarray(getattr(ds, 'ImageOrientationPatient', []), dtype='float32')\n            if pos.size == 3 and ori.size == 6:\n                coord = float(np.dot(pos, np.cross(ori[:3], ori[3:])))\n            else: coord = float(getattr(ds, 'InstanceNumber', 0))\n        except Exception: coord = 0.0\n        rows.append((coord, path))\n    return [path for _, path in sorted(rows, key=lambda item: item[0])]\n\ndef read_slice(path, size=224):\n    ds = pydicom.dcmread(str(path), force=True)\n    arr = ds.pixel_array.astype('float32')\n    arr = arr * float(getattr(ds, 'RescaleSlope', 1.0)) + float(getattr(ds, 'RescaleIntercept', 0.0))\n    lo, hi = np.percentile(arr, [1, 99]) if arr.max() > arr.min() else (arr.min(), arr.min() + 1)\n    arr = np.clip((arr - lo) / (hi - lo), 0, 1)\n    return np.asarray(Image.fromarray((arr * 255).astype('uint8')).resize((size, size)), dtype='uint8')\n\ndef study_image(study, series_map, image_root):\n    images = []\n    for sid in series_map.get(study, []):\n        paths = ordered_dicoms(image_root / str(study) / str(sid))\n        if paths: images.append(read_slice(paths[len(paths) // 2]))\n    while len(images) < 3: images.append(np.zeros((224, 224), dtype='uint8'))\n    return np.stack(images[:3], axis=-1)\n"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"import timm\nfrom timm.data import resolve_data_config, create_transform\n\nmodel = timm.create_model(\n    'convnextv2_tiny.fcmae_ft_in22k_in1k', pretrained=True,\n    num_classes=0, global_pool='avg').to(device).eval()\nconfig = resolve_data_config({}, model=model)\npreprocess = create_transform(**config, is_training=False)\nseries_tr = choose_series(train_series); series_te = choose_series(test_series)\ntrain_root = root.parent / 'train_series'; test_root = root.parent / 'test_series'\n\ndef encode(frame, series_map, image_root):\n    feats = []\n    with torch.no_grad():\n        for j, study in enumerate(frame['StudyInstanceUID']):\n            image = Image.fromarray(study_image(study, series_map, image_root))\n            x = preprocess(image).unsqueeze(0)\n            feats.append(model(x).float().cpu().numpy()[0])\n            if (j + 1) % 20 == 0: print('encoded', j + 1, '/', len(frame))\n    return np.asarray(feats, dtype='float32')\n\nX = encode(gold, series_tr, train_root)\nXt = encode(test, series_te, test_root)\nprint(X.shape, Xt.shape)\n"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"rng = np.random.default_rng(SEED)\nperm = rng.permutation(len(gold))\ncut = max(1, int(len(gold) * 0.8))\ntr_idx, va_idx = perm[:cut], perm[cut:]\nxx = torch.tensor(X)\nyy = torch.tensor(gold[labels].to_numpy(dtype='float32'))\n\nprobe = nn.Linear(X.shape[1], len(labels))\nopt = torch.optim.AdamW(probe.parameters(), lr=5e-4, weight_decay=5e-2)\nbest_loss, best_epoch, best_state, patience = float('inf'), 1, None, 0\nfor epoch in range(250):\n    probe.train()\n    loss = nn.functional.binary_cross_entropy_with_logits(probe(xx[tr_idx]), yy[tr_idx])\n    opt.zero_grad(); loss.backward(); opt.step()\n    probe.eval()\n    with torch.no_grad():\n        val_loss = nn.functional.binary_cross_entropy_with_logits(\n            probe(xx[va_idx]), yy[va_idx]).item()\n    if val_loss < best_loss - 1e-4:\n        best_loss, best_epoch = val_loss, epoch + 1\n        best_state = {k: v.detach().cpu().clone() for k, v in probe.state_dict().items()}\n        patience = 0\n    else:\n        patience += 1\n    if (epoch + 1) % 25 == 0:\n        print('epoch', epoch + 1, 'train', round(float(loss), 4),\n              'val', round(val_loss, 4), 'best_epoch', best_epoch)\n    if patience >= 35:\n        break\nprint('selected epochs:', best_epoch, 'validation loss:', round(best_loss, 4))\n\nhead = nn.Linear(X.shape[1], len(labels))\nopt = torch.optim.AdamW(head.parameters(), lr=5e-4, weight_decay=5e-2)\nfor _ in range(best_epoch):\n    loss = nn.functional.binary_cross_entropy_with_logits(head(xx), yy)\n    opt.zero_grad(); loss.backward(); opt.step()\nwith torch.no_grad():\n    preds = torch.sigmoid(head(torch.tensor(Xt))).numpy()\nsubmission = pd.read_csv(root.parent / 'sample_submission.csv')\nsubmission[labels] = preds\nsubmission.to_csv('/kaggle/working/submission.csv', index=False)\ndisplay(submission.head())\n"}],"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"}},"nbformat":4,"nbformat_minor":5}