| import json |
| from pathlib import Path |
| from typing import Any, Callable, Optional, Tuple |
|
|
| import numpy as np |
| import pydicom |
| import torch |
| from torch import Tensor |
| from torch.utils.data import Dataset |
|
|
|
|
| DTYPE = torch.float16 |
|
|
|
|
| class SyntaxDataset(Dataset): |
| """ |
| Dataset for the RNN/LSTM head on top of the backbone. |
| |
| JSON structure: |
| [ |
| { |
| "study_uid": "...", |
| "syntax_left": 12.5, |
| "syntax_right": 8.2, |
| "videos_left": [ |
| {"path": "../data/anon_data/.../IM-0001-0001.dcm"}, |
| ... |
| ], |
| "videos_right": [ |
| {"path": "../data/anon_data/.../IM-0002-0001.dcm"}, |
| ... |
| ], |
| }, |
| ... |
| ] |
| |
| Important: the "videos_{artery}[i]['path']" fields are DICOM paths |
| relative to the JSON directory (the rnn_folds/ folder). |
| """ |
|
|
| def __init__( |
| self, |
| root: str, |
| meta: str, |
| train: bool, |
| length: int, |
| label: str, |
| artery: str, |
| inference: bool = False, |
| validation: bool = False, |
| transform: Optional[Callable] = None, |
| ) -> None: |
| self.root = Path(root).resolve() |
| self.train = train |
| self.length = int(length) |
| self.label = label |
| self.artery = artery.lower() |
| self.inference = inference |
| self.validation = validation |
| self.transform = transform |
|
|
| meta_path = Path(meta) |
| if not meta_path.is_absolute(): |
| meta_path = self.root / meta_path |
| meta_path = meta_path.resolve() |
|
|
| self.base_dir = meta_path.parent |
|
|
| print(f"RNN Dataset: root={self.root}, meta={meta_path}, base_dir={self.base_dir}") |
|
|
| with open(meta_path, "r", encoding="utf-8") as f: |
| dataset = json.load(f) |
|
|
| if not self.inference: |
| dataset = [rec for rec in dataset if len(rec.get(f"videos_{self.artery}", [])) > 0] |
|
|
| if validation and self.label: |
| dataset = [rec for rec in dataset if float(rec.get(self.label, 0.0)) > 0] |
|
|
| self.dataset = dataset |
| print(f"RNN Dataset loaded: {len(self.dataset)} samples after filtering") |
|
|
| artery_bin = {"left": 0, "right": 1}.get(self.artery) |
| if artery_bin is None: |
| raise ValueError(f"Unknown artery '{artery}', expected 'left' or 'right'") |
| self.artery_bin = artery_bin |
|
|
| def __len__(self) -> int: |
| return len(self.dataset) |
|
|
| def get_sample_weights(self) -> Tensor: |
| """ |
| Return sample weights for WeightedRandomSampler based on score bins. |
| |
| Each artery has its own thresholds, and bin frequency is inverted. |
| """ |
| bin_thresholds = { |
| 0: [0, 5, 10, 15], |
| 1: [0, 2, 5, 8], |
| } |
| thr0, thr1, thr2, thr3 = bin_thresholds[self.artery_bin] |
|
|
| def in_bin(score: float) -> int: |
| if score == thr0: |
| return 0 |
| if thr0 < score <= thr1: |
| return 1 |
| if thr1 < score <= thr2: |
| return 2 |
| if thr2 < score <= thr3: |
| return 3 |
| return 4 |
|
|
| scores = [float(rec.get(self.label, 0.0)) for rec in self.dataset] |
| bins = [in_bin(s) for s in scores] |
|
|
| counts = np.bincount(np.array(bins, dtype=np.int64), minlength=5) |
| total = int(counts.sum()) |
|
|
| weights_by_bin = np.array( |
| [(total / counts[b]) if counts[b] > 0 else 0.0 for b in range(5)], |
| dtype=np.float64, |
| ) |
|
|
| weights = np.array([weights_by_bin[b] for b in bins], dtype=np.float64) |
| print( |
| "RNN sample weights counts:", |
| int(counts[0]), |
| int(counts[1]), |
| int(counts[2]), |
| int(counts[3]), |
| int(counts[4]), |
| ) |
| return torch.as_tensor(weights, dtype=DTYPE) |
|
|
| def __getitem__(self, idx: int) -> Tuple[Tensor, Tensor, Tensor, Any]: |
| """ |
| Return: |
| clips: Tensor stack of clips (N_clips, C, T, H, W) after transform |
| label: Tensor(1,) binary label (0/1) |
| target: Tensor(1,) regression target (log1p(score)) |
| suid: study identifier (study_uid) |
| """ |
| rec = self.dataset[idx] |
| suid = rec["study_uid"] |
|
|
| if self.label: |
| bin_thresholds = { |
| 0: 15, |
| 1: 5, |
| } |
| score = float(rec.get(self.label, 0.0)) |
| label = torch.tensor( |
| [1.0 if score > bin_thresholds[self.artery_bin] else 0.0], |
| dtype=DTYPE, |
| ) |
| target = torch.tensor([np.log1p(score)], dtype=DTYPE) |
| else: |
| label = torch.tensor([0.0], dtype=DTYPE) |
| target = torch.tensor([0.0], dtype=DTYPE) |
|
|
| videos_list = rec.get(f"videos_{self.artery}", []) |
| nv = len(videos_list) |
|
|
| if self.inference: |
| if nv == 0: |
| return torch.zeros(0), label, target, suid |
| seq_indices = range(nv) |
| else: |
| if nv == 0: |
| raise ValueError(f"No videos for artery={self.artery} in record {suid}") |
| seq_indices = torch.randint(low=0, high=nv, size=(4,)) |
|
|
| clips = [] |
| for vi in seq_indices: |
| vi_idx = int(vi) |
| video_rec = videos_list[vi_idx] |
| rel_path = video_rec["path"] |
|
|
| full_path = (self.base_dir / rel_path).resolve() |
| if not full_path.exists(): |
| raise FileNotFoundError( |
| f"DICOM not found: {full_path}\n" |
| f" base_dir={self.base_dir}\n" |
| f" rel_path='{rel_path}'\n" |
| f" study={suid}" |
| ) |
|
|
| video = pydicom.dcmread(str(full_path)).pixel_array |
|
|
| if video.ndim != 3: |
| raise ValueError(f"Expected 3D video, got {video.shape} in {full_path}") |
|
|
| if video.shape[0] > 128 and video.shape[-1] <= 128: |
| video = np.moveaxis(video, -1, 0) |
|
|
| if video.dtype == np.uint16: |
| vmax = int(np.max(video)) |
| if vmax <= 0: |
| raise ValueError(f"Invalid vmax={vmax} in {full_path}") |
| video = (video.astype(np.float32) * (255.0 / vmax)).clip(0, 255).astype(np.uint8) |
| else: |
| video = video.astype(np.uint8) |
|
|
| while video.shape[0] < self.length: |
| video = np.concatenate([video, video], axis=0) |
|
|
| t = int(video.shape[0]) |
| if self.train: |
| begin = torch.randint(low=0, high=t - self.length + 1, size=(1,)).item() |
| else: |
| begin = (t - self.length) // 2 |
| video = video[begin: begin + self.length] |
|
|
| video = torch.from_numpy(np.stack([video, video, video], axis=-1)) |
|
|
| if self.transform is not None: |
| video = self.transform(video) |
|
|
| clips.append(video) |
|
|
| clips = torch.stack(clips, dim=0) if clips else torch.zeros(0, dtype=DTYPE) |
|
|
| return clips, label, target, suid |
|
|