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