Spaces:
Configuration error
Configuration error
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
|