ARotting's picture
Publish Evolino-inspired evolved recurrent forecasting features
9147d47 verified
Raw
History Blame Contribute Delete
2.95 kB
from __future__ import annotations
import json
from pathlib import Path
import gradio as gr
import numpy as np
import plotly.graph_objects as go
from model import RecurrentFeatures
from safetensors.numpy import load_file
from train import make_problem
PROJECT_DIR = Path(__file__).resolve().parent
ARTIFACT_DIR = PROJECT_DIR / "artifacts" / "evolino-pocket"
REPORT = json.loads((ARTIFACT_DIR / "evaluation.json").read_text(encoding="utf-8"))
def load_model(name: str) -> RecurrentFeatures:
tensors = load_file(ARTIFACT_DIR / name)
return RecurrentFeatures(
input_weight=tensors["input_weight"],
recurrent_weight=tensors["recurrent_weight"],
bias=tensors["bias"],
readout=tensors["readout"],
)
EVOLVED = load_model("evolved_features.safetensors")
RANDOM = load_model("random_control.safetensors")
SERIES, TARGETS, SPLIT = make_problem()
EVOLVED_PREDICTIONS = EVOLVED.predict(SERIES)
RANDOM_PREDICTIONS = RANDOM.predict(SERIES)
def inspect_forecast(start: int, horizon: int, points: int) -> tuple[go.Figure, dict]:
horizon_index = int(horizon) - 1
indices = SPLIT["test"][int(start) : int(start) + int(points)]
truth = TARGETS[indices, horizon_index]
evolved = EVOLVED_PREDICTIONS[indices, horizon_index]
random = RANDOM_PREDICTIONS[indices, horizon_index]
figure = go.Figure()
figure.add_scatter(y=truth, name="True future", line={"width": 3})
figure.add_scatter(y=evolved, name="Evolved features")
figure.add_scatter(y=random, name="Random control", line={"dash": "dot"})
figure.update_layout(
template="plotly_dark",
title=f"Mackey–Glass direct forecast, horizon +{int(horizon)}",
xaxis_title="Held-out test step",
yaxis_title="Normalized value",
)
metrics = {
"horizon": int(horizon),
"evolved_window_rmse": float(np.sqrt(np.mean((evolved - truth) ** 2))),
"random_window_rmse": float(np.sqrt(np.mean((random - truth) ** 2))),
**REPORT["aggregate"],
}
return figure, metrics
with gr.Blocks(title="Evolino Pocket") as demo:
gr.Markdown(
"# Evolino Pocket\n"
"Inspect an Evolino-inspired hybrid: mutation search shapes a recurrent "
"feature generator while ridge regression solves its multihorizon "
"readout exactly."
)
with gr.Row():
start = gr.Slider(0, 650, value=20, step=1, label="Test window start")
horizon = gr.Slider(1, 12, value=12, step=1, label="Forecast horizon")
points = gr.Slider(30, 150, value=100, step=10, label="Visible points")
initial = inspect_forecast(20, 12, 100)
chart = gr.Plot(value=initial[0])
metrics = gr.JSON(value=initial[1])
button = gr.Button("Inspect forecast", variant="primary")
button.click(
inspect_forecast,
inputs=[start, horizon, points],
outputs=[chart, metrics],
)
if __name__ == "__main__":
demo.launch()