| 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() | |
| 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() | |