#!/usr/bin/env python3 """Generate structured synthetic ERA5-MODIS monthly pairs for an executable demo.""" from __future__ import annotations import argparse import sys from pathlib import Path import numpy as np import yaml ROOT = Path(__file__).resolve().parents[1] sys.path.insert(0, str(ROOT / "model")) from ml_modis import PRESSURE_LEVELS, PRESSURE_VARIABLES, SINGLE_FEATURES, feature_names, validate_multimodal_keys def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser() parser.add_argument("--config", default=str(ROOT / "conf/config.yaml")) parser.add_argument("--samples", type=int, default=None) parser.add_argument("--output", default=None) return parser.parse_args() def ocean_mask(lat: np.ndarray, lon: np.ndarray) -> np.ndarray: """Analytic North Atlantic mask excluding coarse Greenland/Europe land shapes.""" greenland = (lat > 59) & (lon > -53) & (lon < -20 + 0.55 * (lat - 59)) europe = (lat > 50) & (lon > -10 + 0.35 * (lat - 50)) iceland = (lat > 63) & (lat < 67) & (lon > -25) & (lon < -13) north_america = (lon < -52 + 0.3 * (lat - 45)) return ~(greenland | europe | iceland | north_america) def main() -> None: args = parse_args() config = yaml.safe_load(Path(args.config).read_text()) n = int(args.samples or config["data"]["samples"]) rng = np.random.default_rng(config["runtime"]["seed"]) years = np.asarray(config["data"]["years"], dtype=np.int16) months = np.asarray(config["data"]["months"], dtype=np.int8) platforms = np.asarray(config["data"]["platforms"], dtype="U5") records = [] used = set() while len(records) < n: year = int(rng.choice(years)) month = int(rng.choice(months)) platform = str(rng.choice(platforms)) lat = int(rng.integers(45, 76)) lon = int(rng.integers(-60, 31)) key = (year, month, platform, lat, lon) if key in used or not ocean_mask(np.array([lat]), np.array([lon]))[0]: continue used.add(key) records.append(key) year = np.asarray([r[0] for r in records], dtype=np.int16) month = np.asarray([r[1] for r in records], dtype=np.int8) platform = np.asarray([r[2] for r in records], dtype="U5") lat = np.asarray([r[3] for r in records], dtype=np.float32) lon = np.asarray([r[4] for r in records], dtype=np.float32) hour = np.where(platform == "Terra", 11.0, 13.0).astype(np.float32) phase = np.deg2rad(lon + 25) + (month - 9) * 0.35 maritime = np.cos(np.deg2rad(lat - 58)) * np.cos(np.deg2rad(lon + 25)) synoptic = np.sin(phase * 1.7 + (year - 2001) * 0.43) + 0.45 * np.cos(np.deg2rad(lat * 3)) sst = 286.0 - 0.42 * (lat - 45) + 1.1 * np.cos(phase) - 0.35 * (month - 9) + 0.025 * (year - 2001) surface_pressure = 101300 + 900 * synoptic - 8 * (lat - 55) + rng.normal(0, 160, n) humidity_base = np.clip(0.82 - 0.008 * (lat - 45) + 0.08 * maritime + 0.04 * synoptic, 0.35, 0.98) stability = 0.7 * (lat - 55) - 1.8 * synoptic + rng.normal(0, 0.7, n) x = np.empty((n, 114), dtype=np.float32) column = 0 for variable in PRESSURE_VARIABLES: for level in PRESSURE_LEVELS: z = (1000 - level) / 50.0 if variable == "temperature": value = sst - 1.7 - 3.15 * z + 0.15 * stability elif variable == "specific_humidity": value = 0.010 * humidity_base * np.exp(-0.23 * z) elif variable == "relative_humidity": value = np.clip(humidity_base - 0.025 * z + 0.04 * np.sin(phase + z), 0.05, 1.0) elif variable == "u_wind": value = 5 + 0.8 * z + 2.2 * np.sin(phase) + 0.12 * (lat - 55) elif variable == "v_wind": value = 1.5 + 1.6 * np.cos(phase * 1.3) - 0.25 * z elif variable == "omega": value = -0.025 * synoptic * np.exp(-0.08 * z) elif variable == "geopotential": value = z * 50 * 9.81 + 4 * synoptic elif variable == "cloud_liquid": value = np.maximum(0, 2.2e-4 * (humidity_base - 0.55) * np.exp(-0.18 * z)) else: value = np.clip((humidity_base - 0.55) * 1.8 * np.exp(-0.12 * z), 0, 1) x[:, column] = value + rng.normal(0, max(float(np.std(value)) * 0.035, 1e-6), n) column += 1 cos_sza = np.clip(np.cos(np.deg2rad(lat - 20)) * (0.97 - 0.01 * (hour - 11)), 0, 1) singles = np.column_stack([ sst, surface_pressure, surface_pressure + 35, sst - 0.4, sst - 1.1, sst - (1 - humidity_base) * 12, x[:, 30], x[:, 40], 190 * cos_sza, 315 - 2.5 * (sst - 278), 65 + 18 * synoptic, 18 + 8 * stability, 650 + 120 * humidity_base + 20 * synoptic, 16 + 30 * humidity_base, 0.08 + 0.18 * np.maximum(synoptic, 0), 80 * np.maximum(synoptic, 0), -25 * np.maximum(-synoptic, 0), np.clip(0.25 + 0.45 * humidity_base + 0.05 * synoptic, 0, 1), np.clip((lat - 68) / 8, 0, 1), np.maximum(0, 1.8 + 1.5 * synoptic), cos_sza, lat, lon, hour, ]).astype(np.float32) x[:, 90:] = singles platform_term = np.where(platform == "Aqua", 1.0, -1.0) low_cloud = np.clip(0.22 + 0.55 * humidity_base + 0.035 * stability + 0.025 * synoptic, 0.05, 0.9) nd = 62 + 48 * humidity_base + 5 * synoptic + 0.32 * (lat - 55) + 1.8 * platform_term reff = 18.5 - 0.035 * nd + 0.055 * (sst - 278) - 0.10 * stability lwp = 58 + 115 * low_cloud + 10 * synoptic - 2.0 * stability cf = np.clip(low_cloud + 0.018 * platform_term, 0.03, 0.95) plume = np.exp(-((lat - 60) / 10) ** 2 - ((lon + 20) / 25) ** 2) eruption = (year == 2014).astype(np.float32) * (0.72 + 0.28 * (month == 10)) * plume nd *= 1 + 0.28 * eruption reff *= 1 - 0.08 * eruption lwp *= 1 + 0.008 * eruption cf = np.clip(cf * (1 + 0.11 * eruption), 0.01, 0.99) y = np.column_stack([ nd + rng.normal(0, 3.0, n), reff + rng.normal(0, 0.28, n), lwp + rng.normal(0, 5.0, n), cf + rng.normal(0, 0.018, n), ]).astype(np.float32) y[:, 0:3] = np.maximum(y[:, 0:3], 1e-3) y[:, 3] = np.clip(y[:, 3], 0.001, 0.999) payload = {"X": x, "Y": y, "year": year, "month": month, "platform": platform, "platform_hour": hour, "latitude": lat, "longitude": lon, "feature_names": np.asarray(feature_names()), "target_names": np.asarray(config["data"]["variables"]["targets"]["names"]), "format_version": np.array(config["format_version"]), "is_ocean": np.ones(n, dtype=bool), "eruption_strength": eruption.astype(np.float32)} validate_multimodal_keys(payload) output = ROOT / (args.output or config["data"]["path"]) output.parent.mkdir(parents=True, exist_ok=True) np.savez_compressed(output, **payload) print(f"output={output.relative_to(ROOT)} samples={n} shape={list(x.shape)} " f"eruption_samples={int((year == 2014).sum())}") if __name__ == "__main__": main()