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
0likes279downloads
train.py68 linesDownload Raw Back to sparsewake
1from __future__ import annotations2 3import numpy as np4import torch5from torch.utils.data import DataLoader, TensorDataset6 7from .models import TemporalMLP8 9 10def standardize_train_val_test(x: np.ndarray, train_idx: np.ndarray) -> tuple[np.ndarray, np.ndarray, np.ndarray]:11    mean = x[train_idx].mean(axis=0, keepdims=True)12    std = x[train_idx].std(axis=0, keepdims=True)13    std = np.where(std < 1e-8, 1.0, std)14    return ((x - mean) / std).astype(np.float32), mean.astype(np.float32), std.astype(np.float32)15 16 17def train_temporal_mlp(18    x: np.ndarray,19    y: np.ndarray,20    train_idx: np.ndarray,21    val_idx: np.ndarray,22    output_dim: int = 2,23    epochs: int = 50,24    batch_size: int = 1024,25    lr: float = 1e-3,26    weight_decay: float = 1e-4,27    seed: int = 1,28    device: str = "cpu",29) -> TemporalMLP:30    torch.manual_seed(seed)31    model = TemporalMLP(x.shape[1], output_dim=output_dim)32    model.to(device)33    opt = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=weight_decay)34    loss_fn = torch.nn.MSELoss()35    train_ds = TensorDataset(torch.from_numpy(x[train_idx]), torch.from_numpy(y[train_idx, :output_dim].astype(np.float32)))36    train_loader = DataLoader(train_ds, batch_size=batch_size, shuffle=True)37    best_state = None38    best_val = float("inf")39    patience = 840    stale = 041    for _ in range(epochs):42        model.train()43        for xb, yb in train_loader:44            xb = xb.to(device)45            yb = yb.to(device)46            opt.zero_grad()47            loss = loss_fn(model(xb), yb)48            loss.backward()49            opt.step()50        model.eval()51        with torch.no_grad():52            xv = torch.from_numpy(x[val_idx]).to(device)53            yv = torch.from_numpy(y[val_idx, :output_dim].astype(np.float32)).to(device)54            val = float(loss_fn(model(xv), yv).cpu())55        if val < best_val:56            best_val = val57            best_state = {k: v.detach().cpu().clone() for k, v in model.state_dict().items()}58            stale = 059        else:60            stale += 161            if stale >= patience:62                break63    if best_state is not None:64        model.load_state_dict(best_state)65    model.to("cpu")66    return model67 68