{"nbformat":4,"nbformat_minor":5,"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.10.12"}},"cells":[{"cell_type":"code","execution_count":null,"metadata":{},"source":"# RSNA Knee Abnormality Detection: v31\nimport os\nimport sys\nimport glob\nimport gc\nimport time\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom pydicom.pixel_data_handlers.util import apply_voi_lut\nimport pydicom\nimport cv2\nfrom sklearn.metrics import roc_auc_score\n\nprint(f\"PyTorch: {torch.__version__}, CUDA: {torch.cuda.is_available()}\")\n\nBASE = '/kaggle/input/competitions/rsna-knee-abnormality-detection'\nTARGETS = ['ACL', 'MCL', 'Medial Meniscus', 'Lateral Meniscus', 'Medial OA',\n           'Lateral OA', 'PF OA', 'Effusion', 'Synovitis', \"Baker's\",\n           'Contusion', 'Fracture']\n\nclass CFG:\n    img_size = 256\n    num_slices = 3\n    num_planes = 2\n    batch_size = 8\n    epochs = 7\n    lr = 1e-4\n    weight_decay = 1e-3\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    num_classes = 12\n\ndef find_csv(name, root='/kaggle/input'):\n    for dirpath, _, files in os.walk(root):\n        if name in files:\n            return os.path.join(dirpath, name)\n    raise FileNotFoundError(name)\n\n# ------------- Model -------------\nclass KneeModel(nn.Module):\n    def __init__(self, pretrained=True):\n        super().__init__()\n        from torchvision.models import efficientnet_b0, EfficientNet_B0_Weights\n        if pretrained:\n            wpath = None\n            for rootdir in ['/kaggle/input', '/home/ubuntu']:\n                for dirpath, _, files in os.walk(rootdir):\n                    if any('efficientnet_b0' in f for f in files):\n                        wpath = os.path.join(dirpath, [f for f in files if 'efficientnet_b0' in f][0])\n                        break\n                if wpath:\n                    break\n            if wpath and os.path.exists(wpath):\n                print(f'[Model] loading weights from {wpath}')\n                base = efficientnet_b0(weights=None)\n                sd = torch.load(wpath, map_location='cpu')\n                base.load_state_dict(sd)\n            else:\n                try:\n                    weights = EfficientNet_B0_Weights.IMAGENET1K_V1\n                except Exception:\n                    weights = 'IMAGENET1K_V1'\n                base = efficientnet_b0(weights=weights)\n        else:\n            base = efficientnet_b0(weights=None)\n        base.classifier = nn.Identity()\n        self.backbone = base\n        self.att = nn.Sequential(nn.Linear(1280, 128), nn.ReLU())\n        self.gate = nn.Linear(128, 1)\n        self.head = nn.Sequential(\n            nn.Linear(1280, 512), nn.ReLU(), nn.Dropout(0.3),\n            nn.Linear(512, CFG.num_classes))\n\n    def forward(self, x):\n        x = x.float()\n        B = x.size(0)\n        x = x.view(B, -1, 3, x.size(-2), x.size(-1))\n        K = x.size(1)\n        s = x.view(B * K, 3, x.size(-2), x.size(-1))\n        feats = self.backbone(s).view(B, K, -1)\n        a = self.att(feats)\n        g = torch.softmax(self.gate(a), dim=1)\n        fused = (feats * g).sum(dim=1)\n        return self.head(fused)\n\n# ------------- Image loading + augmentation -------------\ndef prep_stack(files, max_slices=CFG.num_slices, img_size=CFG.img_size, augment=False, rng=None):\n    imgs = []\n    for f in files:\n        try:\n            dcm = pydicom.dcmread(f, stop_before_pixels=False, force=True)\n            px = dcm.pixel_array.astype(np.float32)\n            px = apply_voi_lut(px, dcm)\n            if px.max() > px.min():\n                px = (px - px.min()) / (px.max() - px.min())\n            else:\n                continue\n            px = np.clip(px, 0, 1)\n            if augment and rng is not None:\n                # brightness jitter\n                b = float(rng.uniform(0.85, 1.15))\n                px = np.clip(px * b, 0, 1)\n            px = cv2.resize(px, (img_size, img_size), interpolation=cv2.INTER_AREA)\n            if augment and rng is not None:\n                # rotation ±15°\n                ang = float(rng.uniform(-15, 15))\n                M = cv2.getRotationMatrix2D((img_size // 2, img_size // 2), ang, 1.0)\n                px = cv2.warpAffine(px, M, (img_size, img_size),\n                                    borderMode=cv2.BORDER_REFLECT)\n                # vertical flip (anatomically valid)\n                if rng.rand() < 0.5:\n                    px = px[::-1].copy()\n            imgs.append(px)\n        except Exception:\n            continue\n    if not imgs:\n        return None\n    n = len(imgs)\n    idx = np.linspace(0, n - 1, max_slices, dtype=int) if n >= max_slices else np.arange(n)\n    if augment and rng is not None:\n        idx = rng.choice(n, max_slices, replace=False) if n >= max_slices else idx\n    stack = np.stack([imgs[i] for i in idx], axis=0)\n    rgb = np.repeat(stack[:, None, :, :], 3, axis=1)\n    return torch.from_numpy(rgb.astype(np.float32))\n\n# ------------- Dataset -------------\nclass KneeDataset(Dataset):\n    def __init__(self, df, series_df, img_dir, label_cols=None, augment=False, seed=None):\n        self.ids = df['StudyInstanceUID'].astype(str).str.strip().values\n        self.img_dir = img_dir\n        self.augment = augment\n        self.rng = np.random.RandomState(seed)\n        if label_cols is not None:\n            self.labels = df[label_cols].fillna(0).values.astype(np.float32)\n        else:\n            self.labels = None\n        self.series_map = {}\n        self.fluid = {}\n        for _, r in series_df.iterrows():\n            s_id = str(r.iloc[0]).strip()\n            ser = str(r.iloc[1]).strip()\n            self.series_map.setdefault(s_id, [])\n            if ser not in self.series_map[s_id]:\n                self.series_map[s_id].append(ser)\n            self.fluid[ser] = int(r.iloc[2]) if 'Fluid_Sensitive' in series_df.columns else 0\n        print(f'[Dataset] {len(self.series_map)} studies mapped, len={len(self.ids)}')\n\n    def __len__(self):\n        return len(self.ids)\n\n    def __getitem__(self, i):\n        sid = str(self.ids[i]).strip()\n        series_list = self.series_map.get(sid, [])\n        fluid = [s for s in series_list if self.fluid.get(s, 0) == 1]\n        non_fluid = [s for s in series_list if self.fluid.get(s, 0) == 0]\n        selected = []\n        for pool in [fluid, non_fluid]:\n            for s in pool:\n                if len(selected) >= CFG.num_planes:\n                    break\n                selected.append(s)\n            if len(selected) >= CFG.num_planes:\n                break\n        if not selected:\n            selected = series_list[:CFG.num_planes]\n        if self.augment and self.rng.rand() < 0.3:\n            self.rng.shuffle(selected)\n        stacks = []\n        zero = torch.zeros((CFG.num_slices, 3, CFG.img_size, CFG.img_size))\n        for s in selected:\n            path = os.path.join(self.img_dir, sid, s)\n            if not os.path.exists(path):\n                stacks.append(zero)\n                continue\n            files = sorted(glob.glob(os.path.join(path, '*.dcm')))\n            if not files:\n                stacks.append(zero)\n                continue\n            t = prep_stack(files, augment=self.augment, rng=self.rng)\n            stacks.append(t if t is not None else zero)\n        while len(stacks) < CFG.num_planes:\n            stacks.append(zero)\n        x = torch.stack(stacks[:CFG.num_planes], dim=0)\n        x = x.view(-1, 3, CFG.img_size, CFG.img_size)\n        if self.labels is not None:\n            return x, torch.tensor(self.labels[i], dtype=torch.float32)\n        return x\n\n# ------------- Loss + training -------------\ndef train_epoch(model, loader, optimizer, scheduler):\n    model.train()\n    total_loss = n = 0\n    nan_c = 0\n    for x, lab in loader:\n        x, lab = x.to(CFG.device), lab.to(CFG.device)\n        out = model(x)\n        loss = F.binary_cross_entropy_with_logits(out.clamp(-10, 10), lab)\n        if torch.isnan(loss) or torch.isinf(loss):\n            nan_c += 1\n            continue\n        optimizer.zero_grad()\n        loss.backward()\n        torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n        optimizer.step()\n        total_loss += loss.item() * x.size(0)\n        n += x.size(0)\n    scheduler.step()\n    if nan_c:\n        print(f'[WARN] {nan_c} NaN batches')\n    return total_loss / max(n, 1)\n\ndef evaluate(model, loader):\n    model.eval()\n    preds, labs = [], []\n    with torch.no_grad():\n        for x, lab in loader:\n            x, lab = x.to(CFG.device), lab.to(CFG.device)\n            out = torch.sigmoid(model(x))\n            preds.append(out.cpu().numpy())\n            labs.append(lab.cpu().numpy())\n    preds = np.vstack(preds); labs = np.vstack(labs)\n    aucs = []\n    for i in range(CFG.num_classes):\n        try:\n            aucs.append(roc_auc_score(labs[:, i], preds[:, i]))\n        except Exception:\n            aucs.append(0.5)\n    return float(np.mean(aucs)), aucs\n\n# ------------- Main -------------\ndef main():\n    t0 = time.time()\n    print('[Data] train_series dir:', os.path.exists(os.path.join(BASE, 'train_series')))\n\n    train_series = pd.read_csv(os.path.join(BASE, 'train_series.csv'))\n    print(f'[Data] train_series rows: {len(train_series)}')\n\n    # weak labels v6\n    weak_path = find_csv('train_weak_labels_full.csv')\n    print('[Data] weak labels path:', weak_path)\n    weak_df = pd.read_csv(weak_path)\n    weak_df.columns = [c.strip() for c in weak_df.columns]\n    if 'StudyInstanceUID' not in weak_df.columns:\n        weak_df.columns = ['StudyInstanceUID'] + [c for c in weak_df.columns if c != 'StudyInstanceUID']\n    weak_df['_uid'] = weak_df['StudyInstanceUID'].astype(str).str.strip()\n    # sanity: check targets exist\n    for t in TARGETS:\n        if t not in weak_df.columns:\n            print(f'[WARN] missing target {t}; filling 0')\n            weak_df[t] = 0\n    print(f'[Data] weak labels: {len(weak_df)} rows, mean positives:')\n    print((weak_df[TARGETS].sum() / len(weak_df)).round(3).to_string())\n\n    val_size = int(len(weak_df) * 0.2)\n    rng = np.random.RandomState(42)\n    val_idx = set(rng.choice(len(weak_df), val_size, replace=False))\n    tr_df = weak_df.iloc[[i for i in range(len(weak_df)) if i not in val_idx]].reset_index(drop=True)\n    va_df = weak_df.iloc[[i for i in range(len(weak_df)) if i in val_idx]].reset_index(drop=True)\n    print(f'[Data] train {len(tr_df)} / val {len(va_df)}')\n\n    img_dir = os.path.join(BASE, 'train_series')\n    tr_ds = KneeDataset(tr_df, train_series, img_dir, label_cols=TARGETS, augment=True, seed=42)\n    va_ds = KneeDataset(va_df, train_series, img_dir, label_cols=TARGETS, augment=False, seed=42)\n\n    s0 = tr_ds[0]\n    print(f'[DEBUG] sample0 mean={s0[0].mean():.4f} max={s0[0].max():.4f} (must be >0)')\n    sys.stdout.flush()\n    tr_loader = DataLoader(tr_ds, batch_size=CFG.batch_size, shuffle=True, num_workers=0, drop_last=True)\n    va_loader = DataLoader(va_ds, batch_size=CFG.batch_size, shuffle=False, num_workers=0)\n\n    print('[Model] building EfficientNet-B0...')\n    model = KneeModel(pretrained=True).to(CFG.device)\n    n_p = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    print(f'[Model] {n_p/1e6:.1f}M params on {CFG.device}')\n    optimizer = torch.optim.AdamW(model.parameters(), lr=CFG.lr, weight_decay=CFG.weight_decay)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=CFG.epochs)\n\n    best_auc = 0.0\n    print('[Train] starting...')\n    sys.stdout.flush()\n    for ep in range(CFG.epochs):\n        t1 = time.time()\n        loss = train_epoch(model, tr_loader, optimizer, scheduler)\n        val_auc, aucs = evaluate(model, va_loader)\n        print(f'[Epoch {ep+1}/{CFG.epochs}] loss={loss:.4f} val_auc={val_auc:.4f} time={time.time()-t1:.0f}s')\n        for i, a in enumerate(aucs):\n            print(f'  {TARGETS[i]}: {a:.4f}', end='  ' if (i+1) % 3 else '\\n')\n        print()\n        sys.stdout.flush()\n        if val_auc > best_auc:\n            best_auc = val_auc\n            torch.save(model.state_dict(), '/kaggle/working/best_model.pth')\n            print(f'[BEST] val_auc={best_auc:.4f} saved')\n        if time.time() - t0 > 6.5 * 3600:\n            print('[TIME GUARD] stopping training')\n            break\n        gc.collect()\n\n    print(f'[DONE] best val AUC = {best_auc:.4f} in {time.time()-t0:.0f}s')\n    sys.stdout.flush()\n\n    # inference on ALL test studies (test.csv may have more than 3 now)\n    test_csv = os.path.join(BASE, 'test.csv')\n    if os.path.exists(test_csv):\n        test_df = pd.read_csv(test_csv)\n        print(f'[Inference] {len(test_df)} test studies')\n        test_series = pd.read_csv(find_csv('test_series.csv'))\n        test_ds = KneeDataset(test_df, test_series, img_dir, augment=False)\n        model.load_state_dict(torch.load('/kaggle/working/best_model.pth', map_location=CFG.device))\n        preds = []\n        with torch.no_grad():\n            for x in DataLoader(test_ds, batch_size=4, shuffle=False, num_workers=0):\n                x = x.to(CFG.device)\n                preds.append(torch.sigmoid(model(x)).cpu().numpy())\n        preds = np.vstack(preds)\n        sub = pd.DataFrame(preds, columns=TARGETS)\n        sub.insert(0, 'StudyInstanceUID', test_df['StudyInstanceUID'].astype(str).str.strip().values[:len(sub)])\n        sub.to_csv('/kaggle/working/submission.csv', index=False)\n        print('[Saved] /kaggle/working/submission.csv')\n        print(sub.to_string())\n        sys.stdout.flush()\n    else:\n        print('[Inference] no test.csv found — submission.csv NOT produced')\n\nif __name__ == '__main__':\n    main()\n","outputs":[]}]}