File size: 3,394 Bytes
eb9fb7b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
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 LIFSpikingClassifier, MatchedDenseClassifier
from PIL import Image
from safetensors.torch import load_file

PROJECT_DIR = Path(__file__).resolve().parent
ARTIFACT_DIR = PROJECT_DIR / "artifacts" / "spike-pocket"
FRAME = pd.read_parquet(PROJECT_DIR / "data" / "test.parquet")
REPORT = json.loads((ARTIFACT_DIR / "evaluation.json").read_text(encoding="utf-8"))
SPIKING = LIFSpikingClassifier()
SPIKING.load_state_dict(load_file(ARTIFACT_DIR / "spiking_lif.safetensors"))
SPIKING.eval()
DENSE = MatchedDenseClassifier()
DENSE.load_state_dict(load_file(ARTIFACT_DIR / "matched_dense.safetensors"))
DENSE.eval()


@torch.inference_mode()
def inspect_spikes(
    index: int,
    timesteps: 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) + 2111)
    corrupted = np.clip(pixels + rng.normal(0, noise, size=64), 0, 1).astype(np.float32)
    tensor = torch.from_numpy(corrupted)[None]
    generator = torch.Generator().manual_seed(int(index) + 30_000)
    spike_logits, spike_rate, raster = SPIKING(
        tensor,
        timesteps=int(timesteps),
        generator=generator,
        return_raster=True,
    )
    dense_logits = DENSE(tensor)
    image = Image.fromarray(
        (corrupted.reshape(8, 8) * 255).astype(np.uint8), mode="L"
    ).resize((512, 512), Image.Resampling.NEAREST)
    figure = go.Figure(
        go.Heatmap(
            z=raster[0].numpy().T,
            colorscale=[[0, "#080b14"], [1, "#49e6ff"]],
            showscale=False,
        )
    )
    figure.update_layout(
        template="plotly_dark",
        title="Poisson input-spike raster",
        xaxis_title="Timestep",
        yaxis_title="Input pixel",
    )
    metrics = {
        "true_label": int(row["label"]),
        "spiking_prediction": int(spike_logits.argmax(1)),
        "dense_prediction": int(dense_logits.argmax(1)),
        "hidden_spike_rate": float(spike_rate),
        "verified_clean_accuracy": REPORT["results"]["spiking_lif"]["clean"][
            "accuracy"
        ],
    }
    return image, figure, metrics


with gr.Blocks(title="Spike Pocket") as demo:
    gr.Markdown(
        "# Spike Pocket\n"
        "Inspect Poisson event encoding and a surrogate-gradient leaky-integrate-"
        "and-fire classifier beside its parameter-matched dense control."
    )
    with gr.Row():
        index = gr.Slider(0, len(FRAME) - 1, value=12, step=1, label="Test digit")
        timesteps = gr.Slider(8, 64, value=32, step=4, label="Timesteps")
        noise = gr.Slider(0, 0.4, value=0.0, step=0.02, label="Input noise")
    initial = inspect_spikes(12, 32, 0.0)
    with gr.Row():
        image = gr.Image(value=initial[0], label="Rate-encoded digit")
        raster = gr.Plot(value=initial[1], label="Spike raster")
    metrics = gr.JSON(value=initial[2], label="Neuromorphic readout")
    button = gr.Button("Simulate spikes", variant="primary")
    button.click(
        inspect_spikes,
        inputs=[index, timesteps, noise],
        outputs=[image, raster, metrics],
    )


if __name__ == "__main__":
    demo.launch()