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.
0279
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 