CESM-SeasonalML / scripts /fake_data.py
zhangrenchao's picture
Publish CESM-SeasonalML engineering reproduction
b300acf verified
Raw
History Blame Contribute Delete
4.67 kB
"""Generate dimensionally faithful synthetic EOF sequences and precipitation fields."""
import json
import sys
from pathlib import Path
import numpy as np
import yaml
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from model.cesm_seasonal_ml import CLASS_NAMES, DATA_FORMAT_VERSION, canonical_patterns
VARIABLES = ["SST_TP", "SST_WP", "SST_IO", "SST_NP", "VP200_PW", "U200_NP", "Z500_ENP"]
def make_manifest(size, max_lag):
names = []
lag = 0
while len(names) < size:
for variable in VARIABLES:
for eof in range(1, 5):
names.append(f"{variable}_EOF{eof}_Lag{lag}")
if len(names) == size:
return names
lag = (lag + 1) % (max_lag + 1)
return names
def main():
config = yaml.safe_load((ROOT / "conf/config.yaml").read_text())
rng = np.random.default_rng(config["seed"])
sizes = [config["data"]["train_size"], config["data"]["validation_size"], config["data"]["test_size"]]
n = sum(sizes)
height, width = config["data"]["grid"]
patterns = canonical_patterns(height, width)
eof = np.zeros((2, n, 12, 28), dtype=np.float32)
precipitation = np.zeros((2, n, height, width), dtype=np.float32)
latent_labels = np.zeros((2, n), dtype=np.int64)
yy, xx = np.mgrid[-1:1:complex(height), -1:1:complex(width)]
texture = np.sin(2.5 * xx) * np.cos(2 * yy)
for season in range(2):
enso = rng.normal(size=n)
west_pacific = 0.45 * enso + rng.normal(scale=0.9, size=n)
circulation = -0.55 * enso + 0.35 * west_pacific + rng.normal(scale=0.75, size=n)
score = np.stack([enso - circulation, -enso + 0.2 * west_pacific, enso + 0.3 * circulation, -enso - west_pacific], axis=1)
score += rng.normal(scale=0.7 if season else 0.9, size=score.shape)
labels = score.argmax(axis=1)
latent_labels[season] = labels
for month in range(12):
persistence = 0.82 ** (11 - month)
eof[season, :, month] = rng.normal(scale=0.7, size=(n, 28))
eof[season, :, month, 0] += persistence * enso * 2.0
eof[season, :, month, 4] += persistence * west_pacific * 1.4
eof[season, :, month, 16] += persistence * circulation * 1.2
amplitude = rng.uniform(0.8, 1.5, size=n)
noise = rng.normal(scale=0.38, size=(n, height, width))
precipitation[season] = amplitude[:, None, None] * patterns[labels] + 0.18 * score.max(axis=1)[:, None, None] * texture + noise
rf_manifest = make_manifest(103, 12)
nn_manifest = make_manifest(336, 11)
nn_manifest += [f"engineering_interaction_{index:03d}" for index in range(1, 81)]
rf_features = np.empty((2, n, 103), dtype=np.float32)
nn_features = np.empty((2, n, 416), dtype=np.float32)
flat = eof.reshape(2, n, -1)
rf_features[:] = flat[:, :, :103]
# The reported 416-vector cannot be reconstructed from the published manifest.
# Preserve all 336 sequence values and add deterministic interaction features.
nn_features[:, :, :336] = flat
nn_features[:, :, 336:] = flat[:, :, :80] * flat[:, :, 28:108]
split = np.repeat(np.array([0, 1, 2], dtype=np.int8), sizes)
years = np.concatenate([np.arange(1920, 1920 + sizes[0]), np.arange(1920, 1920 + sizes[1]), np.arange(1981, 1981 + sizes[2])]).astype(np.int32)
output = ROOT / config["data"]["path"]
output.parent.mkdir(parents=True, exist_ok=True)
np.savez_compressed(output, format_version=DATA_FORMAT_VERSION, seasons=np.array(["NDJ", "JFM"]), eof_sequence=eof,
rf_features=rf_features, nn_features=nn_features, precipitation=precipitation,
latent_labels=latent_labels, split=split, years=years,
latitude=np.linspace(31, 49, height), longitude=np.linspace(-125, -102, width),
rf_manifest=np.array(rf_manifest), nn_manifest=np.array(nn_manifest), class_names=np.array(CLASS_NAMES))
metadata = {"samples": n, "split_sizes": sizes, "grid": [height, width], "eof_sequence": [2, n, 12, 28],
"rf_features": [2, n, 103], "nn_features": [2, n, 416], "synthetic": True,
"grid_status": "engineering assumption, not a paper-reported dimension",
"manifest_status": "first 336 structured placeholders plus 80 unique engineering_interaction_* dimension-preserving assumptions; not an exact paper feature list"}
(output.parent / "metadata.json").write_text(json.dumps(metadata, indent=2) + "\n")
print(f"data={output.relative_to(ROOT)} samples={n} grid={height}x{width}")
if __name__ == "__main__":
main()