Spaces:
Sleeping
Sleeping
| 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 | |