Spaces:
Running
Running
File size: 2,121 Bytes
fbaf630 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 | from __future__ import annotations
from typing import List, Tuple, Optional
import os
import numpy as np
import torch
from torch.utils.data import Dataset
from .preprocess import (
read_dicom_uint8,
ensure_length_center_crop,
test_like_transform,
IMAGENET_MEAN,
IMAGENET_STD,
)
DTYPE = torch.float16
class SyntaxInferenceDataset(Dataset):
"""
Загрузчик DICOM для инференса одной артерии.
Возвращает:
videos: (S=1, C, T, H, W) — один клип на элемент, RGB получен дублированием серого канала
label: заглушка (torch.tensor([0], dtype=DTYPE))
target: заглушка (torch.tensor([0.0], dtype=DTYPE))
uid: имя файла
"""
def __init__(
self,
files: List[str],
frames_per_clip: int,
video_size: Tuple[int, int],
mean: Tuple[float, float, float] = IMAGENET_MEAN,
std: Tuple[float, float, float] = IMAGENET_STD,
transform: Optional[torch.nn.Module] = None,
):
self.files = list(files)
self.frames = int(frames_per_clip)
self.video_size = video_size
self.mean = mean
self.std = std
self.transform = transform or test_like_transform(video_size)
def __len__(self) -> int:
return len(self.files)
def __getitem__(self, idx: int):
path = self.files[idx]
uid = os.path.basename(path)
# (T,H,W) uint8 → выравнивание длины
arr = read_dicom_uint8(path)
arr = ensure_length_center_crop(arr, self.frames)
# (T,H,W,3): серый → RGB
vid_thwc = np.stack([arr, arr, arr], axis=-1)
vid_thwc = torch.tensor(vid_thwc)
# (C,T,H,W) → добавляем размерность последовательности S=1
vid_cthw = self.transform(vid_thwc)
videos = vid_cthw.unsqueeze(0)
label = torch.tensor([0], dtype=DTYPE)
target = torch.tensor([0.0], dtype=DTYPE)
return videos, label, target, uid
|