File size: 3,553 Bytes
cb40653
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
import numpy as np
import h5py
import yaml
import torch
from torch.utils.data import Dataset, DataLoader
from torchvision import transforms


class SchwarzDatasetV2(Dataset):

    def __init__(self, h5_path, indices, log_mass_mean=None, log_mass_std=None):
        self.h5_path = h5_path
        self.indices = np.sort(np.array(indices))
        self.log_mass_mean = log_mass_mean
        self.log_mass_std = log_mass_std

        with h5py.File(h5_path, 'r') as hf:
            self.masses = hf['mass_solar'][self.indices]
            self.spins = hf['spin'][self.indices]
            self.inclinations = hf['inclination_deg'][self.indices]
            self.distances = hf['distance_mpc'][self.indices]
            self.fov_uas = hf['fov_uas'][self.indices]

        self.log_masses = np.log(self.masses)

        if self.log_mass_mean is None:
            self.log_mass_mean = float(np.mean(self.log_masses))
            self.log_mass_std = float(np.std(self.log_masses))

        self.transform = transforms.Compose([
            transforms.RandomHorizontalFlip(p=0.5),
            transforms.RandomVerticalFlip(p=0.5),
        ])

    def __len__(self):
        return len(self.indices)

    def __getitem__(self, idx):
        with h5py.File(self.h5_path, 'r') as hf:
            image = hf['images'][self.indices[idx], 0]

        image = torch.from_numpy(image).float()
        image = image.unsqueeze(0).repeat(3, 1, 1)
        image = self.transform(image)

        log_mass_norm = (self.log_masses[idx] - self.log_mass_mean) / self.log_mass_std
        spin = self.spins[idx]
        incl_rad = np.deg2rad(self.inclinations[idx])

        target = torch.tensor([log_mass_norm, spin, incl_rad], dtype=torch.float32)

        log_dist = np.log(self.distances[idx])
        log_fov = np.log(max(self.fov_uas[idx], 1e-6))
        aux = torch.tensor([log_dist, log_fov], dtype=torch.float32)

        return image, target, aux

    def denormalize_mass(self, log_mass_norm):
        return np.exp(log_mass_norm * self.log_mass_std + self.log_mass_mean)


def get_dataloaders_v2(config_path='configs/config_v2.yaml'):
    with open(config_path, 'r') as f:
        config = yaml.safe_load(f)

    h5_path = config['data']['output_path']
    val_split = float(config['data']['val_split'])
    test_split = float(config['data']['test_split'])
    batch_size = int(config['cnn']['batch_size'])

    with h5py.File(h5_path, 'r') as hf:
        n_total = hf['images'].shape[0]

    indices = np.arange(n_total)
    np.random.seed(42)
    np.random.shuffle(indices)

    n_test = int(n_total * test_split)
    n_val = int(n_total * val_split)
    n_train = n_total - n_val - n_test

    train_idx = indices[:n_train]
    val_idx = indices[n_train:n_train + n_val]
    test_idx = indices[n_train + n_val:]

    train_dataset = SchwarzDatasetV2(h5_path, train_idx)

    val_dataset = SchwarzDatasetV2(
        h5_path, val_idx,
        log_mass_mean=train_dataset.log_mass_mean,
        log_mass_std=train_dataset.log_mass_std,
    )

    test_dataset = SchwarzDatasetV2(
        h5_path, test_idx,
        log_mass_mean=train_dataset.log_mass_mean,
        log_mass_std=train_dataset.log_mass_std,
    )

    train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=0)
    val_loader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False, num_workers=0)
    test_loader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False, num_workers=0)

    return train_loader, val_loader, test_loader, train_dataset