ARotting's picture
Publish Parameter-matched Highway Network depth retest
5f70dab verified
Raw
History Blame Contribute Delete
3.07 kB
from __future__ import annotations
import json
from pathlib import Path
import gradio as gr
import numpy as np
import pandas as pd
import plotly.graph_objects as go
import torch
from model import HighwayNetwork, PlainDeepNetwork
from PIL import Image
from safetensors.torch import load_file
PROJECT_DIR = Path(__file__).resolve().parent
ARTIFACT_DIR = PROJECT_DIR / "artifacts" / "highway-depth-pocket"
FRAME = pd.read_parquet(PROJECT_DIR / "data" / "test.parquet")
HIGHWAY = HighwayNetwork()
HIGHWAY.load_state_dict(load_file(ARTIFACT_DIR / "highway.safetensors"))
HIGHWAY.eval()
PLAIN = PlainDeepNetwork()
PLAIN.load_state_dict(load_file(ARTIFACT_DIR / "plain.safetensors"))
PLAIN.eval()
REPORT = json.loads((ARTIFACT_DIR / "evaluation.json").read_text(encoding="utf-8"))
@torch.inference_mode()
def inspect_gates(index: int, noise: float) -> tuple[Image.Image, go.Figure, dict]:
row = FRAME.iloc[int(index) % len(FRAME)]
pixels = np.asarray(row["image"], dtype=np.float32) / 16
rng = np.random.default_rng(int(index) + 2237)
pixels = np.clip(pixels + rng.normal(0, noise, size=64), 0, 1).astype(np.float32)
tensor = torch.from_numpy(pixels)[None]
highway_logits, gates = HIGHWAY(tensor, return_gates=True)
plain_logits = PLAIN(tensor)
gate_means = gates[0].mean(1).numpy()
figure = go.Figure(go.Bar(x=list(range(1, 9)), y=gate_means))
figure.update_layout(
template="plotly_dark",
title="Transform-gate activation by highway layer",
xaxis_title="Layer",
yaxis_title="Mean transform gate",
yaxis_range=[0, 1],
)
image = Image.fromarray(
(pixels.reshape(8, 8) * 255).astype(np.uint8), mode="L"
).resize((512, 512), Image.Resampling.NEAREST)
metrics = {
"true_label": int(row["label"]),
"highway_prediction": int(highway_logits.argmax(1)),
"plain_prediction": int(plain_logits.argmax(1)),
"verified_highway_clean_mean": REPORT["results"]["highway"]["aggregate"][
"clean_accuracy_mean"
],
"verified_plain_clean_mean": REPORT["results"]["plain"]["aggregate"][
"clean_accuracy_mean"
],
}
return image, figure, metrics
with gr.Blocks(title="Highway Depth Pocket") as demo:
gr.Markdown(
"# Highway Depth Pocket\n"
"Inspect learned transform-versus-carry gates in an eight-layer Highway "
"Network beside a sixteen-layer parameter-matched plain network."
)
with gr.Row():
index = gr.Slider(0, len(FRAME) - 1, value=18, step=1, label="Test digit")
noise = gr.Slider(0, 0.4, value=0.0, step=0.02, label="Noise")
initial = inspect_gates(18, 0.0)
with gr.Row():
image = gr.Image(value=initial[0], label="Input")
chart = gr.Plot(value=initial[1], label="Gate profile")
metrics = gr.JSON(value=initial[2])
button = gr.Button("Inspect highway", variant="primary")
button.click(inspect_gates, inputs=[index, noise], outputs=[image, chart, metrics])
if __name__ == "__main__":
demo.launch()