zhangrenchao's picture
Publish CausalModelEvaluation engineering reproduction
20cdc88 verified
Raw
History Blame Contribute Delete
4.9 kB
"""Generate seasonal directed VAR networks and precipitation fields for CME."""
import argparse
from pathlib import Path
import numpy as np
import yaml
ROOT = Path(__file__).resolve().parents[1]
def stable_network(rng, nodes, season_index):
coefficients = np.zeros((nodes, nodes, 10), dtype=np.float32)
coefficients[np.arange(nodes), np.arange(nodes), 0] = 0.48 + 0.04 * season_index
for source in range(nodes):
target = (source + 3 + season_index) % nodes
lag = (source + 2 * season_index) % 4
coefficients[source, target, lag] = (0.13 + 0.03 * (source % 3)) * (-1 if source % 5 == 0 else 1)
if source % 4 == 0:
coefficients[source, (source + 11) % nodes, (lag + 1) % 6] = -0.11
return coefficients
def simulate(rng, coefficients, samples, time_steps):
nodes, max_lag = coefficients.shape[0], coefficients.shape[2]
output = np.zeros((samples, time_steps, nodes), dtype=np.float32)
for sample in range(samples):
series = rng.normal(0, 0.35, (time_steps + max_lag, nodes)).astype(np.float32)
for time in range(max_lag, time_steps + max_lag):
forcing = np.zeros(nodes, dtype=np.float32)
for lag in range(1, max_lag + 1):
forcing += series[time - lag] @ coefficients[:, :, lag - 1]
series[time] = forcing + rng.normal(0, 0.38, nodes)
output[sample] = series[max_lag:]
return output
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())
data, seed = config["data"], int(config["seed"])
path = ROOT / config["paths"]["dataset"]
path.parent.mkdir(parents=True, exist_ok=True)
if path.exists() and not args.force:
print(f"exists={path.relative_to(ROOT)} use --force to regenerate")
return
rng = np.random.default_rng(seed)
seasons, models = len(data["seasons"]), int(data["model_count"])
nodes, samples, steps = int(data["nodes"]), int(data["samples"]), int(data["time_steps"])
reference_coefficients = np.stack([stable_network(rng, nodes, season) for season in range(seasons)])
reference_series = np.stack([simulate(rng, reference_coefficients[s], samples, steps) for s in range(seasons)])
model_series = np.empty((models, seasons, samples, steps, nodes), dtype=np.float32)
model_coefficients = np.empty((models, seasons, nodes, nodes, 10), dtype=np.float32)
quality = np.linspace(0.92, 0.35, models).astype(np.float32)
for model in range(models):
for season in range(seasons):
coefficients = reference_coefficients[season].copy()
cross = ~np.eye(nodes, dtype=bool)
coefficients[cross] *= quality[model]
mutation_count = 3 + 3 * model
for _ in range(mutation_count):
source, target = rng.integers(0, nodes, 2)
if source != target:
coefficients[source, target, rng.integers(0, 7)] = rng.choice([-1, 1]) * rng.uniform(0.08, 0.16)
model_coefficients[model, season] = coefficients
model_series[model, season] = simulate(rng, coefficients, samples, steps)
lat_count, lon_count = map(int, data["grid_shape"])
latitude = np.linspace(-90.0, 90.0, lat_count, dtype=np.float32)
longitude = np.linspace(0, 360, lon_count, endpoint=False, dtype=np.float32)
lat2d, lon2d = np.meshgrid(latitude, longitude, indexing="ij")
reference_precip = (3.0 + 2.1 * np.cos(np.deg2rad(lat2d)) ** 2 +
0.45 * np.sin(np.deg2rad(2 * lon2d))).astype(np.float32)
precip_fields, delta = [], []
for model, q in enumerate(quality):
bias = (1 - q) * (0.8 * np.sin(np.deg2rad(lat2d)) + 0.35 * np.cos(np.deg2rad(lon2d)))
precip_fields.append(reference_precip * (0.88 + 0.12 * q) + bias + rng.normal(0, 0.06 + 0.08 * (1 - q), reference_precip.shape))
delta.append(0.18 + 0.95 * (q - 0.62) ** 2 + rng.normal(0, 0.025))
np.savez_compressed(path, format_version=np.asarray(data["format_version"]), seasons=np.asarray(data["seasons"]),
reference_series=reference_series, model_series=model_series,
reference_coefficients=reference_coefficients, model_coefficients=model_coefficients,
reference_precipitation=reference_precip, model_precipitation=np.asarray(precip_fields, np.float32),
delta_precipitation=np.asarray(delta, np.float32), latitude_degrees=latitude,
longitude_degrees=longitude, model_quality=quality, time_step_days=np.asarray(data["time_step_days"]))
print(f"generated={path.relative_to(ROOT)} reference={reference_series.shape} models={model_series.shape}")
if __name__ == "__main__":
main()