LSTNN-DEMO / app.py
amyyy09's picture
xd
0328a0e verified
Raw
History Blame
5.96 kB
import numpy as np
import torch
import gradio as gr
import matplotlib.pyplot as plt
from matplotlib.lines import Line2D
from lsttn_model import build_model
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
ARTIFACTS_DIR = "lsttn_artifacts"
BG = "#1a1a1a"
GREEN = "#16a34a"
YELLOW = "#eab308"
RED = "#dc2626"
# ---- Carga de artefactos (una sola vez, al iniciar el Space) ----
meta = np.load(f"{ARTIFACTS_DIR}/meta.npy")
NUM_NODES, WINDOW_SIZE, HORIZON = int(meta[0]), int(meta[1]), int(meta[2])
A = np.load(f"{ARTIFACTS_DIR}/adjacency.npy")
mean_flujo, std_flujo = np.load(f"{ARTIFACTS_DIR}/mean_std.npy")
demo_data = np.load(f"{ARTIFACTS_DIR}/demo_samples.npz")
X_demo, Y_demo = demo_data["X"], demo_data["Y"]
N_SAMPLES = X_demo.shape[0]
POS = np.load(f"{ARTIFACTS_DIR}/node_positions.npy")
EDGES = np.load(f"{ARTIFACTS_DIR}/edges.npy")
model = build_model(NUM_NODES, WINDOW_SIZE, HORIZON, A, device=DEVICE)
state = torch.load(f"{ARTIFACTS_DIR}/model_state.pt", map_location=DEVICE)
model.load_state_dict(state)
model.eval()
def _predict_all_nodes(sample_id: int, horizon_step: int = 0):
x = torch.tensor(X_demo[sample_id:sample_id + 1], dtype=torch.float32, device=DEVICE)
with torch.no_grad():
pred = model(x, x) # (1, N, horizon)
flow_norm = pred[0, :, horizon_step].cpu().numpy()
flow_real = np.maximum(0, flow_norm * std_flujo + mean_flujo)
return flow_real
def _flow_color(v, low, high):
if v <= low:
return GREEN
elif v <= high:
return YELLOW
return RED
def plot_network(sample_id: int, horizon_step: int):
sample_id = int(sample_id)
horizon_step = int(horizon_step)
flow_real = _predict_all_nodes(sample_id, horizon_step)
low, high = np.percentile(flow_real, [33, 66])
fig, ax = plt.subplots(figsize=(8, 8))
fig.patch.set_facecolor(BG)
ax.set_facecolor(BG)
for i, j in EDGES:
avg = (flow_real[i] + flow_real[j]) / 2
ax.plot(
[POS[i, 0], POS[j, 0]], [POS[i, 1], POS[j, 1]],
color=_flow_color(avg, low, high), linewidth=2, alpha=0.85, zorder=1,
)
colors = [_flow_color(v, low, high) for v in flow_real]
ax.scatter(POS[:, 0], POS[:, 1], c=colors, s=35, zorder=2, edgecolors="white", linewidths=0.4)
ax.axis("off")
ax.set_title(
f"Red PEMS-04 — flujo previsto a t+{horizon_step + 1} (muestra {sample_id})",
color="white", fontsize=12,
)
legend_handles = [
Line2D([0], [0], color=GREEN, lw=3, label="Flujo bajo"),
Line2D([0], [0], color=YELLOW, lw=3, label="Flujo medio"),
Line2D([0], [0], color=RED, lw=3, label="Flujo alto"),
]
ax.legend(
handles=legend_handles, loc="lower center", bbox_to_anchor=(0.5, -0.05),
ncol=3, frameon=False, labelcolor="white",
)
fig.tight_layout()
return fig
def plot_sensor_detail(sensor_id: int, sample_id: int):
sensor_id, sample_id = int(sensor_id), int(sample_id)
x = torch.tensor(X_demo[sample_id:sample_id + 1], dtype=torch.float32, device=DEVICE)
y_true = Y_demo[sample_id, sensor_id]
with torch.no_grad():
pred = model(x, x)
y_pred = pred[0, sensor_id].cpu().numpy()
true_real = y_true * std_flujo + mean_flujo
pred_real = np.maximum(0, y_pred * std_flujo + mean_flujo)
mae = np.mean(np.abs(true_real - pred_real))
rmse = np.sqrt(np.mean((true_real - pred_real) ** 2))
fig, ax = plt.subplots(figsize=(6, 4))
fig.patch.set_facecolor(BG)
ax.set_facecolor(BG)
steps = np.arange(1, HORIZON + 1)
ax.plot(steps, true_real, marker="o", label="Real", color="white")
ax.plot(steps, pred_real, marker="x", label="Predicción LSTTN", color=GREEN, linestyle="--")
ax.set_xlabel("Paso futuro (x5 min)", color="white")
ax.set_ylabel("Flujo de tráfico (vehículos)", color="white")
ax.set_title(f"Sensor {sensor_id} — muestra {sample_id}", color="white")
ax.tick_params(colors="white")
for spine in ax.spines.values():
spine.set_color("#555555")
ax.legend(labelcolor="white", facecolor=BG, edgecolor="#555555")
ax.grid(True, linestyle=":", alpha=0.3, color="white")
fig.tight_layout()
metrics_md = f"**MAE:** {mae:.2f} veh &nbsp;&nbsp; **RMSE:** {rmse:.2f} veh"
return fig, metrics_md
with gr.Blocks(title="LSTTN — Pronóstico de tráfico PEMS-04", theme=gr.themes.Base()) as demo:
gr.Markdown(
"# LSTTN — Pronóstico de flujo de tráfico\n"
f"Red vial PEMS-04 · {NUM_NODES} sensores · ventana de {WINDOW_SIZE} pasos → horizonte de {HORIZON} pasos."
)
with gr.Tabs():
with gr.Tab("Mapa de red"):
with gr.Row():
sample_slider_net = gr.Slider(0, N_SAMPLES - 1, value=0, step=1, label="Muestra de test")
horizon_slider = gr.Slider(0, HORIZON - 1, value=0, step=1, label="Paso futuro (horizonte)")
net_plot = gr.Plot()
sample_slider_net.change(plot_network, [sample_slider_net, horizon_slider], net_plot)
horizon_slider.change(plot_network, [sample_slider_net, horizon_slider], net_plot)
demo.load(plot_network, [sample_slider_net, horizon_slider], net_plot)
with gr.Tab("Detalle por sensor"):
with gr.Row():
sensor_slider = gr.Slider(0, NUM_NODES - 1, value=0, step=1, label="ID de sensor")
sample_slider_det = gr.Slider(0, N_SAMPLES - 1, value=0, step=1, label="Muestra de test")
detail_plot = gr.Plot()
metrics_out = gr.Markdown()
sensor_slider.change(plot_sensor_detail, [sensor_slider, sample_slider_det], [detail_plot, metrics_out])
sample_slider_det.change(plot_sensor_detail, [sensor_slider, sample_slider_det], [detail_plot, metrics_out])
demo.load(plot_sensor_detail, [sensor_slider, sample_slider_det], [detail_plot, metrics_out])
if __name__ == "__main__":
demo.launch()