ML-MODIS / scripts /fake_data.py
zhangrenchao's picture
Publish ML-MODIS engineering reproduction
73ddb67 verified
Raw
History Blame Contribute Delete
6.87 kB
#!/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()