import sys, os
sys.path.insert(0, os.path.dirname(__file__))
import streamlit as st
import torch
import torch.nn.functional as F
import numpy as np
import pickle
import lmdb
import json
from urllib.parse import quote
from huggingface_hub import snapshot_download
from src.training import models
from src.utils import utils
st.set_page_config(page_title="Draft Visualizer", layout="wide")
MAX_CHOICES = 15
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
BASICS = {"Mountain", "Forest", "Swamp", "Island", "Plains"}
HF_MODEL_REPO = os.environ.get("HF_MODEL_REPO", "")
HF_DATA_REPO = os.environ.get("HF_DATA_REPO", "")
# ---------------------------------------------------------------------------
# Download HF repos once — returns local paths identical to server layout
# ---------------------------------------------------------------------------
@st.cache_resource(show_spinner="Downloading model files…")
def get_model_dir():
return snapshot_download(HF_MODEL_REPO)
@st.cache_resource(show_spinner="Downloading draft data…")
def get_data_dir():
return snapshot_download(HF_DATA_REPO, repo_type="dataset")
def scryfall_image_url(card_name: str) -> str:
return f"https://api.scryfall.com/cards/named?fuzzy={quote(card_name)}&format=image&version=normal"
# ---------------------------------------------------------------------------
# Model + vocab
# ---------------------------------------------------------------------------
@st.cache_resource(show_spinner="Loading model…")
def load_model_and_vocab(checkpoint_path: str, model_dir: str):
config = utils.load_config(os.path.join(model_dir, "config.yaml"))
config["embedding_path"] = os.path.join(model_dir, "card_encodings.pt")
embedding_dict = utils.get_embedding_dict(config["embedding_path"], add_nontransformed=True)
all_vecs = np.array(list(embedding_dict.values()))
mean, std = all_vecs.mean(axis=0), all_vecs.std(axis=0)
std[std == 0] = 1
cards = sorted(embedding_dict.keys())
card_to_idx = {c: i for i, c in enumerate(cards)}
idx_to_card = {i: c for c, i in card_to_idx.items()}
embedding_matrix = torch.tensor(
np.stack([(embedding_dict[c] - mean) / std for c in cards]),
dtype=torch.float32,
)
gih_wr_matrix = torch.full((len(cards),), -1.0)
gih_folder = os.path.join(model_dir, "gih_wr")
if os.path.isdir(gih_folder):
for fname in os.listdir(gih_folder):
if not fname.endswith("_gih.json"):
continue
with open(os.path.join(gih_folder, fname)) as f:
for entry in json.load(f):
wr = entry.get("ever_drawn_win_rate")
if wr is None:
continue
if isinstance(wr, str):
wr = float(wr.rstrip("%")) / 100
name = utils.normalize_card_name(entry["name"])
if name in card_to_idx:
gih_wr_matrix[card_to_idx[name]] = float(wr)
network = models.DecisionDraftTransformer(
**config, embedding_matrix=embedding_matrix, gih_wr_matrix=gih_wr_matrix
)
state = torch.load(checkpoint_path, map_location="cpu")
network.load_state_dict(state)
network.to(DEVICE)
network.eval()
return network, card_to_idx, idx_to_card, config
# ---------------------------------------------------------------------------
# Data loading
# ---------------------------------------------------------------------------
@st.cache_data
def load_draft(lmdb_path: str, draft_idx: int):
env = lmdb.open(lmdb_path, readonly=True, lock=False)
with env.begin() as txn:
cur = txn.cursor()
keys = [bytes(k) for k, _ in cur if k != b"__len__"]
key = keys[draft_idx % len(keys)]
with env.begin() as txn:
data = pickle.loads(txn.get(key))
env.close()
sequence, in_maindeck, wins, losses, u_g, u_wr = data
return sequence, in_maindeck, int(wins), int(losses), int(u_g), float(u_wr), len(keys)
def build_tensors(sequence, card_to_idx):
T = len(sequence)
history_idx = torch.zeros(1, T, dtype=torch.long)
pack_idx = torch.zeros(1, T, MAX_CHOICES, dtype=torch.long)
pack_mask = torch.zeros(1, T, MAX_CHOICES, dtype=torch.bool)
seq_mask = torch.zeros(1, T, dtype=torch.bool)
for t, pack_cards in enumerate(sequence):
history_idx[0, t] = card_to_idx.get(utils.normalize_card_name(pack_cards[0]), 0)
for j, card in enumerate(pack_cards[:MAX_CHOICES]):
pack_idx[0, t, j] = card_to_idx.get(utils.normalize_card_name(card), 0)
pack_mask[0, t, j] = True
return (history_idx.to(DEVICE), pack_idx.to(DEVICE),
pack_mask.to(DEVICE), seq_mask.to(DEVICE))
@torch.no_grad()
def run_model(network, sequence, card_to_idx, skill_target=0.60):
history_idx, pack_idx, pack_mask, seq_mask = build_tensors(sequence, card_to_idx)
B = history_idx.shape[0]
outcome = torch.full((B,), skill_target, device=DEVICE)
player_wr = torch.full((B,), skill_target, device=DEVICE)
logits, play_logits, pick_play_logits, _, _, _ = network(history_idx, pack_idx, pack_mask, seq_mask, outcome, player_wr)
bc_probs = F.softmax(logits[0], dim=-1).cpu()
play_sig = torch.sigmoid(play_logits[0]).cpu()
pick_play_sig = torch.sigmoid(pick_play_logits[0]).cpu()
return bc_probs, play_sig, pick_play_sig, pack_mask[0].cpu(), pack_idx[0].cpu()
# ---------------------------------------------------------------------------
# Sidebar
# ---------------------------------------------------------------------------
if not HF_MODEL_REPO or not HF_DATA_REPO:
st.error("Set HF_MODEL_REPO and HF_DATA_REPO as Space secrets.")
st.stop()
model_dir = get_model_dir()
data_dir = get_data_dir()
with st.sidebar:
st.header("Config")
run_folders = sorted(
[d for d in os.listdir(model_dir) if os.path.isdir(os.path.join(model_dir, d))
and any(f.endswith(".pt") for f in os.listdir(os.path.join(model_dir, d)))],
reverse=True,
)
selected_run = st.selectbox("Run", options=run_folders)
run_dir = os.path.join(model_dir, selected_run) if selected_run else ""
ckpt_files = sorted(
[f for f in os.listdir(run_dir) if f.endswith(".pt")]
if run_dir else [],
reverse=True,
)
selected_ckpt = st.selectbox("Checkpoint", options=ckpt_files)
checkpoint_path = os.path.join(run_dir, selected_ckpt) if selected_ckpt else ""
sets_with_lmdb = sorted([
s for s in os.listdir(data_dir)
if os.path.exists(os.path.join(data_dir, s, "test.lmdb"))
])
selected_set = st.selectbox("Set", options=sets_with_lmdb)
lmdb_split = st.radio("Split", ["test", "train"], horizontal=True)
lmdb_path = os.path.join(data_dir, selected_set, f"{lmdb_split}.lmdb") if selected_set else ""
skill_target = st.slider("Skill level", min_value=0.0, max_value=1.0, value=0.60, step=0.05)
show_images = st.toggle("Show card images", value=True)
if not (checkpoint_path and os.path.exists(lmdb_path)):
st.warning("Select a run, checkpoint, and set.")
st.stop()
network, card_to_idx, idx_to_card, config = load_model_and_vocab(checkpoint_path, model_dir)
st.success(f"Model on {DEVICE}")
draft_idx = st.number_input("Draft index", min_value=0, value=0, step=1)
# ---------------------------------------------------------------------------
# Load draft + run model
# ---------------------------------------------------------------------------
sequence, in_maindeck, wins, losses, u_g, u_wr, n_drafts = load_draft(lmdb_path, draft_idx)
st.sidebar.caption(f"{n_drafts} drafts available")
T = len(sequence)
bc_probs, play_sig, pick_play_sig, pack_mask, pack_idx = run_model(network, sequence, card_to_idx, skill_target)
# ---------------------------------------------------------------------------
# Header
# ---------------------------------------------------------------------------
st.title("MTG Draft Visualizer")
c1, c2, c3, c4, c5 = st.columns(5)
c1.metric("Draft result", f"{wins}W – {losses}L")
c2.metric("Player WR", f"{u_wr:.1%}" if u_wr > 0 else "unknown")
c3.metric("Player games", f"{u_g:,}" if u_g > 0 else "unknown")
c4.metric("Total picks", T)
c5.metric("Set", selected_set)
st.divider()
# ---------------------------------------------------------------------------
# Pick slider
# ---------------------------------------------------------------------------
if "pick_slider" not in st.session_state:
st.session_state.pick_slider = 0
st.session_state.pick_slider = max(0, min(T - 1, st.session_state.pick_slider))
col_prev, col_slider, col_next = st.columns([1, 10, 1])
with col_prev:
if st.button("◀", use_container_width=True):
st.session_state.pick_slider = max(0, st.session_state.pick_slider - 1)
with col_next:
if st.button("▶", use_container_width=True):
st.session_state.pick_slider = min(T - 1, st.session_state.pick_slider + 1)
with col_slider:
pick = st.slider("Pick", min_value=0, max_value=T - 1, key="pick_slider")
pack_num = pick // 15 + 1
pick_in_pack = pick % 15 + 1
pack_cards = sequence[pick]
human_name = utils.normalize_card_name(pack_cards[0])
n_cards = int(pack_mask[pick].sum().item())
bc_top = bc_probs[pick].masked_fill(~pack_mask[pick], float('-inf')).argmax().item()
play_top = play_sig[pick].masked_fill(~pack_mask[pick], float('-inf')).argmax().item()
bc_name = idx_to_card.get(pack_idx[pick, bc_top].item(), "?")
play_name = idx_to_card.get(pack_idx[pick, play_top].item(), "?")
st.subheader(f"Pack {pack_num}, Pick {pick_in_pack} (step {pick}) — skill = {skill_target:.2f}")
m1, m2, m3 = st.columns(3)
m1.metric("👤 Human pick", human_name)
m2.metric("🤖 BC pick", bc_name, delta="✓" if bc_name == human_name else "✗")
m3.metric("▶ Top play rate", play_name, delta="✓" if play_name == human_name else "✗")
# ---------------------------------------------------------------------------
# Card image grid
# ---------------------------------------------------------------------------
if show_images:
valid_slots = [j for j in range(MAX_CHOICES) if pack_mask[pick, j]]
def border_color(name):
is_human = name == human_name
is_bc = name == bc_name
is_play = name == play_name
if is_human and is_bc and is_play: return "#00ff00"
if is_human and is_bc: return "#4488ff"
if is_human and is_play: return "#ff8800"
if is_human: return "#aaccff"
if is_play: return "#ff4400"
if is_bc: return "#aa66ff"
return "#333333"
cols = st.columns(5)
for idx_in_row, j in enumerate(valid_slots):
cidx = pack_idx[pick, j].item()
name = idx_to_card.get(cidx, "?")
bc_p = bc_probs[pick, j].item()
play = play_sig[pick, j].item()
color = border_color(name)
col = cols[idx_in_row % 5]
with col:
url = scryfall_image_url(name)
st.markdown(
f'
',
unsafe_allow_html=True,
)
st.caption(f"bc {bc_p:.3f} | play {play:.2f}")
st.markdown(
"🟢 all agree 🔵 human+BC 🟠 human+play "
"🔴 play only 🟣 BC only 🩵 human only",
unsafe_allow_html=True,
)
st.divider()
# ---------------------------------------------------------------------------
# Deck so far
# ---------------------------------------------------------------------------
with st.expander(f"Deck so far ({pick} cards)", expanded=False):
if pick == 0:
st.caption("No picks yet.")
else:
deck = [utils.normalize_card_name(sequence[t][0]) for t in range(pick)]
cols = st.columns(5)
for i, name in enumerate(deck):
pred_play = pick_play_sig[pick - 1, i].item() if pick > 0 else 0.0
gt_play = in_maindeck[i] if in_maindeck is not None else None
gt_str = (" ✓" if gt_play == 1.0 else " ✗") if gt_play is not None else ""
with cols[i % 5]:
url = scryfall_image_url(name)
st.markdown(
f'
',
unsafe_allow_html=True,
)
st.caption(f"#{i+1} play {pred_play:.2f}{gt_str}")