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