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