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