Download scripts/fake_data.py from OneScience-Group/MassConservingCNN: direct link, hf CLI and curl.
- Browser
- Download file 3.67 kB
-
https://huggingface.co/OneScience-Group/MassConservingCNN/resolve/main/scripts/fake_data.py
- Command line
-
hf download hf://OneScience-Group/MassConservingCNN/scripts/fake_data.py
-
curl -L -o fake_data.py https://huggingface.co/OneScience-Group/MassConservingCNN/resolve/main/scripts/fake_data.py
3.67 kB
| """Generate small, structured modified-shallow-water analysis pairs.""" | |
| import argparse | |
| from pathlib import Path | |
| import numpy as np | |
| import yaml | |
| ROOT = Path(__file__).resolve().parents[1] | |
| def periodic_gaussian(x, center, width): | |
| distance = np.minimum(np.abs(x - center), 1.0 - np.abs(x - center)) | |
| return np.exp(-0.5 * (distance / width) ** 2) | |
| def make_split(path, count, config, seed): | |
| rng = np.random.default_rng(seed) | |
| n = int(config["data"]["grid_points"]) | |
| x = np.arange(n, dtype=np.float32) / n | |
| xa = np.empty((count, 3, n), dtype=np.float32) | |
| target = np.empty_like(xa) | |
| radar = np.empty((count, 1, n), dtype=np.float32) | |
| for sample in range(count): | |
| phase = rng.uniform(0.0, 1.0) | |
| wave = np.sin(2 * np.pi * (x - phase)) | |
| harmonic = np.sin(4 * np.pi * (x - 0.6 * phase)) | |
| convective = periodic_gaussian(x, (phase + 0.23) % 1.0, 0.045) | |
| secondary = periodic_gaussian(x, (phase + 0.66) % 1.0, 0.07) | |
| u_true = 0.75 * wave + 0.22 * harmonic - 0.28 * np.gradient(convective) | |
| h_true = 10.0 + 0.35 * np.cos(2 * np.pi * (x - phase)) + 0.5 * convective | |
| convergence = np.maximum(-np.gradient(u_true), 0.0) | |
| r_true = np.maximum(0.0, 0.7 * convective + 0.28 * convergence - 0.09) | |
| rain_mask = (r_true > 0.08).astype(np.float32) | |
| # Smooth EnKF-like errors are tied to convection and dry-region mass drift. | |
| dry = 1.0 - rain_mask | |
| u_error = 0.11 * secondary - 0.07 * convective + 0.025 * harmonic | |
| h_error = 0.16 * dry + 0.08 * secondary - 0.05 * convective | |
| r_error = 0.13 * secondary * dry - 0.06 * convective | |
| xa[sample, 0] = u_true + u_error | |
| xa[sample, 1] = h_true + h_error | |
| xa[sample, 2] = np.maximum(0.0, r_true + r_error) | |
| target[sample] = np.stack((u_true, h_true, r_true)) | |
| radar[sample, 0] = rain_mask | |
| # Shared synthetic climatology keeps train and validation normalization identical. | |
| means = np.asarray([0.0, 10.0], dtype=np.float32) | |
| stds = np.asarray([0.6, 0.4, 0.3], dtype=np.float32) | |
| normalized_x = xa.copy() | |
| normalized_y = target.copy() | |
| normalized_x[:, :2] = (xa[:, :2] - means[None, :, None]) / stds[None, :2, None] | |
| normalized_y[:, :2] = (target[:, :2] - means[None, :, None]) / stds[None, :2, None] | |
| normalized_x[:, 2] = xa[:, 2] / stds[2] | |
| normalized_y[:, 2] = target[:, 2] / stds[2] | |
| inputs = np.concatenate((normalized_x, radar), axis=1).astype(np.float32) | |
| np.savez_compressed( | |
| path, inputs=inputs, targets=normalized_y.astype(np.float32), xa=xa, | |
| targets_physical=target, radar=radar, climate_mean_uh=means, | |
| climate_std_uhr=stds, format_version=np.asarray(config["data"]["format_version"]), | |
| variable_order=np.asarray(["u", "h", "r"]), input_layout=np.asarray("BCX"), | |
| data_source=np.asarray("structured_synthetic_msw"), | |
| ) | |
| def main(): | |
| parser = argparse.ArgumentParser() | |
| parser.add_argument("--force", action="store_true") | |
| args = parser.parse_args() | |
| config = yaml.safe_load((ROOT / "conf/config.yaml").read_text()) | |
| output = ROOT / config["data"]["root"] | |
| output.mkdir(parents=True, exist_ok=True) | |
| splits = (("train.npz", int(config["data"]["train_samples"])), | |
| ("validation.npz", int(config["data"]["validation_samples"]))) | |
| for offset, (name, count) in enumerate(splits): | |
| path = output / name | |
| if args.force or not path.exists(): | |
| make_split(path, count, config, int(config["seed"]) + offset) | |
| print(f"generated={path.relative_to(ROOT)} samples={count} shape=({count},4,250)") | |
| if __name__ == "__main__": | |
| main() | |