from __future__ import annotations import json import random from pathlib import Path import numpy as np import pandas as pd import torch import trackio from model import ChronosMicroGRU, parameter_count from safetensors.torch import save_file from torch.nn import functional as F from torch.utils.data import DataLoader, TensorDataset PROJECT_DIR = Path(__file__).resolve().parent ROOT_DIR = PROJECT_DIR.parents[1] DATA_DIR = ROOT_DIR / "projects" / "edge-sentinel-ml" / "data" ARTIFACT_DIR = PROJECT_DIR / "artifacts" / "chronos-microgru" CHANNELS = [ "temperature", "pressure", "vibration", "current", "flow", "packet_rate", ] CONTEXT = 32 HORIZON = 8 STRIDE = 8 def seed_everything(seed: int) -> None: random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) def load_frame(name: str) -> pd.DataFrame: path = DATA_DIR / f"{name}.parquet" if not path.exists(): raise FileNotFoundError( f"{path} is missing. Generate the Edge Sentinel dataset first." ) return pd.read_parquet(path) def build_windows( frame: pd.DataFrame, mean: np.ndarray, scale: np.ndarray, ) -> tuple[torch.Tensor, torch.Tensor]: contexts = [] targets = [] total_length = CONTEXT + HORIZON for _, device in frame.groupby("device_id", sort=False): values = (device[CHANNELS].to_numpy(dtype=np.float32) - mean) / scale labels = device["label"].to_numpy(dtype=np.int64) for start in range(0, len(device) - total_length + 1, STRIDE): stop = start + total_length if labels[start:stop].any(): continue contexts.append(values[start : start + CONTEXT]) targets.append(values[start + CONTEXT : stop]) return ( torch.from_numpy(np.stack(contexts).astype(np.float32)), torch.from_numpy(np.stack(targets).astype(np.float32)), ) def gaussian_nll( mean: torch.Tensor, log_variance: torch.Tensor, target: torch.Tensor, ) -> torch.Tensor: return 0.5 * (log_variance + (target - mean).square() / log_variance.exp()).mean() @torch.inference_mode() def predict( model: ChronosMicroGRU, contexts: torch.Tensor, ) -> tuple[np.ndarray, np.ndarray]: model.eval() means = [] log_variances = [] for start in range(0, len(contexts), 512): mean, log_variance = model(contexts[start : start + 512]) means.append(mean.numpy()) log_variances.append(log_variance.numpy()) return np.concatenate(means), np.concatenate(log_variances) def fit_variance_scale( mean: np.ndarray, log_variance: np.ndarray, target: np.ndarray, ) -> float: candidates = np.linspace(0.25, 4.0, 301) variance = np.exp(log_variance) losses = [ np.mean( 0.5 * (np.log(variance * candidate) + (target - mean) ** 2 / (variance * candidate)) ) for candidate in candidates ] return float(candidates[int(np.argmin(losses))]) def forecast_metrics( prediction: np.ndarray, log_variance: np.ndarray, target: np.ndarray, variance_scale: float, ) -> dict: error = prediction - target variance = np.exp(log_variance) * variance_scale standard_deviation = np.sqrt(variance) lower = prediction - 1.644854 * standard_deviation upper = prediction + 1.644854 * standard_deviation return { "normalized_rmse": float(np.sqrt(np.mean(error**2))), "normalized_mae": float(np.mean(np.abs(error))), "gaussian_nll": float(np.mean(0.5 * (np.log(variance) + error**2 / variance))), "interval_90_coverage": float(np.mean((target >= lower) & (target <= upper))), "mean_interval_width": float(np.mean(upper - lower)), } def channel_metrics( prediction: np.ndarray, target: np.ndarray, scale: np.ndarray, ) -> dict: error = (prediction - target) * scale.reshape(1, 1, -1) return { channel: { "rmse": float(np.sqrt(np.mean(error[..., index] ** 2))), "mae": float(np.mean(np.abs(error[..., index]))), } for index, channel in enumerate(CHANNELS) } def main() -> None: seed_everything(2037) train_frame = load_frame("train") validation_frame = load_frame("validation") test_frame = load_frame("test") normal_train = train_frame[train_frame["label"] == 0] mean = normal_train[CHANNELS].to_numpy(dtype=np.float32).mean(axis=0) scale = np.maximum( normal_train[CHANNELS].to_numpy(dtype=np.float32).std(axis=0), 1e-5, ) train_context, train_target = build_windows(train_frame, mean, scale) validation_context, validation_target = build_windows( validation_frame, mean, scale, ) test_context, test_target = build_windows(test_frame, mean, scale) loader = DataLoader( TensorDataset(train_context, train_target), batch_size=128, shuffle=True, generator=torch.Generator().manual_seed(2037), ) model = ChronosMicroGRU(channels=len(CHANNELS), horizon=HORIZON) optimizer = torch.optim.AdamW(model.parameters(), lr=0.002, weight_decay=0.001) epochs = 55 scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=epochs) best_validation_nll = float("inf") best_epoch = 0 best_state = None trackio.init( project="chronos-microgru", name="probabilistic-eight-step-v1", config={ "parameters": parameter_count(model), "context": CONTEXT, "horizon": HORIZON, "channels": CHANNELS, "train_windows": len(train_context), }, ) for epoch in range(1, epochs + 1): model.train() running_loss = 0.0 examples = 0 for context, target in loader: prediction, log_variance = model(context) nll = gaussian_nll(prediction, log_variance, target) mean_loss = F.mse_loss(prediction, target) loss = nll + 0.08 * mean_loss optimizer.zero_grad(set_to_none=True) loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() running_loss += loss.item() * len(context) examples += len(context) scheduler.step() validation_mean, validation_log_variance = predict( model, validation_context, ) validation_nll = float( gaussian_nll( torch.from_numpy(validation_mean), torch.from_numpy(validation_log_variance), validation_target, ) ) if validation_nll < best_validation_nll: best_validation_nll = validation_nll best_epoch = epoch best_state = { key: value.detach().cpu().clone() for key, value in model.state_dict().items() } trackio.log( { "epoch": epoch, "train_loss": running_loss / examples, "validation_gaussian_nll": validation_nll, "learning_rate": scheduler.get_last_lr()[0], } ) assert best_state is not None model.load_state_dict(best_state) validation_mean, validation_log_variance = predict(model, validation_context) variance_scale = fit_variance_scale( validation_mean, validation_log_variance, validation_target.numpy(), ) test_mean, test_log_variance = predict(model, test_context) persistence = np.repeat( test_context[:, -1:, :].numpy(), HORIZON, axis=1, ) model_metrics = forecast_metrics( test_mean, test_log_variance, test_target.numpy(), variance_scale, ) persistence_rmse = float(np.sqrt(np.mean((persistence - test_target.numpy()) ** 2))) persistence_mae = float(np.mean(np.abs(persistence - test_target.numpy()))) results = { "model": "Chronos MicroGRU", "parameters": parameter_count(model), "context": CONTEXT, "horizon": HORIZON, "channels": CHANNELS, "train_windows": len(train_context), "validation_windows": len(validation_context), "test_windows": len(test_context), "best_epoch": best_epoch, "variance_scale": variance_scale, "model_test": model_metrics, "persistence_test": { "normalized_rmse": persistence_rmse, "normalized_mae": persistence_mae, }, "rmse_improvement_percent": 100 * (persistence_rmse - model_metrics["normalized_rmse"]) / persistence_rmse, "model_channel_metrics": channel_metrics( test_mean, test_target.numpy(), scale, ), "persistence_channel_metrics": channel_metrics( persistence, test_target.numpy(), scale, ), } trackio.log( { "test_normalized_rmse": model_metrics["normalized_rmse"], "test_interval_90_coverage": model_metrics["interval_90_coverage"], "persistence_normalized_rmse": persistence_rmse, } ) trackio.finish() ARTIFACT_DIR.mkdir(parents=True, exist_ok=True) save_file(model.state_dict(), ARTIFACT_DIR / "model.safetensors") np.savez( ARTIFACT_DIR / "normalization.npz", mean=mean, scale=scale, variance_scale=variance_scale, ) (ARTIFACT_DIR / "evaluation.json").write_text( json.dumps(results, indent=2), encoding="utf-8", ) print(json.dumps(results, indent=2)) if __name__ == "__main__": main()