schwarznet / data /preprocessing_v2.py
That-Random-Coder
feat: SchwarzNet v2 release - Kerr ray tracer, 5-model deep ensemble, PINN GR verification, Gradio web app
cb40653
Raw History Blame Contribute Delete
3.55 kB
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