ARotting's picture
Publish Interactive learned-versus-reference OS transition comparison
3425207 verified
Raw
History Blame Contribute Delete
2.8 kB
from __future__ import annotations
from pathlib import Path
import gradio as gr
import torch
from model import NeuralModelMachine
from safetensors.torch import load_file
from schema import ACTIONS, CAPABILITY_NAMES, STATE_NAMES, reference_transition
ARTIFACT_DIR = (
Path(__file__).resolve().parent / "artifacts" / "kernelmind-model-machine"
)
MODEL = NeuralModelMachine()
MODEL.load_state_dict(load_file(ARTIFACT_DIR / "model_machine.safetensors"))
MODEL.eval()
@torch.inference_mode()
def run_transition(
action: str,
active_state: list[str],
active_capabilities: list[str],
) -> tuple[dict, dict]:
state = [int(name in active_state) for name in STATE_NAMES]
capabilities = [int(name in active_capabilities) for name in CAPABILITY_NAMES]
action_index = ACTIONS.index(action)
predicted, blocked_probability = MODEL.transition(
torch.tensor([state], dtype=torch.float32),
torch.tensor([capabilities], dtype=torch.float32),
torch.tensor([action_index]),
)
reference_state, reference_blocked = reference_transition(
state, capabilities, action
)
learned_state = predicted[0].int().tolist()
learned = {
"next_state": dict(zip(STATE_NAMES, learned_state, strict=True)),
"blocked_probability": float(blocked_probability[0]),
}
reference = {
"next_state": dict(zip(STATE_NAMES, reference_state, strict=True)),
"blocked": bool(reference_blocked),
"exact_match": learned_state == reference_state
and int(blocked_probability[0] >= 0.5) == reference_blocked,
}
return learned, reference
with gr.Blocks(title="KernelMind Neural Model Machine") as demo:
gr.Markdown(
"# KernelMind Neural Model Machine\n"
"A trained world model predicts how an OS action changes virtual machine "
"state. Compare its transition directly with the deterministic reference."
)
action = gr.Dropdown(ACTIONS, value="WRITE_BACKUP", label="Action")
state = gr.CheckboxGroup(
STATE_NAMES,
value=["file_exists", "service_running"],
label="Current state bits",
)
capabilities = gr.CheckboxGroup(
CAPABILITY_NAMES,
value=["network"],
label="Capabilities",
)
initial = run_transition(
"WRITE_BACKUP", ["file_exists", "service_running"], ["network"]
)
with gr.Row():
learned = gr.JSON(value=initial[0], label="Learned transition")
reference = gr.JSON(value=initial[1], label="Reference transition")
button = gr.Button("Advance the model machine", variant="primary")
button.click(
run_transition,
inputs=[action, state, capabilities],
outputs=[learned, reference],
)
if __name__ == "__main__":
demo.launch()