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