mtg-draft-viz / src /training /train_DT.py
TimoBertram's picture
Upload src/training/train_DT.py with huggingface_hub
4824f76 verified
Raw
History Blame Contribute Delete
13.1 kB
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)