mtg-draft-viz / src /utils /utils.py
TimoBertram's picture
Upload src/utils/utils.py with huggingface_hub
14d9438 verified
Raw
History Blame Contribute Delete
6.77 kB
import json
from collections import defaultdict
import torch
import numpy as np
import os
import pickle
import yaml
basics = ['Mountain', 'Forest', 'Swamp', 'Island', 'Plains']
def normalize_card_name(name):
"""Canonical card name: front face only, no A- prefix, underscores→spaces, basic variants collapsed."""
name = name.split(' // ')[0] # DFC: keep front face only
name = basic_card_name(name) # Mountain_1 → Mountain (must precede underscore replace)
name = name.replace('_', ' ') # CSV column artifact
if name.startswith('A-'):
name = name[2:] # Alchemy rebalanced cards
name = name.replace("Sol'kanar", "Sol'Kanar") # capitalization mismatch in source data
return name
def load_config(path):
with open(path, 'r') as f:
return yaml.safe_load(f)
def unwrap(net):
return net.module if hasattr(net, "module") else net
def get_embedding_dict(path, add_nontransformed=False):
with open(path, 'rb') as f:
embedding_dict = pickle.load(f)
if add_nontransformed:
# Normalize all keys: deduplicate DFCs, strip A- prefixes, collapse basic variants.
# Earlier entries win so the raw full-name embedding isn't silently dropped.
normalized = {}
for k, v in embedding_dict.items():
nk = normalize_card_name(k)
if nk not in normalized:
normalized[nk] = v
return normalized
return embedding_dict
def dump_embedding_dict(embedding_dict, name, embedding_folder = ''):
with open(name, 'wb') as f:
pickle.dump(embedding_dict,f)
def _download_scryfall_cards(path):
import requests
from datetime import date
os.makedirs(path, exist_ok=True)
headers = {"User-Agent": "MTG-IL/1.0", "Accept": "application/json"}
print("Fetching Scryfall bulk data index...")
resp = requests.get("https://api.scryfall.com/bulk-data", headers=headers)
resp.raise_for_status()
bulk_index = resp.json()
if "data" not in bulk_index:
raise RuntimeError(f"Unexpected Scryfall response: {bulk_index}")
entry = next((e for e in bulk_index["data"] if e["type"] == "default_cards"), None)
if entry is None:
raise RuntimeError(f"No default_cards entry. Types: {[e['type'] for e in bulk_index['data']]}")
download_uri = entry["download_uri"]
out_path = os.path.join(path, f"cards_{date.today()}.json")
print(f"Downloading Scryfall card data to {out_path}...")
with requests.get(download_uri, stream=True, headers=headers) as r:
r.raise_for_status()
with open(out_path, "wb") as f:
for chunk in r.iter_content(chunk_size=1 << 17):
f.write(chunk)
print("Scryfall download complete.")
return out_path
def get_card_json(path = './data/', transform = False):
card_files = [f for f in os.listdir(path) if 'cards_' in f]
card_files.sort(key=lambda x: os.path.getmtime(os.path.join(path, x)))
if len(card_files) == 0:
card_path = _download_scryfall_cards(path)
else:
card_path = os.path.join(path, card_files[-1])
print(f'Using cards file: {card_path}')
cards = json.load(open(card_path, encoding='utf8'))
if transform:
card_dict_split = {}
for card in cards:
card_dict_split[normalize_card_name(card['name'])] = card
return card_dict_split
return cards
def load_winrates(path):
import csv
with open(path, 'r') as f:
reader = csv.reader(f)
next(reader)
winrates = {row[0]: float(row[15].strip('%')) if row[15] else None for row in reader}
min_winrate = min([winrate for winrate in winrates.values() if winrate is not None])
for card in winrates:
if winrates[card] is None:
winrates[card] = min_winrate
return winrates
def clean_card(card):
card = basic_card_name(card)
if 'A-' in card:
card = card.replace('A-','')
elif "Sol'kanar" in card:
card = card.replace("Sol'kanar","Sol'Kanar")
elif '_' in card:
card = card.replace('_',' ')
return card
def basic_card_name(card_name):
ints = ['1','2','3','4','5']
for b in basics:
if b in card_name:
for i in ints:
if card_name == f'{b}_{i}':
return b
return card_name
def get_embedding_of_card(card_name, embedding_dict):
try:
key = normalize_card_name(card_name)
if key in embedding_dict:
return embedding_dict[key], False
raise Exception(f'Could not find {card_name!r} (normalized: {key!r})')
except Exception as e:
print(e)
raise e
def get_card_embeddings(card_names, embedding_dict, embedding_size=1330):
embeddings = []
for card in card_names:
if card == '':
embeddings.append([])
elif card == []:
if type(embedding_size) == tuple:
channels, height, width = embedding_size
new_embedding = torch.zeros(1,channels, height, width)
else:
new_embedding = torch.zeros(1,embedding_size)
embeddings.append(new_embedding)
elif isinstance(card, list):
if len(card) == 0:
embeddings.append(None)
continue
deck_embedding = []
for c in card:
embedding, got_new = get_embedding_of_card(c, embedding_dict)
deck_embedding.append(embedding)
try:
num_cards = len(deck_embedding)
deck_embedding = torch.stack(deck_embedding)
if type(embedding_size) == tuple:
channels, height, width = embedding_size
deck_embedding = deck_embedding.view(num_cards,channels, height, width)
else:
deck_embedding = deck_embedding.view(num_cards,-1)
except Exception as e:
raise e
embeddings.append(deck_embedding)
else:
embedding, got_new = get_embedding_of_card(card, embedding_dict)
embeddings.append(embedding)
return embeddings
def get_original_name(transformed_name):
cards_original = transformed_name.replace('_', ' ')
for suffix in [' 1', ' 2', ' 3', ' 4', ' 5']:
cards_original = cards_original.replace(suffix, '')
if 'Nicol Bolas' in cards_original:
cards_original = 'Nicol Bolas, the Ravager // Nicol Bolas, the Arisen'
return cards_original