import json import math import numpy as np import torch import torch.nn as nn import torch.nn.functional as F import wandb from functools import partial from torch.utils.data import DataLoader, ConcatDataset import tqdm from torch.utils.data.distributed import DistributedSampler from torch.nn.parallel import DistributedDataParallel as DDP from torch.distributed import destroy_process_group import os from src.training import models from src.utils import utils from src.utils import ddp from src.eval.metrics import evaluate as pick_metrics from src.preprocessing.preprocess_pipeline import ensure_set_ready MAX_CHOICES = 15 card_to_idx = None # shared with workers: card_name -> int index def wlog(data: dict, step: int | None = None): if wandb.run is not None: wandb.log(data, step=step) def worker_init_fn(shared_card_to_idx, worker_id): global card_to_idx card_to_idx = shared_card_to_idx def collate_fn(batch): batch_size = len(batch) max_T = max(len(seq) for seq, *_ in batch) history_idx = torch.zeros(batch_size, max_T, dtype=torch.long) pack_idx = torch.zeros(batch_size, max_T, MAX_CHOICES, dtype=torch.long) pack_mask = torch.zeros(batch_size, max_T, MAX_CHOICES, dtype=torch.bool) seq_mask = torch.ones(batch_size, max_T, dtype=torch.bool) # True = padding wins_t = torch.zeros(batch_size) losses_t = torch.zeros(batch_size) user_wr_t = torch.zeros(batch_size) user_games_t = torch.zeros(batch_size) play_target = torch.zeros(batch_size, max_T) # 1.0 = in maindeck play_known = torch.zeros(batch_size, max_T, dtype=torch.bool) # False for old-format data for i, item in enumerate(batch): if len(item) == 6: # new format: (sequence, in_maindeck, wins, losses, u_g, u_wr) sequence, in_maindeck, wins, losses, u_g, u_wr = item else: # old format: (sequence, wins, losses, u_g, u_wr) sequence, wins, losses, u_g, u_wr = item in_maindeck = None T = len(sequence) seq_mask[i, :T] = False for t, pack_cards in enumerate(sequence): history_idx[i, t] = card_to_idx.get(utils.normalize_card_name(pack_cards[0]), 0) for j, card in enumerate(pack_cards[:MAX_CHOICES]): pack_idx[i, t, j] = card_to_idx.get(utils.normalize_card_name(card), 0) pack_mask[i, t, j] = True if in_maindeck is not None: play_target[i, :T] = torch.tensor(in_maindeck[:T], dtype=torch.float) play_known[i, :T] = True wins_t[i] = wins losses_t[i] = losses user_wr_t[i] = u_wr user_games_t[i] = u_g return history_idx, pack_idx, pack_mask, seq_mask, wins_t, losses_t, user_wr_t, user_games_t, play_target, play_known def _v_loss_iql(values, q_picked, valid, tau=0.7): """IQL expectile regression: V(s) ← τ-expectile of Q(s, a_human). τ > 0.5 pushes V toward the upper end of Q so that good picks produce positive advantages Q(s,a) - V(s).""" v_sig = torch.sigmoid(values) q_sig = torch.sigmoid(q_picked.detach()) diff = q_sig - v_sig # positive when Q > V weight = torch.where(diff >= 0, diff.new_full(diff.shape, tau), diff.new_full(diff.shape, 1.0 - tau)) return (weight * diff.pow(2) * valid.float()).sum() / valid.float().sum().clamp(min=1) def _q_loss_iql(q_picked, values, wins, losses, device, valid, seq_mask): """IQL Bellman backup for Q. Non-terminal steps: MSE( sigmoid(Q(t,0)), sigmoid(V(t+1)).detach() ) Terminal step: MSE( sigmoid(Q(T-1,0)), wins/(wins+losses) ) This forces Q and V onto the same scale without querying OOD actions.""" B = q_picked.shape[0] 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 # [B, T] non_terminal_mask = valid & ~next_is_end # [B, T] total = q_picked.new_zeros(()) if non_terminal_mask.any(): q_nt = torch.sigmoid(q_picked[:, :-1]) # [B, T-1] v_next = torch.sigmoid(values[:, 1:]).detach() # [B, T-1] m = non_terminal_mask[:, :-1] total = total + ((q_nt - v_next).pow(2) * m.float()).sum() / m.float().sum().clamp(min=1) if terminal_mask.any(): W, L = wins.to(device).float(), losses.to(device).float() true_wr = (W / (W + L).clamp(min=1)).unsqueeze(1) q_sig = torch.sigmoid(q_picked) total = total + ((q_sig - true_wr).pow(2) * terminal_mask.float()).sum() / terminal_mask.float().sum().clamp(min=1) return total def _advantage_weights(q_values, values, valid, beta=2.0): q_wr = torch.sigmoid(q_values[..., 0]).masked_fill(~valid, 0.0) v_wr = torch.sigmoid(values).masked_fill(~valid, 0.0) advantage = (q_wr - v_wr).detach() weights = torch.exp((beta * advantage).clamp(-5, 5)) * valid.float() n_valid = valid.float().sum().clamp(min=1) return weights / (weights.sum() / n_valid).clamp(min=1e-8) def training_step(network, batch, optimizer, scheduler, device, scaler, lam_gih=1.0, lam_play=1.0, tau=0.7): history_idx, pack_idx, pack_mask, seq_mask, wins, losses, user_wr, user_games, play_target, play_known = batch optimizer.zero_grad() history_idx = history_idx.to(device, non_blocking=True) pack_idx = 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 # Skill weighting: upweight picks from high win-rate players. # u_wr=0 means unknown — treat as average (replaced by batch mean of known). skill_w = user_wr.to(device) # [B] known = skill_w[skill_w > 0] fallback = known.mean() if known.numel() > 0 else skill_w.new_tensor(0.5) skill_w = torch.where(skill_w > 0, skill_w, fallback) skill_w = (skill_w / skill_w.mean().clamp(min=1e-8)).unsqueeze(1) # [B, 1] with torch.autocast(device_type='cuda'): logits, q_values, values, play_logits, pick_play_logits, gih_pred, gih_target, gih_known = network( history_idx, pack_idx, pack_mask, seq_mask) # Advantage weights: sigmoid(Q) - sigmoid(V) weights = _advantage_weights(q_values, values, valid) log_probs = F.log_softmax(logits, dim=-1)[..., 0].masked_fill(~valid, 0.0) bc_loss = -(weights * skill_w * log_probs * valid.float()).sum() / valid.float().sum().clamp(min=1) q_loss = _q_loss_iql(q_values[..., 0], values, wins, losses, device, valid, seq_mask) v_loss = _v_loss_iql(values, q_values[..., 0], valid, tau=tau) # Playability loss: BCE over all (t, s) pairs where s <= t and pick s has a label. # pick_play_logits[b, t, s] = P(pick_s in maindeck | deck context at step t). # Evaluating past picks at every future step gives ~T/2 × more signal per draft. T = seq_mask.shape[1] play_known_d = play_known.to(device) # [B, T] play_target_d = play_target.to(device) # [B, T] # Valid entry: step t not padding, pick s not padding, pick s has a label pick_play_mask = ( (~seq_mask).unsqueeze(2) # t valid [B, T, 1] & (~seq_mask).unsqueeze(1) # s valid [B, 1, T] & play_known_d.unsqueeze(1) # s has label [B, 1, T] & ~torch.triu(torch.ones(T, T, dtype=torch.bool, device=device), diagonal=1).unsqueeze(0) ) # [B, T, T], lower triangular including diagonal if pick_play_mask.any(): labels_exp = play_target_d.unsqueeze(1).expand(-1, T, -1) # [B, T, T] play_loss = F.binary_cross_entropy_with_logits( pick_play_logits[pick_play_mask], labels_exp[pick_play_mask]) else: play_loss = pick_play_logits.sum() * 0 # Auxiliary GIH loss: MSE on cards with known win-rate targets if gih_known.any(): gih_loss = F.mse_loss(gih_pred[gih_known], gih_target[gih_known].to(gih_pred.dtype)) else: gih_loss = gih_pred.sum() * 0 # zero, keeps grad graph loss = bc_loss + q_loss + v_loss + lam_play * play_loss + lam_gih * gih_loss play_diag = torch.sigmoid(pick_play_logits.diagonal(dim1=1, dim2=2)) # [B, T] _play_w_std = play_diag[~seq_mask].std().item() q_sig_all = torch.sigmoid(q_values) # [B, T, P] v_sig_all = torch.sigmoid(values).unsqueeze(-1) # [B, T, 1] adv_all = (q_sig_all - v_sig_all).masked_fill(~pack_mask, 0.0) _adv_std = adv_all[pack_mask].std().item() # NaN diagnostic: padding steps have -inf logits by design and are masked out; # NaN here means a real forward-pass bug in a VALID step. _bc, _q, _v, _play, _gih = (bc_loss.item(), q_loss.item(), v_loss.item(), play_loss.item(), gih_loss.item()) if any(math.isnan(x) for x in (_bc, _q, _v, _play, _gih)): nan_logits_valid = torch.isnan(logits[valid]).any().item() if valid.any() else False nan_qval_valid = torch.isnan(q_values[valid]).any().item() if valid.any() else False nan_packs_valid = torch.isnan(q_values[pack_mask]).any().item() if pack_mask.any() else False print(f"[NaN] bc={_bc:.4f} q={_q:.4f} v={_v:.4f} play={_play:.4f} gih={_gih:.4f} " f"| logits(valid)={nan_logits_valid} q(valid)={nan_qval_valid} q(pack_mask)={nan_packs_valid}") scaler.scale(loss).backward() scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(network.parameters(), max_norm=1.0) scaler.step(optimizer) scaler.update() scheduler.step() return _bc, _q, _v, _play, _gih, _play_w_std, _adv_std def train(rank, local_rank, network, train_loader, eval_loader, eval_test_loader, config, use_ddp, idx_to_card=None): is_master = rank == 0 global_step = 0 if use_ddp: network = DDP(network, device_ids=[local_rank], find_unused_parameters=False, broadcast_buffers=False) scaler = torch.amp.GradScaler('cuda') max_epochs = config['max_epochs'] lr = config['lr'] excluded = [p for n, p in network.named_parameters() if "gamma" in n] optimizer = torch.optim.AdamW([ {"params": [p for n, p in network.named_parameters() if "gamma" not in n], "weight_decay": 1e-4}, {"params": excluded, "weight_decay": 0.0}, ], lr=lr) warmup_steps = config['warmup_steps'] total_steps = max_epochs * warmup_steps def lr_lambda(step): if step < warmup_steps: return 0.01 + 0.99 * step / warmup_steps progress = (step - warmup_steps) / max(1, total_steps - warmup_steps) return 0.01 + 0.5 * 0.99 * (1 + math.cos(math.pi * progress)) scheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda) for epoch in range(1,max_epochs+1): total_iterations = len(train_loader) if use_ddp: train_loader.sampler.set_epoch(epoch) eval_loader.sampler.set_epoch(epoch) if eval_test_loader is not None: eval_test_loader.sampler.set_epoch(epoch) # Training loop network.train() if is_master: bar = tqdm.tqdm(enumerate(train_loader), total = total_iterations, mininterval = 1, desc = 'Training') else: bar = enumerate(train_loader) lam_gih = config.get('lambda_gih', 1.0) lam_play = config.get('lambda_play', 1.0) tau = config.get('tau', 0.7) for i, batch in bar: bc_loss, q_loss, v_loss, play_loss, gih_loss, play_w_std, adv_std = training_step( network, batch, optimizer, scheduler, device=torch.device(f'cuda:{local_rank}'), scaler=scaler, lam_gih=lam_gih, lam_play=lam_play, tau=tau) if is_master: wlog({ "train/bc_loss": bc_loss, "train/q_loss": q_loss, "train/v_loss": v_loss, "train/play_loss": play_loss, "train/gih_loss": gih_loss, "train/lr": scheduler.get_last_lr()[0], "train/play_weight_std": play_w_std, "train/adv_std": adv_std, }, step=global_step) global_step += 1 if is_master: metrics = pick_metrics(network, eval_loader, device=torch.device(f'cuda:{local_rank}'), prefix="eval", idx_to_card=idx_to_card) wlog(metrics, step=global_step) if eval_test_loader is not None: test_metrics = pick_metrics(network, eval_test_loader, device=torch.device(f'cuda:{local_rank}'), prefix="eval_test", idx_to_card=idx_to_card) wlog(test_metrics, step=global_step) global_step += 1 raw = network.module if isinstance(network, DDP) else network run_id = config.get('_run_id', config.get('net_name', 'run')) ckpt_dir = os.path.join(config.get('checkpoint_dir', 'checkpoints'), run_id) os.makedirs(ckpt_dir, exist_ok=True) torch.save(raw.state_dict(), os.path.join(ckpt_dir, f'epoch{epoch}.pt')) if use_ddp: destroy_process_group() def main(card_to_idx_shared, embedding_matrix, gih_wr_matrix, train_sets, test_sets, config, idx_to_card=None): global card_to_idx card_to_idx = card_to_idx_shared init_fn = partial(worker_init_fn, card_to_idx_shared) rank, world_size, local_rank, use_ddp = ddp.ddp_setup_from_env() if use_ddp: import torch.distributed as dist dist.barrier() is_master = rank == 0 if not is_master: os.environ["WANDB_MODE"] = "disabled" run_id = None if is_master: wandb.init(entity="tibert97", project="Drafting IL", config=config, ) run_id = f"{wandb.run.name}-{wandb.run.id}" # e.g. "golden-river-42-3ix738nu" config['_run_id'] = run_id batch_size = config['batch_size'] num_workers = config['num_workers'] if torch.cuda.is_available() else 0 pin_memory = config['pin_memory'] persistent_workers = config['persistent_workers'] and torch.cuda.is_available() prefetch_factor = config['prefetch_factor'] # Training data: train_sets only train_data = ConcatDataset([models.LMDBDataset(db_path=f'{config["super_folder"]}/{s}/train.lmdb') for s in train_sets]) train_sampler = DistributedSampler(train_data, num_replicas=world_size, rank=rank, shuffle=True) if use_ddp else None train_loader = DataLoader(train_data, batch_size=batch_size, num_workers=num_workers, pin_memory=pin_memory, collate_fn=collate_fn, persistent_workers=persistent_workers, prefetch_factor=prefetch_factor, sampler=train_sampler, shuffle=(train_sampler is None), worker_init_fn=init_fn) # Eval on train_sets (in-distribution) eval_data = ConcatDataset([models.LMDBDataset(db_path=f'{config["super_folder"]}/{s}/test.lmdb') for s in train_sets]) eval_sampler = DistributedSampler(eval_data, num_replicas=world_size, rank=rank, shuffle=False) if use_ddp else None eval_loader = DataLoader(eval_data, batch_size=batch_size, num_workers=num_workers, pin_memory=pin_memory, collate_fn=collate_fn, persistent_workers=persistent_workers, prefetch_factor=prefetch_factor, sampler=eval_sampler, shuffle=False, worker_init_fn=init_fn) # Eval on test_sets (held-out, never trained on) eval_test_loader = None if test_sets: eval_test_data = ConcatDataset([models.LMDBDataset(db_path=f'{config["super_folder"]}/{s}/test.lmdb') for s in test_sets]) eval_test_sampler = DistributedSampler(eval_test_data, num_replicas=world_size, rank=rank, shuffle=False) if use_ddp else None eval_test_loader = DataLoader(eval_test_data, batch_size=batch_size, num_workers=num_workers, pin_memory=pin_memory, collate_fn=collate_fn, persistent_workers=persistent_workers, prefetch_factor=prefetch_factor, sampler=eval_test_sampler, shuffle=False, worker_init_fn=init_fn) config['warmup_steps'] = config.get('warmup_epochs', 1) * len(train_loader) network = models.DraftTransformer(**config, embedding_matrix=embedding_matrix, gih_wr_matrix=gih_wr_matrix).cuda(local_rank) train(rank=rank, local_rank=local_rank, network=network, train_loader=train_loader, eval_loader=eval_loader, eval_test_loader=eval_test_loader, config=config, use_ddp=use_ddp, idx_to_card=idx_to_card) if __name__ == "__main__": import argparse parser = argparse.ArgumentParser() parser.add_argument('--lr', type=float, default=None) parser.add_argument('--batch_size', type=int, default=None) parser.add_argument('--dropout', type=float, default=None) parser.add_argument('--lambda_gih', type=float, default=None) parser.add_argument('--max_epochs', type=int, default=None) args = parser.parse_args() config = utils.load_config('src/configs/config.yaml') for key, val in vars(args).items(): if val is not None: config[key] = val train_sets = config['train_sets'] test_sets = config.get('test_sets', []) embedding_path = config['embedding_path'] print(f'Train sets: {train_sets}') print(f'Test-only sets: {test_sets}') if int(os.environ.get('LOCAL_RANK', 0)) == 0: for tag in train_sets + test_sets: ensure_set_ready(tag, config) # Build normalized embedding matrix + card vocab embedding_dict = utils.get_embedding_dict(embedding_path, add_nontransformed=True) all_vecs = np.array(list(embedding_dict.values())) mean = all_vecs.mean(axis=0) std = 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, ) print(f"Embedding matrix: {embedding_matrix.shape} ({embedding_matrix.numel()*4/1e6:.1f} MB)") # Build per-card GIH WR vector from downloaded 17lands data (-1 = unknown) gih_folder = os.path.join(os.path.dirname(embedding_path), 'gih_wr') gih_card_data = {} # card_name -> list of win rates across all sets for tag in train_sets + test_sets: gih_path = os.path.join(gih_folder, f'{tag}_gih.json') if not os.path.exists(gih_path): continue with open(gih_path) 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 gih_card_data.setdefault(utils.normalize_card_name(entry['name']), []).append(float(wr)) gih_wr_matrix = torch.full((len(cards),), -1.0) for card, idx in card_to_idx.items(): if card in gih_card_data: gih_wr_matrix[idx] = sum(gih_card_data[card]) / len(gih_card_data[card]) n_known = (gih_wr_matrix >= 0).sum().item() print(f"GIH WR known for {n_known} / {len(cards)} cards ({100*n_known/len(cards):.1f}%)") main(card_to_idx, embedding_matrix, gih_wr_matrix, train_sets, test_sets, config, idx_to_card=idx_to_card)