Team Ai
Datasetpublic

SparseWake/sparsewake

SparseWake SparseWake is a synthetic benchmark for sparse temporal hydrodynamic sensing. ICLR 2027 release The expanded release adds controlled multi-source mixtures and common-prior nearest-source tasks, with complete core data banks, reference checkpoints, a small review supplement, and reproduction code with a frozen wake-library input. Download release iclr2027-v1.0rc2 The version page lists the three archives, exact sizes, checksums, extraction instructions… See the full description on the dataset page: https://huggingface.co/datasets/SparseWake/sparsewake.

sourceHugging Facecc-by-4.0updated 14d agoView on Hugging Face
0likes275downloads
train_temporal_mlp.py66 linesDownload Raw Back to scripts
1from __future__ import annotations2 3import argparse4import json5from pathlib import Path6import sys7 8import numpy as np9import yaml10 11ROOT = Path(__file__).resolve().parents[1]12sys.path.insert(0, str(ROOT / "src"))13 14from sparsewake.data import load_h515from sparsewake.evaluate import evaluate_predictions, predict16from sparsewake.features import build_design_matrix17from sparsewake.splits import pose_holdout_split18from sparsewake.train import standardize_train_val_test, train_temporal_mlp19 20 21def main() -> None:22    parser = argparse.ArgumentParser()23    parser.add_argument("--config", required=True)24    parser.add_argument("--data", default=None)25    parser.add_argument("--quick", action="store_true")26    parser.add_argument("--out", default="tables/quick_train_metrics.json")27    args = parser.parse_args()28    cfg = yaml.safe_load(Path(args.config).read_text())29    data_path = Path(args.data) if args.data else ROOT / cfg["data"]30    data = load_h5(data_path, input_key=cfg.get("input_key", "X_raw"))31    history = 4 if args.quick else int(cfg.get("history", 24))32    x, idx = build_design_matrix(data, feature_set=cfg.get("feature_set", "raw_norm"), history=history)33    y = data["target"][idx]34    pose_id = data["pose_id"][idx]35    train_idx, val_idx, test_idx = pose_holdout_split(pose_id, seed=int(cfg.get("seed", 1)))36    if args.quick:37        train_idx = train_idx[: min(len(train_idx), 512)]38        val_idx = val_idx[: min(len(val_idx), 128)]39        test_idx = test_idx[: min(len(test_idx), 128)]40    x, _, _ = standardize_train_val_test(x, train_idx)41    output_dim = 3 if cfg.get("target", "location") == "location_theta" else 242    model = train_temporal_mlp(43        x,44        y,45        train_idx,46        val_idx,47        output_dim=output_dim,48        epochs=3 if args.quick else int(cfg.get("epochs", 50)),49        batch_size=int(cfg.get("batch_size", 1024)),50        seed=int(cfg.get("seed", 1)),51    )52    pred = predict(model, x[test_idx])53    metrics = evaluate_predictions(y[test_idx, :output_dim], pred)54    metrics["quick_mode"] = bool(args.quick)55    metrics["n_train"] = int(len(train_idx))56    metrics["n_test"] = int(len(test_idx))57    out = ROOT / args.out58    out.parent.mkdir(parents=True, exist_ok=True)59    out.write_text(json.dumps(metrics, indent=2) + "\n")60    print(json.dumps(metrics, indent=2))61 62 63if __name__ == "__main__":64    main()65 66