mtg-draft-viz / src /eval /metrics.py
TimoBertram's picture
Upload src/eval/metrics.py with huggingface_hub
40c03eb verified
Raw
History Blame Contribute Delete
14.6 kB
import torch
import torch.nn.functional as F
import wandb
def evaluate(network, loader, device, prefix="eval", idx_to_card=None, n_log=500):
"""Single-pass eval: losses, BC accuracy/distance, Q accuracy/distance, Q calibration.
If idx_to_card is provided, logs a W&B table of Q-vs-BC disagreements and a
card ranking table sorted by predicted GIH win rate (descending)."""
network.eval()
bc_losses, q_losses, v_losses, gih_losses, play_losses = [], [], [], [], []
bc_correct = bc_dist = q_correct = q_dist = n = 0
q_terminal_all, actual_wr_all = [], []
disagreements = [] # rows for disagreement W&B table
card_gih_pred = {} # card_idx -> predicted gih value (first occurrence)
card_gih_target = {} # card_idx -> actual gih target (-1 if unknown)
with torch.no_grad(), torch.autocast(device_type='cuda'):
for batch in loader:
history_idx, pack_idx, pack_mask, seq_mask, wins, losses, user_wr, user_games, play_target, play_known = batch
history_cpu = history_idx # keep CPU copy for deck lookup
history_idx = history_idx.to(device, non_blocking=True)
pack_idx_gpu = pack_idx.to(device, non_blocking=True)
pack_mask = pack_mask.to(device, non_blocking=True)
seq_mask = seq_mask.to(device, non_blocking=True)
valid = ~seq_mask # [B, T]
last_t = (valid.long().sum(dim=1) - 1).clamp(min=0)
logits, q_values, values, play_logits, pick_play_logits, gih_pred, gih_target, gih_known = network(
history_idx, pack_idx_gpu, pack_mask, seq_mask)
B, T, P = logits.shape
# --- BC loss ---
log_probs = F.log_softmax(logits, dim=-1)[..., 0].masked_fill(~valid, 0.0)
bc_loss = -(log_probs * valid.float()).sum() / valid.float().sum().clamp(min=1)
# --- Q loss (IQL Bellman MSE, matching training) ---
wins_d = wins.to(device).float(); losses_d = losses.to(device).float()
next_is_end = torch.cat([seq_mask[:, 1:],
torch.ones(B, 1, dtype=torch.bool, device=device)], dim=1)
terminal_mask = valid & next_is_end
non_terminal_mask = valid & ~next_is_end
q_picked = q_values[..., 0]
q_loss = q_picked.new_zeros(())
if non_terminal_mask.any():
q_nt = torch.sigmoid(q_picked[:, :-1])
v_next = torch.sigmoid(values[:, 1:])
m = non_terminal_mask[:, :-1]
q_loss = q_loss + ((q_nt - v_next).pow(2) * m.float()).sum() / m.float().sum().clamp(min=1)
if terminal_mask.any():
true_wr = (wins_d / (wins_d + losses_d).clamp(min=1)).unsqueeze(1)
q_sig = torch.sigmoid(q_picked)
q_loss = q_loss + ((q_sig - true_wr).pow(2) * terminal_mask.float()).sum() / terminal_mask.float().sum().clamp(min=1)
# --- V loss (IQL expectile MSE, matching training) ---
tau = 0.7
v_sig = torch.sigmoid(values)
q_sig0 = torch.sigmoid(q_picked)
diff = q_sig0 - v_sig
weight = torch.where(diff >= 0,
diff.new_full(diff.shape, tau),
diff.new_full(diff.shape, 1.0 - tau))
v_loss = (weight * diff.pow(2) * valid.float()).sum() / valid.float().sum().clamp(min=1)
bc_losses.append(bc_loss.item())
q_losses.append(q_loss.item())
v_losses.append(v_loss.item())
# --- GIH loss ---
if gih_known.any():
gih_losses.append(F.mse_loss(
gih_pred[gih_known],
gih_target[gih_known].to(gih_pred.dtype),
).item())
# --- Playability loss (same [B,T,T] cross-product as training) ---
play_known_d = play_known.to(device)
play_target_d = play_target.to(device)
pick_play_mask = (
valid.unsqueeze(2)
& valid.unsqueeze(1)
& play_known_d.unsqueeze(1)
& ~torch.triu(torch.ones(T, T, dtype=torch.bool, device=device), diagonal=1).unsqueeze(0)
)
if pick_play_mask.any():
labels_exp = play_target_d.unsqueeze(1).expand(-1, T, -1)
play_losses.append(F.binary_cross_entropy_with_logits(
pick_play_logits[pick_play_mask],
labels_exp[pick_play_mask],
).item())
# --- Collect per-card GIH predictions for ranking ---
# gih_pred is context-free (computed before self-attention), so the same
# card always produces the same prediction in eval mode.
pmask_cpu = pack_mask.cpu() # [B, T, P]
pidx_cpu = pack_idx # [B, T, P] already CPU
gpred_cpu = gih_pred.cpu().float() # [B, T, P]
gtgt_cpu = gih_target.cpu().float() # [B, T, P]
valid_slots = pmask_cpu.view(-1)
cidxs = pidx_cpu.view(-1)[valid_slots].tolist()
preds = gpred_cpu.view(-1)[valid_slots].tolist()
tgts = gtgt_cpu.view(-1)[valid_slots].tolist()
for cidx, pred, tgt in zip(cidxs, preds, tgts):
if cidx not in card_gih_pred:
card_gih_pred[cidx] = pred
card_gih_target[cidx] = tgt
# --- V calibration: terminal V(s) vs actual win rate (V predicts absolute win rate) ---
v_term = torch.sigmoid(values).gather(1, last_t.unsqueeze(1)).squeeze(1).cpu().float()
n_games = wins + losses
has_games = n_games > 0
if has_games.any():
actual_wr = (wins / n_games.clamp(min=1))[has_games].float()
q_terminal_all.append(v_term[has_games])
actual_wr_all.append(actual_wr)
# --- accuracy & distance — flatten to valid steps only ---
valid_flat = valid.view(B * T)
logits_flat = logits.view(B * T, P)
q_flat = q_values.view(B * T, P)
pack_mask_flat = pack_mask.view(B * T, P)
# Advantage: sigmoid(Q) - sigmoid(V), both predict absolute win rate
q_sig_flat = torch.sigmoid(q_flat)
v_sig_flat = torch.sigmoid(values.view(B * T)).unsqueeze(-1)
q_adv_flat = (q_sig_flat - v_sig_flat).masked_fill(~pack_mask_flat, float('-inf'))
bc_rank = (logits_flat > logits_flat[:, :1]).sum(dim=-1)
bc_rank = bc_rank[valid_flat].float()
q_rank = (q_adv_flat > q_adv_flat[:, :1]).sum(dim=-1)
q_rank = q_rank[valid_flat].float()
bc_correct += (bc_rank == 0).sum().item()
bc_dist += bc_rank.sum().item()
q_correct += (q_rank == 0).sum().item()
q_dist += q_rank.sum().item()
n += bc_rank.numel()
# --- All pick examples (BC + Q picks for every valid step) ---
if idx_to_card is not None and len(disagreements) < n_log:
pack_idx_flat = pack_idx.view(B * T, P) # CPU
play_logits_flat = play_logits.view(B * T, P).cpu().float()
play_sig_flat_gpu = torch.sigmoid(play_logits.view(B * T, P))
q_combined_flat = (q_sig_flat * play_sig_flat_gpu).masked_fill(~pack_mask_flat, float('-inf'))
bc_top = logits_flat.argmax(dim=-1) # [B*T]
adv_flat = q_sig_flat - v_sig_flat # [B*T, P]
adv_min = adv_flat.masked_fill(~pack_mask_flat, float('inf')).min(dim=-1, keepdim=True).values
adv_shifted = (adv_flat - adv_min).masked_fill(~pack_mask_flat, float('-inf'))
q_adv_play = (adv_shifted * play_sig_flat_gpu).masked_fill(~pack_mask_flat, float('-inf'))
q_top = q_adv_play.argmax(dim=-1) # [B*T] shifted-adv × play
q_play_top = q_combined_flat.argmax(dim=-1) # [B*T] sigmoid(Q) × play (logged for comparison)
for pos in valid_flat.nonzero(as_tuple=True)[0].tolist():
if len(disagreements) >= n_log:
break
b_idx = pos // T
t_idx = pos % T
slots = pack_mask_flat[pos].cpu()
card_idxs = pack_idx_flat[pos]
names = [idx_to_card.get(card_idxs[j].item(), "?")
for j in range(P) if slots[j]]
human = idx_to_card.get(card_idxs[0].item(), "?")
bc_p = idx_to_card.get(card_idxs[bc_top[pos].item()].item(), "?")
q_p = idx_to_card.get(card_idxs[q_top[pos].item()].item(), "?")
q_play_p = idx_to_card.get(card_idxs[q_play_top[pos].item()].item(), "?")
deck_idxs = history_cpu[b_idx, :t_idx]
deck_str = ", ".join(idx_to_card.get(i.item(), "?") for i in deck_idxs) \
if t_idx > 0 else "(empty)"
q_val_human = torch.sigmoid(q_flat[pos][0]).item()
q_val_q_pick = torch.sigmoid(q_flat[pos][q_top[pos].item()]).item()
adv_human = q_adv_flat[pos][0].item()
adv_q_pick = q_adv_flat[pos][q_top[pos].item()].item()
play_human = torch.sigmoid(play_logits_flat[pos][0]).item()
play_bc = torch.sigmoid(play_logits_flat[pos][bc_top[pos].item()]).item()
play_q = torch.sigmoid(play_logits_flat[pos][q_top[pos].item()]).item()
play_q_play = torch.sigmoid(play_logits_flat[pos][q_play_top[pos].item()]).item()
q_val_q_play = torch.sigmoid(q_flat[pos][q_play_top[pos].item()]).item()
disagreements.append([
human, bc_p, q_p, q_play_p,
human == bc_p, human == q_p,
round(q_val_human, 3),
round(q_val_q_pick, 3),
round(q_val_q_play, 3),
round(adv_human, 4),
round(adv_q_pick, 4),
round(play_human, 3),
round(play_bc, 3),
round(play_q, 3),
round(play_q_play, 3),
", ".join(names),
deck_str,
])
if q_terminal_all:
q_terminal_all = torch.cat(q_terminal_all)
actual_wr_all = torch.cat(actual_wr_all)
q_calib = torch.corrcoef(torch.stack([q_terminal_all, actual_wr_all]))[0, 1].item()
else:
q_calib = float("nan")
# Spearman rank correlation between predicted GIH and actual GIH
known = [(card_gih_pred[c], card_gih_target[c])
for c in card_gih_pred if card_gih_target.get(c, -1.0) >= 0]
if len(known) >= 2:
pred_t = torch.tensor([x[0] for x in known])
actual_t = torch.tensor([x[1] for x in known])
pred_ranks = pred_t.argsort().argsort().float()
actual_ranks = actual_t.argsort().argsort().float()
gih_rank_corr = torch.corrcoef(torch.stack([pred_ranks, actual_ranks]))[0, 1].item()
else:
gih_rank_corr = float("nan")
metrics = {
f"{prefix}/bc_loss": sum(bc_losses) / len(bc_losses),
f"{prefix}/q_loss": sum(q_losses) / len(q_losses),
f"{prefix}/v_loss": sum(v_losses) / len(v_losses),
f"{prefix}/gih_loss": sum(gih_losses) / len(gih_losses) if gih_losses else float("nan"),
f"{prefix}/play_loss": sum(play_losses) / len(play_losses) if play_losses else float("nan"),
f"{prefix}/bc_acc": bc_correct / n if n else float("nan"),
f"{prefix}/bc_dist": bc_dist / n if n else float("nan"),
f"{prefix}/q_acc": q_correct / n if n else float("nan"),
f"{prefix}/q_dist": q_dist / n if n else float("nan"),
f"{prefix}/q_calibration": q_calib,
f"{prefix}/gih_rank_corr": gih_rank_corr,
}
# Log tables directly (no step) to avoid corrupting the scalar step counter
if wandb.run is not None and idx_to_card is not None:
tables = {}
if disagreements:
tables[f"{prefix}/pick_disagreements"] = wandb.Table(
columns=["human_pick", "bc_pick", "q_pick", "q_play_pick",
"bc_correct", "q_correct",
"q_human", "q_q_pick", "q_play_pick_val",
"adv_human", "adv_q_pick",
"play_human", "play_bc", "play_q", "play_q_play",
"pack", "deck"],
data=disagreements,
)
if card_gih_pred:
known_cidxs = [c for c, t in card_gih_target.items() if t >= 0 and c in card_gih_pred]
actual_order = sorted(known_cidxs, key=lambda c: card_gih_target[c], reverse=True)
actual_rank_map = {c: r for r, c in enumerate(actual_order, 1)}
pred_order = sorted(known_cidxs, key=lambda c: card_gih_pred[c], reverse=True)
pred_rank_map = {c: r for r, c in enumerate(pred_order, 1)}
rows = []
for cidx in actual_order:
name = idx_to_card.get(cidx, "?")
rows.append([actual_rank_map[cidx], pred_rank_map[cidx], name,
round(card_gih_pred[cidx], 4), round(card_gih_target[cidx], 4),
round(card_gih_pred[cidx] - card_gih_target[cidx], 4)])
for cidx, pred in card_gih_pred.items():
if card_gih_target.get(cidx, -1.0) < 0:
rows.append([None, pred_rank_map.get(cidx), idx_to_card.get(cidx, "?"),
round(pred, 4), None, None])
tables[f"{prefix}/card_ranking"] = wandb.Table(
columns=["actual_rank", "pred_rank", "card", "pred_gih", "actual_gih", "delta"],
data=rows,
)
metrics.update(tables)
return metrics