Spaces:
Configuration error
Configuration error
That-Random-Coder
feat: SchwarzNet v2 release - Kerr ray tracer, 5-model deep ensemble, PINN GR verification, Gradio web app
cb40653 Download data/preprocessing_v2.py from HarshNarodey/schwarznet: direct link, hf CLI and curl.
- Browser
- Download file 3.55 kB
-
https://huggingface.co/spaces/HarshNarodey/schwarznet/resolve/main/data/preprocessing_v2.py
- Command line
-
hf download hf://spaces/HarshNarodey/schwarznet/data/preprocessing_v2.py
-
curl -L -o preprocessing_v2.py https://huggingface.co/spaces/HarshNarodey/schwarznet/resolve/main/data/preprocessing_v2.py
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 | |