Spaces:
Sleeping
Sleeping
File size: 20,362 Bytes
aea7828 1b75d58 aea7828 1b75d58 aea7828 1b75d58 aea7828 ed681ce aea7828 ed681ce 1b75d58 aea7828 1b75d58 aea7828 ed681ce aea7828 1b75d58 aea7828 030c370 aea7828 030c370 aea7828 030c370 aea7828 030c370 aea7828 1b75d58 ed681ce aea7828 030c370 aea7828 ed681ce aea7828 030c370 aea7828 ed681ce aea7828 030c370 aea7828 030c370 aea7828 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 | 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)
|