{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[],"dockerImageVersionId":28755,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import re\nimport warnings\nfrom pathlib import Path\nfrom typing import Dict, List, Optional, Tuple, Any\n\nimport numpy as np\nimport pydicom\nimport torch\nimport torch.nn.functional as F\n\nwarnings.filterwarnings(\"ignore\", category=UserWarning)\n\nTARGET_COLS = [\n    \"ACL\", \"MCL\", \"Medial Meniscus\", \"Lateral Meniscus\",\n    \"Medial OA\", \"Lateral OA\", \"PF OA\", \"Effusion\",\n    \"Synovitis\", \"Baker's\", \"Contusion\", \"Fracture\"\n]\n\nSLOTS = [\n    ('SAG_FLUID_FS', 'Sagittal', True, True),\n    ('COR_FLUID_FS', 'Coronal', True, True),\n    ('AX_FLUID_FS',  'Axial', True, True),\n    ('SAG_FLUID_NOFS', 'Sagittal', True, False),\n    ('COR_T1', 'Coronal', False, False),\n    ('SAG_T1', 'Sagittal', False, False),\n]\n\nFATSAT_OPTS = {'FS', 'FATSAT', 'FAT_SAT', 'FSAT'}\n_SEP = re.compile(r'[_\\-.]')\n_FATSAT_RX = re.compile(r'\\bfs\\b|fatsat|fat sat|\\bstir\\b|\\bspair\\b|\\bspir\\b|\\bwe\\b|water excit|\\btirm\\b|\\bsting\\b|\\bfatsup\\b')\n_T1_RX = re.compile(r'\\bt1\\b|\\bt1w\\b')\n_T2_RX = re.compile(r'\\bt2\\b|\\bt2w\\b')\n_PD_RX = re.compile(r'\\bpd\\b|\\bpdw\\b|proton|\\bdp\\b|dens')\n\n\nclass KneeImageProcessor:\n    def __init__(\n        self,\n        img_size: int = 224,\n        crop_mm: float = 160.0,\n        n_best_slices: int = 12\n    ):\n        self.img_size = img_size\n        self.crop_mm = crop_mm\n        self.n_best_slices = n_best_slices\n\n    def parse_dicom_meta(self, dcm_path: Path) -> Dict[str, Any]:\n        try:\n            ds = pydicom.dcmread(dcm_path, stop_before_pixels=True, force=True)\n            \n            series_desc = str(getattr(ds, 'SeriesDescription', '') or '')\n            seq_name = str(getattr(ds, 'SequenceName', '') or '')\n            scan_opts = str(getattr(ds, 'ScanOptions', '') or '')\n            scan_seq = str(getattr(ds, 'ScanningSequence', '') or '')\n            \n            tr = getattr(ds, 'RepetitionTime', None)\n            te = getattr(ds, 'EchoTime', None)\n            laterality = str(getattr(ds, 'Laterality', '') or '').strip().upper()\n            \n            slice_thickness = getattr(ds, 'SliceThickness', None)\n            flip_angle = getattr(ds, 'FlipAngle', None)\n            mfg_raw = str(getattr(ds, 'Manufacturer', '') or '').upper()\n            \n            if 'SIEMENS' in mfg_raw:\n                manufacturer_code = 1.0\n            elif 'GE' in mfg_raw or 'GENERAL ELECTRIC' in mfg_raw:\n                manufacturer_code = 0.5\n            elif 'PHILIPS' in mfg_raw:\n                manufacturer_code = 0.0\n            else:\n                manufacturer_code = -0.5\n            \n            sex = str(getattr(ds, 'PatientSex', 'O') or 'O').strip().upper()\n            age_raw = str(getattr(ds, 'PatientAge', '') or '')\n            field_strength = getattr(ds, 'MagneticFieldStrength', 1.5)\n            \n            age = 50.0\n            if age_raw:\n                match = re.match(r'(\\d+)([YMWD])', age_raw)\n                if match:\n                    val, unit = match.groups()\n                    val = float(val)\n                    if unit == 'Y': age = val\n                    elif unit == 'M': age = val / 12.0\n                    elif unit == 'W': age = val / 52.0\n                    elif unit == 'D': age = val / 365.0\n            \n            pixel_spacing = getattr(ds, 'PixelSpacing', None)\n            px = float(pixel_spacing[0]) if pixel_spacing else 0.45 \n            \n            return {\n                'SeriesDescription': series_desc,\n                'SequenceName': seq_name,\n                'ScanOptions': scan_opts,\n                'ScanningSequence': scan_seq,\n                'RepetitionTime': float(tr) if tr is not None else None,\n                'EchoTime': float(te) if te is not None else None,\n                'Laterality': laterality[0] if laterality and laterality[0] in ('L', 'R') else None,\n                'PixelSpacing': px,\n                'PatientSex': 1.0 if sex == 'M' else (0.0 if sex == 'F' else 0.5),\n                'PatientAge': age,\n                'MagneticFieldStrength': float(field_strength) if field_strength is not None else 1.5,\n                'SliceThickness': float(slice_thickness) if slice_thickness is not None else None,\n                'FlipAngle': float(flip_angle) if flip_angle is not None else None,\n                'Manufacturer': manufacturer_code\n            }\n        except Exception:\n            return {}\n\n    def classify_series(self, meta: Dict[str, Any], plane: str) -> Tuple[bool, bool]:\n        desc = (meta.get('SeriesDescription', '') + ' ' + meta.get('SequenceName', '')).lower()\n        desc = re.sub(_SEP, ' ', desc)\n        \n        opts_fs = any(t.strip() in FATSAT_OPTS for t in meta.get('ScanOptions', '').upper().split('|'))\n        fatsat = bool(re.search(_FATSAT_RX, desc) or opts_fs)\n        \n        tr = meta.get('RepetitionTime')\n        te = meta.get('EchoTime')\n        gre = 'GR' in meta.get('ScanningSequence', '').upper()\n        \n        t1 = bool(re.search(_T1_RX, desc))\n        t2 = bool(re.search(_T2_RX, desc))\n        pdw = bool(re.search(_PD_RX, desc))\n        \n        if t1 and not t2 and not pdw:\n            weight = 'T1'\n        elif t2 and not pdw:\n            weight = 'T2'\n        elif pdw:\n            weight = 'PD'\n        elif gre:\n            weight = 'GRE'\n        elif tr is not None and tr < 800:\n            weight = 'T1'\n        elif te is not None and te > 60:\n            weight = 'T2'\n        elif tr is not None and tr >= 800:\n            weight = 'PD'\n        else:\n            weight = 'UNK'\n            \n        fluid = weight in ['PD', 'T2']\n        return fluid, fatsat\n\n    def get_physically_sorted_slices(self, series_dir: Path, plane: str) -> List[Tuple[float, Path]]:\n        files = list(series_dir.glob(\"*.dcm\"))\n        if not files:\n            return []\n            \n        slice_positions = []\n        for f in files:\n            try:\n                ds = pydicom.dcmread(f, stop_before_pixels=True, force=True)\n                ipp = getattr(ds, 'ImagePositionPatient', None)\n                if ipp is not None:\n                    if plane == 'Sagittal': pos = float(ipp[0])\n                    elif plane == 'Coronal': pos = float(ipp[1])\n                    else: pos = float(ipp[2])\n                else:\n                    pos = float(getattr(ds, 'SliceLocation', 0.0))\n            except Exception:\n                num_match = re.findall(r'\\d+', f.name)\n                pos = float(num_match[-1]) if num_match else 0.0\n            slice_positions.append((pos, f))\n            \n        slice_positions.sort(key=lambda x: x[0])\n        return slice_positions\n\n    def process_volume(\n        self, \n        sorted_slices: List[Tuple[float, Path]], \n        plane: str, \n        laterality: Optional[str], \n        pixel_spacing: float\n    ) -> np.ndarray:\n        n = len(sorted_slices)\n        if n == 0:\n            return np.zeros((self.n_best_slices, self.img_size, self.img_size), dtype=np.uint8)\n\n        lo, hi = int(0.10 * (n - 1)), int(0.90 * (n - 1))\n        if hi > lo:\n            idx = np.unique(np.linspace(lo, hi, self.n_best_slices).astype(int))\n        else:\n            idx = np.array([n // 2])\n            \n        while len(idx) < self.n_best_slices:\n            idx = np.append(idx, idx[-1])\n            \n        planes = []\n        for i in idx[:self.n_best_slices]:\n            _, fpath = sorted_slices[int(i)]\n            try:\n                ds = pydicom.dcmread(fpath, force=True)\n                a = ds.pixel_array.astype(np.float32)\n                \n                sl = float(getattr(ds, 'RescaleSlope', 1.0) or 1.0)\n                ic = float(getattr(ds, 'RescaleIntercept', 0.0) or 0.0)\n                a = a * sl + ic\n                \n                if str(getattr(ds, 'PhotometricInterpretation', '')).upper() == 'MONOCHROME1':\n                    a = a.max() - a\n            except Exception:\n                a = np.zeros((self.img_size, self.img_size), dtype=np.float32)\n            planes.append(a)\n\n        shp = planes[0].shape\n        planes = [p if p.shape == shp else np.zeros(shp, np.float32) for p in planes]\n        vol = np.stack(planes)\n\n        if pixel_spacing > 0:\n            want = int(round(self.crop_mm / pixel_spacing))\n            h, w = shp\n            if 16 < want < min(h, w):\n                cy, cx = h // 2, w // 2\n                half = want // 2\n                vol = vol[:, max(0, cy - half):cy + half, max(0, cx - half):cx + half]\n\n        lo_v, hi_v = np.percentile(vol, [1, 99])\n        vol = np.clip((vol - lo_v) / max(hi_v - lo_v, 1e-6), 0.0, 1.0)\n\n        t = torch.from_numpy(np.ascontiguousarray(vol)).unsqueeze(0)\n        t = F.interpolate(t, size=(self.img_size, self.img_size), mode='bilinear', align_corners=False).squeeze(0)\n\n        return (t * 255.0).round().clamp(0, 255).to(torch.uint8).numpy()\n\n    def _compile_meta_features(self, parsed_series: List[Dict[str, Any]], study_mask: np.ndarray) -> np.ndarray:\n        if not parsed_series:\n            return np.array([0.5, 0.5, 0.5, 0.6, 0.45, -0.5, 0.25, 0.5, 0.5, 0.0], dtype=np.float32)\n\n        patient_sex = parsed_series[0]['PatientSex']\n        patient_age = parsed_series[0]['PatientAge']\n        field_strength = parsed_series[0]['MagneticFieldStrength']\n\n        slice_thicknesses = [s['SliceThickness'] for s in parsed_series if s.get('SliceThickness') is not None]\n        pixel_spacings = [s['PixelSpacing'] for s in parsed_series if s.get('PixelSpacing') is not None]\n        manufacturers = [s['Manufacturer'] for s in parsed_series if s.get('Manufacturer') is not None]\n        tes = [s['EchoTime'] for s in parsed_series if s.get('EchoTime') is not None]\n        trs = [s['RepetitionTime'] for s in parsed_series if s.get('RepetitionTime') is not None]\n        flip_angles = [s['FlipAngle'] for s in parsed_series if s.get('FlipAngle') is not None]\n\n        mean_thickness = np.mean(slice_thicknesses) if slice_thicknesses else 3.0\n        mean_spacing = np.mean(pixel_spacings) if pixel_spacings else 0.45\n        mfg_code = max(set(manufacturers), key=manufacturers.count) if manufacturers else -0.5\n        mean_te = np.mean(tes) if tes else 30.0\n        mean_tr = np.mean(trs) if trs else 2000.0\n        mean_flip = np.mean(flip_angles) if flip_angles else 90.0\n\n        slot_completion_ratio = float(np.sum(study_mask) / len(study_mask))\n\n        meta_arr = np.array([\n            patient_sex,                                         # [0] 性别 (M=1.0, F=0.0, O=0.5)\n            np.clip(patient_age / 100.0, 0.0, 1.0),              # [1] 年龄 (归一化到 0-1)\n            np.clip(field_strength / 3.0, 0.0, 1.5),             # [2] 场强 (0.5=1.5T, 1.0=3T)\n            np.clip(mean_thickness / 5.0, 0.0, 2.0),             # [3] 层厚 (典型值 3-4mm -> 0.6-0.8)\n            np.clip(mean_spacing / 1.0, 0.0, 2.0),               # [4] PixelSpacing (典型值 0.45)\n            mfg_code,                                            # [5] 扫描仪品牌编码 (-0.5~1.0)\n            np.clip(mean_te / 120.0, 0.0, 2.0),                  # [6] EchoTime\n            np.clip(mean_tr / 4000.0, 0.0, 2.0),                 # [7] RepetitionTime\n            np.clip(mean_flip / 180.0, 0.0, 1.0),                # [8] FlipAngle 翻转角\n            slot_completion_ratio                                # [9] 有效插槽比例 (0.0~1.0)\n        ], dtype=np.float32)\n\n        return meta_arr\n\n    def process_study(self, study_dir: Path) -> Tuple[np.ndarray, np.ndarray, np.ndarray]:\n        study_data = np.zeros((len(SLOTS), self.n_best_slices, self.img_size, self.img_size), dtype=np.uint8)\n        study_mask = np.zeros(len(SLOTS), dtype=np.float32)\n        \n        series_dirs = [p for p in study_dir.iterdir() if p.is_dir()]\n        \n        parsed_series = []\n        for s_dir in series_dirs:\n            dcms = list(s_dir.glob(\"*.dcm\"))\n            if not dcms:\n                continue\n            mid_dcm = dcms[len(dcms) // 2]\n            meta = self.parse_dicom_meta(mid_dcm)\n            if not meta:\n                continue\n            \n            plane = self._infer_anatomical_plane(mid_dcm)\n            fluid, fatsat = self.classify_series(meta, plane)\n            \n            meta.update({\n                'plane': plane,\n                'fluid': fluid,\n                'fatsat': fatsat,\n                'dir': s_dir,\n                'n_slices': len(dcms)\n            })\n            parsed_series.append(meta)\n            \n        laterality = None\n        if parsed_series:\n            lats = [s['Laterality'] for s in parsed_series if s['Laterality']]\n            if lats:\n                laterality = max(set(lats), key=lats.count)\n\n        for k, (slot_name, plane, fluid, fs) in enumerate(SLOTS):\n            candidates = [\n                s for s in parsed_series \n                if s['plane'] == plane and s['fatsat'] == fs and s['fluid'] == fluid\n            ]\n            if not candidates and not fluid:\n                candidates = [s for s in parsed_series if s['plane'] == plane and not s['fatsat']]\n                \n            if candidates:\n                best_s = sorted(candidates, key=lambda x: x['n_slices'], reverse=True)[0]\n                sorted_slices = self.get_physically_sorted_slices(best_s['dir'], plane)\n                vol = self.process_volume(sorted_slices, plane, laterality, best_s['PixelSpacing'])\n                \n                study_data[k] = vol\n                study_mask[k] = 1.0\n\n        meta_features = self._compile_meta_features(parsed_series, study_mask)\n        \n        return study_data, study_mask, meta_features\n\n    def _infer_anatomical_plane(self, dcm_path: Path) -> str:\n        try:\n            ds = pydicom.dcmread(dcm_path, stop_before_pixels=True, force=True)\n            iop = ds.ImageOrientationPatient\n            if iop is None or len(iop) < 6:\n                return 'Sagittal'\n                \n            row = np.array(iop[0:3])\n            col = np.array(iop[3:6])\n            normal = np.cross(row, col)\n            \n            abs_normal = np.abs(normal)\n            max_idx = np.argmax(abs_normal)\n            \n            if max_idx == 0: return 'Sagittal'\n            elif max_idx == 1: return 'Coronal'\n            else: return 'Axial'\n        except Exception:\n            return 'Sagittal'","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport gc\nimport re\nimport time\nimport warnings\nfrom pathlib import Path\nfrom concurrent.futures import ThreadPoolExecutor\nimport numpy as np\nimport pandas as pd\nimport pydicom\nimport torch\nimport torch.nn.functional as F\nfrom tqdm import tqdm\n\nwarnings.filterwarnings(\"ignore\", category=UserWarning)\n\nIMG_SIZE = 224\nCROP_MM = 160.0\nN_BEST_SLICES = 12\nSTUDIES_PER_CHUNK = 400\nSTART_CHUNK = 0\n\nSLOTS = [\n    ('SAG_FLUID_FS', 'Sagittal', True, True),\n    ('COR_FLUID_FS', 'Coronal', True, True),\n    ('AX_FLUID_FS',  'Axial', True, True),\n    ('SAG_FLUID_NOFS', 'Sagittal', True, False),\n    ('COR_T1', 'Coronal', False, False),\n    ('SAG_T1', 'Sagittal', False, False),\n]\n\nFATSAT_OPTS = {'FS', 'FATSAT', 'FAT_SAT', 'FSAT'}\n_SEP = re.compile(r'[_\\-.]')\n_FATSAT_RX = re.compile(r'\\bfs\\b|fatsat|fat sat|\\bstir\\b|\\bspair\\b|\\bspir\\b|\\bwe\\b|water excit|\\btirm\\b|\\bsting\\b|\\bfatsup\\b')\n_T1_RX = re.compile(r'\\bt1\\b|\\bt1w\\b')\n_T2_RX = re.compile(r'\\bt2\\b|\\bt2w\\b')\n_PD_RX = re.compile(r'\\bpd\\b|\\bpdw\\b|proton|\\bdp\\b|dens')\n\ndef find_root():\n    for c in [Path('/kaggle/input/competitions/rsna-knee-abnormality-detection'),\n              Path('/kaggle/input/rsna-knee-abnormality-detection')]:\n        if (c / 'train.csv').is_file(): return c\n    raise FileNotFoundError('RSNA Knee Dataset not found!')\n\nROOT = find_root()\nOUTPUT_DIR = Path('/kaggle/working/rsna_knee_processed')\nOUTPUT_DIR.mkdir(parents=True, exist_ok=True)\nN_WORKERS = os.cpu_count() or 4\n\nprint(f\"[Init] Input root: {ROOT}\")\nprint(f\"[Init] Output directory: {OUTPUT_DIR}\")\nprint(f\"[Init] CPU cores available: {N_WORKERS}\")\n\ndef parse_age(age_str: str) -> float:\n    if not age_str:\n        return 50.0\n    match = re.match(r'(\\d+)([YMWD])', age_str.strip().upper())\n    if match:\n        val, unit = match.groups()\n        val = float(val)\n        if unit == 'Y': return val\n        elif unit == 'M': return val / 12.0\n        elif unit == 'W': return val / 52.0\n        elif unit == 'D': return val / 365.0\n    return 50.0\n\ndef parse_single_series_meta(item) -> dict:\n    study_uid, series_uid, path = item\n    row = {\n        'StudyInstanceUID': study_uid, \n        'SeriesInstanceUID': series_uid, \n        'dir': str(path),\n        'PixelSpacing': 0.45,\n        'PatientSex': 0.5,\n        'PatientAge': 50.0,\n        'MagneticFieldStrength': 1.5,\n        'Laterality': None,\n        'SliceThickness': None,\n        'FlipAngle': None,\n        'Manufacturer': -0.5\n    }\n    try:\n        files = sorted([e.name for e in os.scandir(path) if e.name.endswith('.dcm')])\n        row['n_slices'] = len(files)\n        row['files'] = files\n        if files:\n            mid_file = os.path.join(path, files[len(files) // 2])\n            ds = pydicom.dcmread(mid_file, stop_before_pixels=True, force=True)\n            \n            row['SeriesDescription'] = str(getattr(ds, 'SeriesDescription', '') or '')\n            row['SequenceName'] = str(getattr(ds, 'SequenceName', '') or '')\n            row['ScanOptions'] = str(getattr(ds, 'ScanOptions', '') or '')\n            row['ScanningSequence'] = str(getattr(ds, 'ScanningSequence', '') or '')\n            row['RepetitionTime'] = getattr(ds, 'RepetitionTime', None)\n            row['EchoTime'] = getattr(ds, 'EchoTime', None)\n            \n            lat = str(getattr(ds, 'Laterality', '') or '').strip().upper()\n            if lat and lat[0] in ('L', 'R'):\n                row['Laterality'] = lat[0]\n                \n            sex = str(getattr(ds, 'PatientSex', 'O') or 'O').strip().upper()\n            row['PatientSex'] = 1.0 if sex == 'M' else (0.0 if sex == 'F' else 0.5)\n            row['PatientAge'] = parse_age(str(getattr(ds, 'PatientAge', '')))\n            row['MagneticFieldStrength'] = float(getattr(ds, 'MagneticFieldStrength', 1.5) or 1.5)\n            \n            row['SliceThickness'] = getattr(ds, 'SliceThickness', None)\n            row['FlipAngle'] = getattr(ds, 'FlipAngle', None)\n            \n            mfg = str(getattr(ds, 'Manufacturer', '') or '').upper()\n            if 'SIEMENS' in mfg: row['Manufacturer'] = 1.0\n            elif 'GE' in mfg or 'GENERAL ELECTRIC' in mfg: row['Manufacturer'] = 0.5\n            elif 'PHILIPS' in mfg: row['Manufacturer'] = 0.0\n            else: row['Manufacturer'] = -0.5\n            \n            px_val = getattr(ds, 'PixelSpacing', None)\n            if px_val:\n                row['PixelSpacing'] = float(px_val[0])\n    except Exception as e:\n        row['err'] = str(e)[:50]\n        row['n_slices'] = 0\n    return row\n\ndef annotate_series(df: pd.DataFrame) -> pd.DataFrame:\n    desc = (df['SeriesDescription'].fillna('') + ' ' + df['SequenceName'].fillna('')).str.lower().str.replace(_SEP, ' ', regex=True)\n    opts_fs = df['ScanOptions'].fillna('').str.upper().str.split('|').apply(lambda ts: any(t.strip() in FATSAT_OPTS for t in ts))\n    df['fatsat'] = desc.str.contains(_FATSAT_RX) | opts_fs\n\n    tr = pd.to_numeric(df['RepetitionTime'], errors='coerce')\n    te = pd.to_numeric(df['EchoTime'], errors='coerce')\n    gre = df['ScanningSequence'].fillna('').str.upper().str.contains('GR')\n    t1 = desc.str.contains(_T1_RX)\n    t2 = desc.str.contains(_T2_RX)\n    pdw = desc.str.contains(_PD_RX)\n\n    df['weight'] = np.where(t1 & ~t2 & ~pdw, 'T1',\n                    np.where(t2 & ~pdw, 'T2',\n                      np.where(pdw, 'PD',\n                        np.where(gre, 'GRE',\n                          np.where(tr < 800, 'T1',\n                            np.where(te > 60, 'T2',\n                              np.where(tr >= 800, 'PD', 'UNK')))))))\n    df['fluid'] = np.isin(df['weight'], ['PD', 'T2'])\n    return df\n\ndef get_physically_sorted_slices(series_dir: Path, files: list, plane: str) -> list:\n    slice_positions = []\n    for f in files:\n        fpath = series_dir / f\n        try:\n            ds = pydicom.dcmread(fpath, stop_before_pixels=True, force=True)\n            ipp = getattr(ds, 'ImagePositionPatient', None)\n            if ipp is not None:\n                if plane == 'Sagittal': pos = float(ipp[0])\n                elif plane == 'Coronal': pos = float(ipp[1])\n                else: pos = float(ipp[2])\n            else:\n                pos = float(getattr(ds, 'SliceLocation', 0.0))\n        except Exception:\n            num_match = re.findall(r'\\d+', f)\n            pos = float(num_match[-1]) if num_match else 0.0\n        slice_positions.append((pos, fpath))\n    slice_positions.sort(key=lambda x: x[0])\n    return slice_positions\n\ndef process_slot_volume(rec: dict, plane: str, lat: str) -> np.ndarray:\n    if not rec or rec.get('n_slices', 0) == 0:\n        return None\n\n    files = rec['files']\n    s_dir = Path(rec['dir'])\n    n = len(files)\n    px = rec['PixelSpacing']\n\n    lo, hi = int(0.10 * (n - 1)), int(0.90 * (n - 1))\n    if hi > lo:\n        idx = np.unique(np.linspace(lo, hi, N_BEST_SLICES).astype(int))\n    else:\n        idx = np.array([n // 2])\n    while len(idx) < N_BEST_SLICES:\n        idx = np.append(idx, idx[-1])\n\n    sorted_slices = get_physically_sorted_slices(s_dir, files, plane)\n    \n    planes = []\n    for i in idx[:N_BEST_SLICES]:\n        try:\n            _, fpath = sorted_slices[int(i)]\n            ds = pydicom.dcmread(fpath, force=True)\n            a = ds.pixel_array.astype(np.float32)\n            sl = float(getattr(ds, 'RescaleSlope', 1) or 1)\n            ic = float(getattr(ds, 'RescaleIntercept', 0) or 0)\n            a = a * sl + ic\n            if str(getattr(ds, 'PhotometricInterpretation', '')).upper() == 'MONOCHROME1':\n                a = a.max() - a\n        except Exception:\n            a = np.zeros((IMG_SIZE, IMG_SIZE), dtype=np.float32)\n        planes.append(a)\n\n    shp = planes[0].shape\n    planes = [p if p.shape == shp else np.zeros(shp, np.float32) for p in planes]\n    vol = np.stack(planes)\n\n    if px > 0:\n        want = int(round(CROP_MM / px))\n        h, w = shp\n        if 16 < want < min(h, w):\n            cy, cx = h // 2, w // 2\n            half = want // 2\n            vol = vol[:, max(0, cy - half):cy + half, max(0, cx - half):cx + half]\n\n    lo_v, hi_v = np.percentile(vol, [1, 99])\n    vol = np.clip((vol - lo_v) / max(hi_v - lo_v, 1e-6), 0.0, 1.0)\n\n    t = torch.from_numpy(np.ascontiguousarray(vol)).unsqueeze(0)\n    t = F.interpolate(t, size=(IMG_SIZE, IMG_SIZE), mode='bilinear', align_corners=False).squeeze(0)\n\n    return (t * 255.0).round().clamp(0, 255).to(torch.uint8).numpy()\n\ndef process_single_study(args) -> tuple:\n    st_uid, study_slots, lat_val, meta_feats = args\n    study_data = np.zeros((len(SLOTS), N_BEST_SLICES, IMG_SIZE, IMG_SIZE), dtype=np.uint8)\n    study_mask = np.zeros(len(SLOTS), dtype=np.float32)\n\n    for k, (slot_name, plane, _, _) in enumerate(SLOTS):\n        if slot_name in study_slots:\n            vol = process_slot_volume(study_slots[slot_name], plane, lat_val)\n            if vol is not None:\n                study_data[k] = vol\n                study_mask[k] = 1.0\n\n    return st_uid, study_data, study_mask, meta_feats\n\ndef main():\n    t0 = time.time()\n    \n    print(\"\\n=== [Step 1] Scanning DICOM structure ===\")\n    train_series_csv = pd.read_csv(ROOT / 'train_series.csv')\n    plane_map = dict(zip(train_series_csv['SeriesInstanceUID'], train_series_csv['Anatomical_Plane']))\n\n    items = []\n    base_dir = ROOT / 'train_series'\n    for study in os.scandir(base_dir):\n        if study.is_dir():\n            for series in os.scandir(study.path):\n                if series.is_dir():\n                    items.append((study.name, series.name, series.path))\n\n    print(\"\\n=== [Step 2] Reading raw DICOM headers ===\")\n    with ThreadPoolExecutor(max_workers=N_WORKERS * 2) as pool:\n        rows = list(tqdm(pool.map(parse_single_series_meta, items), total=len(items), desc='Parsing Headers'))\n\n    df = pd.DataFrame(rows)\n    df = df[df['n_slices'] > 0].reset_index(drop=True)\n    df['plane'] = df['SeriesInstanceUID'].map(plane_map)\n    df = annotate_series(df)\n\n    lat_map = {}\n    for st, g in df.groupby('StudyInstanceUID'):\n        v = [x for x in g['Laterality'].dropna()]\n        lat_map[st] = max(set(v), key=v.count) if v else None\n\n    print(\"\\n=== [Step 3] Matching series to clinic slots ===\")\n    matched_slots = {}\n    study_metas = {}\n    \n    temp_processor = KneeImageProcessor(img_size=IMG_SIZE, crop_mm=CROP_MM, n_best_slices=N_BEST_SLICES)\n\n    for study, g in tqdm(df.groupby('StudyInstanceUID'), desc=\"Matching Slots\"):\n        chosen = {}\n        parsed_list = []\n        for name, plane, fluid, fs in SLOTS:\n            sel = (g['plane'] == plane) & (g['fatsat'] == fs) & (g['fluid'] == fluid)\n            cand = g[sel]\n            if len(cand) == 0 and not fluid:\n                cand = g[(g['plane'] == plane) & (~g['fatsat'])]\n            if len(cand) > 0:\n                rec_dict = cand.sort_values('n_slices', ascending=False).iloc[0].to_dict()\n                chosen[name] = rec_dict\n                parsed_list.append(rec_dict)\n        matched_slots[study] = chosen\n        \n        study_mask = np.zeros(len(SLOTS), dtype=np.float32)\n        for k, (slot_name, _, _, _) in enumerate(SLOTS):\n            if slot_name in chosen:\n                study_mask[k] = 1.0\n\n        study_metas[study] = temp_processor._compile_meta_features(parsed_list, study_mask)\n\n    print(\"\\n=== [Step 4] Preprocessing & Saving Chunks ===\")\n    studies = sorted(matched_slots.keys())\n    num_chunks = int(np.ceil(len(studies) / STUDIES_PER_CHUNK))\n    manifest = []\n\n    for chunk_idx in range(num_chunks):\n        c_start = chunk_idx * STUDIES_PER_CHUNK\n        c_end = min((chunk_idx + 1) * STUDIES_PER_CHUNK, len(studies))\n        chunk_studies = studies[c_start:c_end]\n        chunk_file_name = f'train_chunk_{chunk_idx:02d}.npz'\n\n        if chunk_idx < START_CHUNK:\n            print(f'[{time.time()-t0:.1f}s] [Skipped] Chunk {chunk_idx:02d} ({chunk_file_name})')\n            for uid in chunk_studies:\n                manifest.append({'StudyInstanceUID': uid, 'chunk_file': chunk_file_name})\n            continue\n\n        job_args = [\n            (uid, matched_slots[uid], lat_map.get(uid), study_metas[uid]) \n            for uid in chunk_studies\n        ]\n        \n        chunk_uids, chunk_imgs, chunk_masks, chunk_metas = [], [], [], []\n\n        print(f'\\nProcessing Chunk {chunk_idx:02d}/{num_chunks-1} (Studies {c_start} ~ {c_end - 1})...')\n        with ThreadPoolExecutor(max_workers=N_WORKERS) as pool:\n            results = list(tqdm(pool.map(process_single_study, job_args), total=len(job_args), desc=f'Chunk {chunk_idx:02d}'))\n\n        for uid, imgs, mask, meta in results:\n            chunk_uids.append(uid)\n            chunk_imgs.append(imgs)\n            chunk_masks.append(mask)\n            chunk_metas.append(meta)\n            manifest.append({'StudyInstanceUID': uid, 'chunk_file': chunk_file_name})\n\n        chunk_file_path = OUTPUT_DIR / chunk_file_name\n        np.savez_compressed(\n            chunk_file_path,\n            uids=np.array(chunk_uids),\n            images=np.stack(chunk_imgs),\n            masks=np.stack(chunk_masks),\n            metas=np.stack(chunk_metas)\n        )\n        size_mb = chunk_file_path.stat().st_size / (1024 * 1024)\n        print(f'[{time.time()-t0:.1f}s] Saved: {chunk_file_name} ({size_mb:.1f} MB)')\n        \n        gc.collect()\n\n    manifest_df = pd.DataFrame(manifest)\n    manifest_df.to_csv(OUTPUT_DIR / 'manifest.csv', index=False)\n    print(f'\\n=== Success ===\\nManifest saved to manifest.csv. Total studies: {len(manifest_df)}')\n    print(f'Total time elapsed: {time.time()-t0:.1f} seconds.')\n\nif __name__ == '__main__':\n    main()","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}