mtg-draft-viz / src /training /models.py
TimoBertram's picture
Upload src/training/models.py with huggingface_hub
61222b2 verified
Raw
History Blame Contribute Delete
24 kB
import torch
import torch.nn as nn
import torch.nn.functional as F
import pickle
import os
import lmdb
from torch.utils.data import Dataset
class LMDBDataset(Dataset):
def __init__(self, db_path):
self.db_path = db_path
self._env = None
self._keys = None
self._length = None
self._pid = None
def _open(self):
pid = os.getpid()
if self._env is None or self._pid != pid:
if self._env is not None:
self._env.close()
self._env = lmdb.open(
self.db_path,
readonly=True, lock=False, readahead=True, max_readers=8192
)
self._pid = pid
def _ensure_keys(self):
if self._keys is None:
self._open()
with self._env.begin() as txn:
cur = txn.cursor()
self._keys = [bytes(k) for k, _ in cur if k != b"__len__"]
self._length = len(self._keys)
def __len__(self):
if self._length is not None:
return self._length
self._open()
with self._env.begin() as txn:
n = txn.get(b"__len__")
if n is not None:
self._length = int(n.decode())
return self._length
self._ensure_keys()
return self._length
def __getitem__(self, idx):
self._ensure_keys()
k = self._keys[idx]
with self._env.begin() as txn:
v = txn.get(k)
return pickle.loads(v)
def __getstate__(self):
state = self.__dict__.copy()
state["_env"] = None
return state
def __del__(self):
try:
if self._env is not None:
self._env.close()
except Exception:
pass
class Card_Preprocessing(nn.Module):
def __init__(self, num_layers, input_size, output_size, nonlinearity=nn.GELU, internal_size=1024, dropout=0):
super(Card_Preprocessing, self).__init__()
self.internal_size = internal_size
self.input = nn.Sequential(
nn.Linear(input_size, internal_size, bias=False),
nonlinearity(),
nn.LayerNorm(internal_size, bias=False),
nn.Dropout(dropout),
)
self.hidden_layers = nn.ModuleList()
self.dropout_rate = dropout
for _ in range(num_layers):
self.hidden_layers.append(nn.Sequential(
nn.Linear(internal_size, internal_size, bias=False),
nonlinearity(),
nn.LayerNorm(internal_size, bias=False),
nn.Dropout(dropout),
))
self.output = nn.Sequential(
nn.Linear(internal_size, output_size, bias=False),
nonlinearity(),
nn.LayerNorm(output_size, bias=False),
)
self.gammas = nn.ParameterList([
torch.nn.Parameter(torch.ones(1, internal_size), requires_grad=True)
for _ in range(num_layers)
])
def forward(self, x):
x = self.input(x)
for i, layer in enumerate(self.hidden_layers):
gamma = torch.sigmoid(self.gammas[i])
x = gamma * x + (1 - gamma) * layer(x)
x = self.output(x)
return x
class CrossAttnBlock(nn.Module):
def __init__(self, d_model: int, n_heads: int, dropout: float):
super().__init__()
self.ln_q = nn.LayerNorm(d_model)
self.ln_k = nn.LayerNorm(d_model)
self.ln_v = nn.LayerNorm(d_model)
self.xattn = nn.MultiheadAttention(d_model, n_heads, dropout=dropout, batch_first=True)
self.ln_ff = nn.LayerNorm(d_model)
self.ffn = nn.Sequential(
nn.Linear(d_model, 4 * d_model),
nn.GELU(),
nn.Dropout(dropout),
nn.Linear(4 * d_model, d_model),
nn.Dropout(dropout),
)
self.dropout_attn = nn.Dropout(dropout)
def forward(self, cards, deck, attn_mask=None, key_padding_mask=None):
q = self.ln_q(cards)
k = self.ln_k(deck)
v = self.ln_v(deck)
attn_out, _ = self.xattn(q, k, v, attn_mask=attn_mask, key_padding_mask=key_padding_mask)
x = cards + self.dropout_attn(attn_out)
y = self.ffn(self.ln_ff(x))
return x + y
class SelfAttnBlock(nn.Module):
def __init__(self, d_model: int, n_heads: int, dropout: float):
super().__init__()
self.ln_q = nn.LayerNorm(d_model)
self.ln_k = nn.LayerNorm(d_model)
self.ln_v = nn.LayerNorm(d_model)
self.xattn = nn.MultiheadAttention(d_model, n_heads, dropout=dropout, batch_first=True)
self.ln_ff = nn.LayerNorm(d_model)
self.ffn = nn.Sequential(
nn.Linear(d_model, 4 * d_model),
nn.GELU(),
nn.Dropout(dropout),
nn.Linear(4 * d_model, d_model),
nn.Dropout(dropout),
)
self.dropout_attn = nn.Dropout(dropout)
def forward(self, x, key_padding_mask=None, attn_mask=None):
q = self.ln_q(x)
k = self.ln_k(x)
v = self.ln_v(x)
attn_out, _ = self.xattn(q, k, v, key_padding_mask=key_padding_mask,
attn_mask=attn_mask)
x = x + self.dropout_attn(attn_out)
y = self.ffn(self.ln_ff(x))
return x + y
class DecisionDraftTransformer(nn.Module):
"""DraftTransformer conditioned on return-to-go (desired win rate).
No Q/V heads β€” policy is learned directly via BC conditioned on RTG."""
def __init__(self, input_size, num_card_layers, card_output_dim, dropout,
embedding_matrix=None, gih_wr_matrix=None, **kwargs):
super().__init__()
if embedding_matrix is not None:
self.register_buffer('embedding_matrix', embedding_matrix)
else:
self.embedding_matrix = None
if gih_wr_matrix is not None:
self.register_buffer('gih_wr_buffer', gih_wr_matrix)
else:
self.register_buffer('gih_wr_buffer', None)
self.card_encoder = Card_Preprocessing(
num_card_layers, input_size=input_size,
internal_size=1024, output_size=card_output_dim, dropout=dropout,
)
self.pos_embedding = nn.Embedding(128, card_output_dim)
self.outcome_proj = nn.Linear(1, card_output_dim) # draft outcome: wins/(wins+losses)
self.player_proj = nn.Linear(1, card_output_dim) # player skill: historical win rate
self.history_layers = nn.ModuleList([
SelfAttnBlock(card_output_dim, n_heads=8, dropout=dropout)
for _ in range(3)
])
self.pack_self_layers = nn.ModuleList([
SelfAttnBlock(card_output_dim, n_heads=8, dropout=dropout)
for _ in range(1)
])
self.pack_layers = nn.ModuleList([
CrossAttnBlock(card_output_dim, n_heads=8, dropout=dropout)
for _ in range(5)
])
self.output_layer = nn.Sequential(
nn.Linear(card_output_dim, card_output_dim * 2), nn.ReLU(),
nn.LayerNorm(card_output_dim * 2, bias=False), nn.Dropout(dropout),
nn.Linear(card_output_dim * 2, card_output_dim), nn.ReLU(),
nn.LayerNorm(card_output_dim, bias=False),
nn.Linear(card_output_dim, 1),
)
self.playability_head = nn.Sequential(
nn.Linear(card_output_dim * 2, card_output_dim), nn.ReLU(),
nn.LayerNorm(card_output_dim, bias=False), nn.Dropout(dropout),
nn.Linear(card_output_dim, 1),
)
self.gih_head = nn.Linear(card_output_dim, 1)
self.soft_deck_proj = nn.Linear(card_output_dim, card_output_dim)
if kwargs.get('path'):
self.load_state_dict(torch.load(f"{kwargs['path']}/network.pt", map_location='cpu'))
print(f"Loaded model from {kwargs['path']}/network.pt")
def forward(self, history_idx, pack_idx, pack_mask, seq_mask, outcome, player_wr):
"""
outcome : [B] β€” this draft's win rate: wins/(wins+losses)
player_wr : [B] β€” player's historical win rate across all drafts
Returns: logits [B,T,P], play_logits [B,T,P], pick_play_logits [B,T,T],
gih_pred [B,T,P], gih_target [B,T,P], gih_known [B,T,P]
"""
B, T = history_idx.shape
P = pack_idx.shape[2]
device = history_idx.device
pos = torch.arange(T, device=device)
pos_enc = self.pos_embedding(pos)
history_picks = self.embedding_matrix[history_idx]
packs = self.embedding_matrix[pack_idx]
picks_enc = self.card_encoder(history_picks)
cond = (self.outcome_proj(outcome.view(B, 1, 1))
+ self.player_proj(player_wr.view(B, 1, 1))) # [B, 1, D]
start = cond
history = torch.cat([start, picks_enc[:, :-1]], dim=1)
history = history + pos_enc.unsqueeze(0)
history = history + cond # re-inject at every position
causal_mask = torch.triu(torch.ones(T, T, device=device), diagonal=1).bool()
for layer in self.history_layers:
history = layer(history, key_padding_mask=seq_mask, attn_mask=causal_mask)
# Build pick_play_logits from post-attention history (causally valid: history[t]
# only attends to picks 0..t-1 via causal mask, so pick_play_logits[t,s] for s<=t is fine)
hist_exp2 = history.unsqueeze(2).expand(-1, -1, T, -1)
picks_exp2 = picks_enc.unsqueeze(1).expand(-1, T, -1, -1)
pick_play_logits = self.playability_head(
torch.cat([hist_exp2, picks_exp2], dim=-1)).squeeze(-1)
triu_mask = torch.triu(torch.ones(T, T, dtype=torch.bool, device=device), diagonal=1)
pick_play_logits = pick_play_logits.masked_fill(triu_mask.unsqueeze(0), float('-inf'))
pick_play_logits = pick_play_logits.masked_fill(seq_mask.unsqueeze(2), float('-inf'))
pick_play_logits = pick_play_logits.masked_fill(seq_mask.unsqueeze(1), float('-inf'))
# Soft deck: playability-weighted cumulative mean of picks, shifted right (causal)
play_w = torch.sigmoid(pick_play_logits.diagonal(dim1=1, dim2=2).clone())
play_w = play_w.masked_fill(seq_mask, 0.0)
weighted_picks = picks_enc * play_w.unsqueeze(-1)
soft_deck = torch.cat([torch.zeros(B, 1, picks_enc.shape[-1], device=device),
torch.cumsum(weighted_picks, dim=1)[:, :-1]], dim=1)
soft_w = torch.cat([torch.zeros(B, 1, device=device),
torch.cumsum(play_w, dim=1)[:, :-1]], dim=1)
soft_deck = soft_deck / soft_w.clamp(min=1e-8).unsqueeze(-1)
# Augment history with deck state before pack cross-attention
history = history + self.soft_deck_proj(soft_deck)
# Encode packs
packs_enc = self.card_encoder(packs.view(B * T, P, -1))
gih_pred = torch.sigmoid(self.gih_head(packs_enc)).view(B, T, P)
if self.gih_wr_buffer is not None:
gih_target = self.gih_wr_buffer[pack_idx]
gih_known = (gih_target >= 0) & pack_mask
else:
gih_target = torch.zeros_like(gih_pred)
gih_known = torch.zeros(B, T, P, dtype=torch.bool, device=device)
pack_slot_mask = ~pack_mask.view(B * T, P)
all_masked = pack_slot_mask.all(dim=-1)
if all_masked.any():
pack_slot_mask = pack_slot_mask.clone()
pack_slot_mask[all_masked, 0] = False
for layer in self.pack_self_layers:
packs_enc = layer(packs_enc, key_padding_mask=pack_slot_mask)
packs_enc = packs_enc.view(B, T, P, -1)
packs_enc = packs_enc + pos_enc.unsqueeze(0).unsqueeze(2)
packs_enc = packs_enc.view(B, T * P, -1)
pack_causal_mask = torch.triu(
torch.ones(T, T, device=device, dtype=torch.bool), diagonal=1
).repeat_interleave(P, dim=0)
for layer in self.pack_layers:
packs_enc = layer(packs_enc, history,
attn_mask=pack_causal_mask,
key_padding_mask=seq_mask)
packs_enc = packs_enc.view(B, T, P, -1)
logits = self.output_layer(packs_enc) \
.masked_fill(~pack_mask.unsqueeze(-1), float('-inf')) \
.squeeze(-1)
hist_exp = history.unsqueeze(2).expand(-1, -1, P, -1)
play_logits = self.playability_head(torch.cat([hist_exp, packs_enc], dim=-1)).squeeze(-1)
play_logits = play_logits.masked_fill(~pack_mask, float('-inf'))
return logits, play_logits, pick_play_logits, gih_pred, gih_target, gih_known
class DraftTransformer(nn.Module):
def __init__(self, input_size, num_card_layers, card_output_dim, dropout,
embedding_matrix=None, gih_wr_matrix=None, **kwargs):
super().__init__()
# Fixed LLaMA embedding lookup β€” not trained, lives on GPU permanently
if embedding_matrix is not None:
self.register_buffer('embedding_matrix', embedding_matrix)
else:
self.embedding_matrix = None
# Per-card GIH win rate targets for auxiliary supervision (-1 = unknown)
if gih_wr_matrix is not None:
self.register_buffer('gih_wr_buffer', gih_wr_matrix)
else:
self.register_buffer('gih_wr_buffer', None)
self.card_encoder = Card_Preprocessing(
num_card_layers, input_size=input_size,
internal_size=1024, output_size=card_output_dim, dropout=dropout,
)
# Learned positional encoding shared by history and pack queries
self.pos_embedding = nn.Embedding(128, card_output_dim)
# Learnable start-of-draft token
self.start_token = nn.Parameter(torch.zeros(1, 1, card_output_dim))
# Causal self-attention over pick history
self.history_layers = nn.ModuleList([
SelfAttnBlock(card_output_dim, n_heads=8, dropout=dropout)
for _ in range(3)
])
# Within-pack self-attention: cards in the same pack compare against each other
self.pack_self_layers = nn.ModuleList([
SelfAttnBlock(card_output_dim, n_heads=8, dropout=dropout)
for _ in range(1)
])
# Pack cards cross-attend to the history state at the current step
self.pack_layers = nn.ModuleList([
CrossAttnBlock(card_output_dim, n_heads=8, dropout=dropout)
for _ in range(5)
])
self.output_layer = nn.Sequential(
nn.Linear(card_output_dim, card_output_dim * 2),
nn.ReLU(),
nn.LayerNorm(card_output_dim * 2, bias=False),
nn.Dropout(dropout),
nn.Linear(card_output_dim * 2, card_output_dim),
nn.ReLU(),
nn.LayerNorm(card_output_dim, bias=False),
nn.Linear(card_output_dim, 1),
)
self.q_head = nn.Sequential(
nn.Linear(card_output_dim * 2, card_output_dim * 2),
nn.ReLU(),
nn.LayerNorm(card_output_dim * 2, bias=False),
nn.Dropout(dropout),
nn.Linear(card_output_dim * 2, card_output_dim),
nn.ReLU(),
nn.LayerNorm(card_output_dim, bias=False),
nn.Linear(card_output_dim, 1),
)
# Playability head: P(card in maindeck) given deck context + card encoding.
# Input: cat(history[t], card_enc[t, j]) β€” 2*d dimensional. At training,
# only slot 0 (the picked card) is supervised; all P slots are computed at inference.
self.playability_head = nn.Sequential(
nn.Linear(card_output_dim * 2, card_output_dim),
nn.ReLU(),
nn.LayerNorm(card_output_dim, bias=False),
nn.Dropout(dropout),
nn.Linear(card_output_dim, 1),
)
# Value head: predicts win rate from playability-weighted soft deck.
# soft_deck[t] = Ξ£_{s<t} sigmoid(play[s]) * picks_enc[s] / Ξ£_{s<t} sigmoid(play[s])
self.value_head = nn.Sequential(
nn.Linear(card_output_dim, card_output_dim),
nn.ReLU(),
nn.LayerNorm(card_output_dim, bias=False),
nn.Dropout(dropout),
nn.Linear(card_output_dim, 1),
)
# Predicts GIH WR from raw card encoding (before any context)
self.gih_head = nn.Linear(card_output_dim, 1)
if kwargs.get('path'):
self.load_state_dict(torch.load(f"{kwargs['path']}/network.pt", map_location='cpu'))
print(f"Loaded model from {kwargs['path']}/network.pt")
def forward(self, history_idx, pack_idx, pack_mask, seq_mask):
"""
history_idx : [B, T] β€” int64 indices of picked cards
pack_idx : [B, T, P] β€” int64 indices of pack cards at each step
pack_mask : [B, T, P] β€” bool, True where card slot is valid
seq_mask : [B, T] β€” bool, True where step is padding
Returns : logits [B, T, P], q_values [B, T, P], values [B, T],
play_logits [B, T, P], pick_play_logits [B, T, T],
gih_pred [B, T, P], gih_target [B, T, P], gih_known [B, T, P]
"""
B, T = history_idx.shape
P = pack_idx.shape[2]
device = history_idx.device
# Positional encoding shared by history and pack (same index = same step)
pos = torch.arange(T, device=device)
pos_enc = self.pos_embedding(pos) # [T, d]
# GPU embedding lookup
history_picks = self.embedding_matrix[history_idx] # [B, T, E]
packs = self.embedding_matrix[pack_idx] # [B, T, P, E]
# Encode picked cards, shift right, prepend start token, add positional encoding
picks_enc = self.card_encoder(history_picks) # [B, T, d]
start = self.start_token.expand(B, -1, -1) # [B, 1, d]
history = torch.cat([start, picks_enc[:, :-1]], dim=1) # [B, T, d]
history = history + pos_enc.unsqueeze(0) # [B, T, d]
# Causal self-attention over history
causal_mask = torch.triu(torch.ones(T, T, device=device), diagonal=1).bool()
for layer in self.history_layers:
history = layer(history, key_padding_mask=seq_mask, attn_mask=causal_mask)
# Encode pack cards: [B*T, P, d]
packs_enc = self.card_encoder(packs.view(B * T, P, -1)) # [B*T, P, d]
# GIH auxiliary: predict intrinsic card quality before any context is added
gih_pred = torch.sigmoid(self.gih_head(packs_enc)).view(B, T, P)
if self.gih_wr_buffer is not None:
gih_target = self.gih_wr_buffer[pack_idx] # [B, T, P]
gih_known = (gih_target >= 0) & pack_mask # [B, T, P]
else:
gih_target = torch.zeros_like(gih_pred)
gih_known = torch.zeros(B, T, P, dtype=torch.bool, device=device)
# Within-pack self-attention: cards in the same pack compare against each other
pack_slot_mask = ~pack_mask.view(B * T, P) # True = invalid slot
# Padding steps have ALL slots masked β†’ all-masked softmax β†’ NaN.
# Fix at source: unmask slot 0 for those rows so softmax always has β‰₯1 valid key.
# Padding steps have no loss contribution (seq_mask=True), so the dummy slot is harmless.
all_masked = pack_slot_mask.all(dim=-1)
if all_masked.any():
pack_slot_mask = pack_slot_mask.clone()
pack_slot_mask[all_masked, 0] = False
for layer in self.pack_self_layers:
packs_enc = layer(packs_enc, key_padding_mask=pack_slot_mask)
# Add step positional encoding so pack cards know which pick they belong to
packs_enc = packs_enc.view(B, T, P, -1)
packs_enc = packs_enc + pos_enc.unsqueeze(0).unsqueeze(2) # [B, T, P, d]
packs_enc = packs_enc.view(B, T * P, -1) # [B, T*P, d]
# Causal cross-attention: pack card at step t attends to history 0..t only
pack_causal_mask = torch.triu(
torch.ones(T, T, device=device, dtype=torch.bool), diagonal=1
).repeat_interleave(P, dim=0) # [T*P, T]
for layer in self.pack_layers:
packs_enc = layer(packs_enc, history,
attn_mask=pack_causal_mask,
key_padding_mask=seq_mask)
packs_enc = packs_enc.view(B, T, P, -1) # [B, T, P, d]
# Logits
logits = self.output_layer(packs_enc) \
.masked_fill(~pack_mask.unsqueeze(-1), float('-inf')) \
.squeeze(-1) # [B, T, P]
# Pack-card playability [B, T, P] β€” used at inference to show per-card play probability.
hist_exp_play = history.unsqueeze(2).expand(-1, -1, P, -1) # [B, T, P, d]
play_input = torch.cat([hist_exp_play, packs_enc], dim=-1) # [B, T, P, 2d]
play_logits = self.playability_head(play_input).squeeze(-1) # [B, T, P]
play_logits = play_logits.masked_fill(~pack_mask, float('-inf'))
# Historical-pick playability [B, T, T] β€” for training and soft deck.
# pick_play_logits[b, t, s] = P(pick_s in maindeck | deck context at step t), for s <= t.
hist_exp2 = history.unsqueeze(2).expand(-1, -1, T, -1) # [B, T, T, d]
picks_exp2 = picks_enc.unsqueeze(1).expand(-1, T, -1, -1) # [B, T, T, d]
pick_play_input = torch.cat([hist_exp2, picks_exp2], dim=-1) # [B, T, T, 2d]
pick_play_logits = self.playability_head(pick_play_input).squeeze(-1) # [B, T, T]
triu_mask = torch.triu(torch.ones(T, T, dtype=torch.bool, device=device), diagonal=1)
pick_play_logits = pick_play_logits.masked_fill(triu_mask.unsqueeze(0), float('-inf'))
pick_play_logits = pick_play_logits.masked_fill(seq_mask.unsqueeze(2), float('-inf'))
pick_play_logits = pick_play_logits.masked_fill(seq_mask.unsqueeze(1), float('-inf'))
# Soft deck: playability-weighted cumulative mean of picks (causal, shifted right).
play_w = pick_play_logits.diagonal(dim1=1, dim2=2).clone() # [B, T]
play_w = torch.sigmoid(play_w).masked_fill(seq_mask, 0.0)
weighted_picks = picks_enc * play_w.unsqueeze(-1) # [B, T, d]
cum_w_picks = torch.cumsum(weighted_picks, dim=1) # [B, T, d]
cum_w = torch.cumsum(play_w, dim=1) # [B, T]
soft_deck = torch.cat([torch.zeros(B, 1, picks_enc.shape[-1], device=device),
cum_w_picks[:, :-1]], dim=1) # [B, T, d]
soft_w = torch.cat([torch.zeros(B, 1, device=device),
cum_w[:, :-1]], dim=1) # [B, T]
soft_deck = soft_deck / soft_w.clamp(min=1e-8).unsqueeze(-1) # [B, T, d]
# Value head reads from soft deck
values = self.value_head(soft_deck).squeeze(-1) # [B, T]
values = values.masked_fill(seq_mask, float('-inf'))
# Q-values: soft deck state (what we've built) + pack card (what we'd add)
soft_exp = soft_deck.unsqueeze(2).expand(-1, -1, P, -1) # [B, T, P, d]
q_input = torch.cat([soft_exp, packs_enc], dim=-1) # [B, T, P, 2d]
q_values = self.q_head(q_input).squeeze(-1) # [B, T, P]
q_values = q_values.masked_fill(~pack_mask, float('-inf'))
return logits, q_values, values, play_logits, pick_play_logits, gih_pred, gih_target, gih_known