import json import math import numpy as np import torch 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_dt.metrics import evaluate as pick_metrics from src.preprocessing.preprocess_pipeline import ensure_set_ready from src.training.train_IL import collate_fn, worker_init_fn, wlog MAX_CHOICES = 15 card_to_idx = None def training_step(network, batch, optimizer, scheduler, device, scaler, lam_gih=1.0, lam_play=1.0): 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 W = wins.to(device).float() L = losses.to(device).float() outcome = W / (W + L).clamp(min=1) # [B] this draft's WR player_wr_d = user_wr.to(device) # [B] historical WR # Skill weighting: upweight picks from high win-rate drafters skill_w = (player_wr_d / player_wr_d.mean().clamp(min=1e-8)).unsqueeze(1) # [B, 1] with torch.autocast(device_type='cuda'): logits, play_logits, pick_play_logits, gih_pred, gih_target, gih_known = network( history_idx, pack_idx, pack_mask, seq_mask, outcome, player_wr_d) log_probs = F.log_softmax(logits, dim=-1)[..., 0].masked_fill(~valid, 0.0) bc_loss = -(skill_w * log_probs * valid.float()).sum() / valid.float().sum().clamp(min=1) # Playability loss T = seq_mask.shape[1] play_known_d = play_known.to(device) play_target_d = play_target.to(device) pick_play_mask = ( (~seq_mask).unsqueeze(2) & (~seq_mask).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_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 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 loss = bc_loss + lam_play * play_loss + lam_gih * gih_loss play_diag = torch.sigmoid(pick_play_logits.diagonal(dim1=1, dim2=2)) _play_w_std = play_diag[~seq_mask].std().item() _bc, _play, _gih = bc_loss.item(), play_loss.item(), gih_loss.item() if any(math.isnan(x) for x in (_bc, _play, _gih)): print(f"[NaN] bc={_bc:.4f} play={_play:.4f} gih={_gih:.4f}") 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, _play, _gih, _play_w_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) network.train() bar = tqdm.tqdm(enumerate(train_loader), total=total_iterations, mininterval=1, desc='Training') if is_master else enumerate(train_loader) lam_gih = config.get('lambda_gih', 1.0) lam_play = config.get('lambda_play', 1.0) for i, batch in bar: bc_loss, play_loss, gih_loss, play_w_std = training_step( network, batch, optimizer, scheduler, device=torch.device(f'cuda:{local_rank}'), scaler=scaler, lam_gih=lam_gih, lam_play=lam_play) if is_master: wlog({ "train/bc_loss": bc_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, }, 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) global_step += 1 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" if is_master: wandb.init(entity="tibert97", project="Drafting IL", config=config) run_id = f"{wandb.run.name}-{wandb.run.id}" 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'] train_data = ConcatDataset([models.LMDBDataset(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_data = ConcatDataset([models.LMDBDataset(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_test_loader = None if test_sets: eval_test_data = ConcatDataset([models.LMDBDataset(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.DecisionDraftTransformer(**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) 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}") gih_folder = os.path.join(os.path.dirname(embedding_path), 'gih_wr') gih_card_data = {} 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]) print(f"GIH WR known for {(gih_wr_matrix >= 0).sum().item()} / {len(cards)} cards") main(card_to_idx, embedding_matrix, gih_wr_matrix, train_sets, test_sets, config, idx_to_card=idx_to_card)