Contest2 / app.py
ffzeroHua's picture
Update app.py
41ca4e3 verified
Raw
History Blame Contribute Delete
31.4 kB
import os
import orjson
import json
import gzip
import concurrent.futures
import random
import torch
import threading
import time
import uuid
import glob
import requests
import traceback
import pathlib
import secrets
import importlib.util
import sys
import numpy as np
import pandas as pd
import matplotlib.pyplot as plt
import gradio as gr
from huggingface_hub import snapshot_download, hf_hub_download, HfApi
# 👇 强制关闭底层可能引发 NaN 的 Flash/MemEff 注意力后端,锁定安全数学模式
import torch.backends.cuda
torch.backends.cuda.enable_flash_sdp(False)
torch.backends.cuda.enable_mem_efficient_sdp(False)
torch.backends.cuda.enable_math_sdp(True)
from torch import nn, Tensor
from torch.nn import functional as F
from torch.distributions import Normal, Categorical
from typing import *
from functools import partial
from itertools import permutations
# 🚀 设定要从云端拉取并进行对抗的两个模型名称与架构 ('resnet' 或 'transformer')
TEST_MODEL = "4zv4_2.pth"
TEST_ARCH = "transformer"
EXAMINER_MODEL = "Elite4z9070.pth"
EXAMINER_ARCH = "resnet"
# ==========================================
# 0. 底层引擎与分布式多开配置
# ==========================================
try:
from libriichi3p.mjai import Bot
from libriichi3p.consts import obs_shape, oracle_obs_shape, ACTION_SPACE, GRP_SIZE
except:
# ⚠️ 请确保此路径指向正确的 .so 文件
SO_FILE_PATH = "/content/drive/MyDrive/MahjongTest/libriichi3p.so"
if not os.path.exists(SO_FILE_PATH):
print(f"❌ 致命错误:在路径 {SO_FILE_PATH} 下根本找不到文件!请检查路径拼写。")
else:
print(f"✅ 找到文件: {SO_FILE_PATH},正在尝试强行加载...")
try:
spec = importlib.util.spec_from_file_location("libriichi3p", SO_FILE_PATH)
libriichi3p_module = importlib.util.module_from_spec(spec)
sys.modules["libriichi3p"] = libriichi3p_module
spec.loader.exec_module(libriichi3p_module)
from libriichi3p.mjai import Bot
from libriichi3p.consts import obs_shape, oracle_obs_shape, ACTION_SPACE
print("🎉 强行导入成功!")
except Exception as e:
print(f"❌ 导入失败,暴露出真实报错: {e}")
try:
from riichienv import RiichiEnv, GameRule
except ImportError:
print("⚠️ 警告:找不到 riichienv 模块,请确保环境已安装该依赖。")
# ========== Online Server =========== #
OT_REQUEST_TIMEOUT = 2
ot_settings = {
"server": "http://example.com",
"online": False,
"api_key": "example_api_key",
}
is_online = False
def online_settings_init():
global ot_settings
if (pathlib.Path(__file__).parent / 'ot_settings.json').exists():
with open(pathlib.Path(__file__).parent / 'ot_settings.json', 'r') as f:
ot_settings = json.load(f)
online_settings_init()
# ==================================== #
DATA_REPO_ID = "ffzeroHua/mj-eval-results"
MODEL_REPO_ID = "ffzeroHua/Riichi-Model-Repo"
HF_TOKEN = os.getenv("HF_TOKEN")
WORKER_ID = os.getenv("WORKER_ID", str(uuid.uuid4())[:6])
REPORT_FILE_PREFIX = 'Step40800P42998_vs_9070_Fair_eval_report'
REPORT_FILE = f"{REPORT_FILE_PREFIX}_{WORKER_ID}.txt"
api = HfApi()
EVAL_RUNNING = True
# ==========================================
# 1. 共享函数
# ==========================================
def sample_top_p(logits, p):
if p >= 1:
return Categorical(logits=logits).sample()
if p <= 0:
return logits.argmax(-1)
probs = logits.softmax(-1)
probs_sort, probs_idx = probs.sort(-1, descending=True)
probs_sum = probs_sort.cumsum(-1)
mask = probs_sum - probs_sort > p
probs_sort[mask] = 0.
sampled = probs_idx.gather(-1, probs_sort.multinomial(1)).squeeze(-1)
return sampled
# ==========================================
# 2. ResNet 架构定义 (经典)
# ==========================================
class ChannelAttention(nn.Module):
def __init__(self, channels, ratio=16, actv_builder=nn.ReLU, bias=True):
super().__init__()
self.shared_mlp = nn.Sequential(
nn.Linear(channels, channels // ratio, bias=bias),
actv_builder(),
nn.Linear(channels // ratio, channels, bias=bias),
)
if bias:
for mod in self.modules():
if isinstance(mod, nn.Linear):
nn.init.constant_(mod.bias, 0)
def forward(self, x: Tensor):
avg_out = self.shared_mlp(x.mean(-1))
max_out = self.shared_mlp(x.amax(-1))
weight = (avg_out + max_out).sigmoid()
x = weight.unsqueeze(-1) * x
return x
class ResBlock(nn.Module):
def __init__(self, channels, *, norm_builder = nn.Identity, actv_builder = nn.ReLU, pre_actv = False):
super().__init__()
self.pre_actv = pre_actv
if pre_actv:
self.res_unit = nn.Sequential(
norm_builder(), actv_builder(),
nn.Conv1d(channels, channels, kernel_size=3, padding=1, bias=False),
norm_builder(), actv_builder(),
nn.Conv1d(channels, channels, kernel_size=3, padding=1, bias=False),
)
else:
self.res_unit = nn.Sequential(
nn.Conv1d(channels, channels, kernel_size=3, padding=1, bias=False),
norm_builder(), actv_builder(),
nn.Conv1d(channels, channels, kernel_size=3, padding=1, bias=False),
norm_builder(),
)
self.actv = actv_builder()
self.ca = ChannelAttention(channels, actv_builder=actv_builder, bias=True)
def forward(self, x):
out = self.res_unit(x)
out = self.ca(out)
out = out + x
if not self.pre_actv:
out = self.actv(out)
return out
class ResNetCore(nn.Module):
def __init__(self, in_channels, conv_channels, num_blocks, *, norm_builder = nn.Identity, actv_builder = nn.ReLU, pre_actv = False):
super().__init__()
blocks = [ResBlock(conv_channels, norm_builder=norm_builder, actv_builder=actv_builder, pre_actv=pre_actv) for _ in range(num_blocks)]
layers = [nn.Conv1d(in_channels, conv_channels, kernel_size=3, padding=1, bias=False)]
if pre_actv: layers += [*blocks, norm_builder(), actv_builder()]
else: layers += [norm_builder(), actv_builder(), *blocks]
layers += [
nn.Conv1d(conv_channels, 32, kernel_size=3, padding=1),
actv_builder(),
nn.Flatten(),
nn.Linear(32 * 34, 1024),
]
self.net = nn.Sequential(*layers)
def forward(self, x):
return self.net(x)
class ResNetBrain(nn.Module):
def __init__(self, *, conv_channels, num_blocks, is_oracle=False, version=1):
super().__init__()
self.is_oracle = is_oracle
self.version = version
in_channels = obs_shape(version)[0]
if is_oracle: in_channels += oracle_obs_shape(version)[0]
norm_builder = partial(nn.BatchNorm1d, conv_channels, momentum=0.01)
actv_builder = partial(nn.Mish, inplace=True)
pre_actv = True
if version == 1:
actv_builder = partial(nn.ReLU, inplace=True)
pre_actv = False
self.latent_net = nn.Sequential(nn.Linear(1024, 512), nn.ReLU(inplace=True))
self.mu_head = nn.Linear(512, 512)
self.logsig_head = nn.Linear(512, 512)
elif version in (3, 4):
norm_builder = partial(nn.BatchNorm1d, conv_channels, momentum=0.01, eps=1e-3)
self.encoder = ResNetCore(
in_channels=in_channels, conv_channels=conv_channels, num_blocks=num_blocks,
norm_builder=norm_builder, actv_builder=actv_builder, pre_actv=pre_actv,
)
self.actv = actv_builder()
self._freeze_bn = False
def forward(self, obs: Tensor, invisible_obs: Optional[Tensor] = None) -> Union[Tuple[Tensor, Tensor], Tensor]:
if self.is_oracle:
assert invisible_obs is not None
obs = torch.cat((obs, invisible_obs), dim=1)
phi = self.encoder(obs)
phi = F.dropout(phi, p=0.1, training=self.training)
if self.version == 1:
latent_out = self.latent_net(phi)
mu = self.mu_head(latent_out)
logsig = self.logsig_head(latent_out)
return mu, logsig
return self.actv(phi)
class ResNetDQN(nn.Module):
def __init__(self, *, version=1):
super().__init__()
self.version = version
if version == 1:
self.v_head = nn.Linear(512, 1)
self.a_head = nn.Linear(512, ACTION_SPACE)
elif version in (2, 3):
hidden_size = 512 if version == 2 else 256
self.v_head = nn.Sequential(nn.Linear(1024, hidden_size), nn.Mish(inplace=True), nn.Linear(hidden_size, 1))
self.a_head = nn.Sequential(nn.Linear(1024, hidden_size), nn.Mish(inplace=True), nn.Linear(hidden_size, ACTION_SPACE))
elif version == 4:
self.net = nn.Linear(1024, 1 + ACTION_SPACE)
nn.init.constant_(self.net.bias, 0)
def forward(self, phi, mask):
if self.version == 4:
v, a = self.net(phi).split((1, ACTION_SPACE), dim=-1)
else:
v = self.v_head(phi)
a = self.a_head(phi)
a_sum = a.masked_fill(~mask, 0.).sum(-1, keepdim=True)
mask_sum = mask.sum(-1, keepdim=True)
a_mean = a_sum / mask_sum
q = (v + a - a_mean).masked_fill(~mask, -1e9)
return q
class ResNetMortalEngine:
def __init__(self, brain, dqn, is_oracle, version, device=None, stochastic_latent=False, enable_amp=False, enable_quick_eval=True, enable_rule_based_agari_guard=False, name='NoName', boltzmann_epsilon=0, boltzmann_temp=1, top_p=1):
self.engine_type = 'mortal'
self.device = device or torch.device('cpu')
self.brain = brain.to(self.device).eval()
self.dqn = dqn.to(self.device).eval()
self.is_oracle, self.version, self.stochastic_latent = is_oracle, version, stochastic_latent
self.enable_amp, self.enable_quick_eval, self.enable_rule_based_agari_guard, self.name = enable_amp, enable_quick_eval, enable_rule_based_agari_guard, name
self.boltzmann_epsilon, self.boltzmann_temp, self.top_p = boltzmann_epsilon, boltzmann_temp, top_p
def react_batch(self, obs, masks, invisible_obs):
global ot_settings, is_online
if ot_settings['online']:
try:
list_obs, list_masks = [o.tolist() for o in obs], [m.tolist() for m in masks]
data = gzip.compress(json.dumps({'obs': list_obs, 'masks': list_masks}, separators=(',', ':')).encode('utf-8'))
headers = {'Authorization': ot_settings['api_key'], 'Content-Encoding': 'gzip'}
r = requests.post(f'{ot_settings["server"]}/react_batch_3p', headers=headers, data=data, timeout=OT_REQUEST_TIMEOUT)
assert r.status_code == 200
is_online = True
r_json = r.json()
return r_json['actions'], r_json['q_out'], r_json['masks'], r_json['is_greedy']
except:
is_online = False
try:
with torch.autocast(self.device.type, enabled=self.enable_amp), torch.inference_mode():
return self._react_batch(obs, masks, invisible_obs)
except Exception as ex:
raise Exception(f'{ex}\n{traceback.format_exc()}')
def _react_batch(self, obs, masks, invisible_obs):
obs = torch.as_tensor(np.stack(obs, axis=0), device=self.device)
masks = torch.as_tensor(np.stack(masks, axis=0), device=self.device)
invisible_obs = None
if self.is_oracle: invisible_obs = torch.as_tensor(np.stack(invisible_obs, axis=0), device=self.device)
batch_size = obs.shape[0]
if self.version == 1:
mu, logsig = self.brain(obs, invisible_obs)
latent = Normal(mu, logsig.exp() + 1e-6).sample() if self.stochastic_latent else mu
q_out = self.dqn(latent, masks)
elif self.version in (2, 3, 4):
phi = self.brain(obs)
q_out = self.dqn(phi, masks)
if self.boltzmann_epsilon > 0:
is_greedy = torch.full((batch_size,), 1-self.boltzmann_epsilon, device=self.device).bernoulli().to(torch.bool)
logits = (q_out / self.boltzmann_temp).masked_fill(~masks, -torch.inf)
sampled = sample_top_p(logits, self.top_p)
actions = torch.where(is_greedy, q_out.argmax(-1), sampled)
else:
is_greedy = torch.ones(batch_size, dtype=torch.bool, device=self.device)
actions = q_out.argmax(-1)
return actions.tolist(), q_out.tolist(), masks.tolist(), is_greedy.tolist()
# ==========================================
# 3. Transformer 架构定义
# ==========================================
class MortalToTransformerAdapter(nn.Module):
def __init__(self, original_channels: int, d_model: int = 256):
super().__init__()
self.sanma_indices = [0, 8] + list(range(9, 34))
self.seq_len = len(self.sanma_indices)
self.input_proj = nn.Linear(original_channels, d_model)
self.pos_embedding = nn.Parameter(torch.randn(1, self.seq_len, d_model) * 0.02)
self.cls_token = nn.Parameter(torch.randn(1, 1, d_model) * 0.02)
def forward(self, obs_34x_c: Tensor) -> Tensor:
x = obs_34x_c.transpose(1, 2)[:, self.sanma_indices, :]
x = self.input_proj(x) + self.pos_embedding
batch_size = x.size(0)
cls_tokens = self.cls_token.expand(batch_size, -1, -1)
return torch.cat([cls_tokens, x], dim=1)
def build_sanma_distance_bias(num_heads: int, max_distance: int = 9) -> torch.Tensor:
ranks = [1, 9] + list(range(1, 10)) + list(range(1, 10)) + [0] * 7
suits = [0, 0] + [1] * 9 + [2] * 9 + [3] * 7
seq_len = 28
dist_matrix = torch.zeros((seq_len, seq_len), dtype=torch.long)
for i in range(1, seq_len):
for j in range(1, seq_len):
tile_i, tile_j = i - 1, j - 1
if suits[tile_i] == suits[tile_j] and suits[tile_i] in (1, 2):
dist_matrix[i, j] = abs(ranks[tile_i] - ranks[tile_j])
else:
dist_matrix[i, j] = max_distance
dist_matrix[0, :] = max_distance + 1
dist_matrix[:, 0] = max_distance + 1
return dist_matrix
class TransformerBrain(nn.Module):
def __init__(self, *, conv_channels, num_blocks, is_oracle=False, version=4):
super().__init__()
self.is_oracle = is_oracle
self.version = version
d_model = conv_channels
num_layers = num_blocks
nhead = 8
in_channels = obs_shape(version)[0]
if is_oracle: in_channels += oracle_obs_shape(version)[0]
self.adapter = MortalToTransformerAdapter(original_channels=in_channels, d_model=d_model)
self.distance_embedding = nn.Embedding(11, nhead)
self.register_buffer("distance_indices", build_sanma_distance_bias(nhead))
encoder_layer = nn.TransformerEncoderLayer(d_model=d_model, nhead=nhead, dim_feedforward=d_model * 4, batch_first=True, norm_first=True, activation="gelu")
self.transformer = nn.TransformerEncoder(encoder_layer, num_layers=num_layers)
def forward(self, obs: Tensor, invisible_obs: Optional[Tensor] = None) -> Union[Tuple[Tensor, Tensor], Tensor]:
if self.is_oracle:
assert invisible_obs is not None
obs = torch.cat((obs, invisible_obs), dim=1)
batch_size = obs.size(0)
x = self.adapter(obs)
attn_bias = self.distance_embedding(self.distance_indices)
num_heads = attn_bias.size(-1)
attn_bias = attn_bias.permute(2, 0, 1).unsqueeze(0).expand(batch_size, -1, -1, -1).reshape(batch_size * num_heads, 28, 28).contiguous()
# 防 NaN Hook
def PREVENT_NAN_HOOK(module, input, output): return
hooks = [module.register_forward_hook(PREVENT_NAN_HOOK) for name, module in self.transformer.named_modules()]
out = self.transformer(x, mask=attn_bias)
for h in hooks: h.remove()
return out[:, 0, :]
class TransformerDQN(nn.Module):
def __init__(self, *, version=4, d_model=256):
super().__init__()
self.version = version
if version == 4:
self.net = nn.Linear(d_model, 1 + ACTION_SPACE)
nn.init.constant_(self.net.bias, 0)
else:
raise ValueError(f'Unexpected version {self.version} for this backend')
def forward(self, phi, mask):
v, a = self.net(phi).split((1, ACTION_SPACE), dim=-1)
return a.masked_fill(~mask, -1e9)
class TransformerMortalEngine:
def __init__(self, brain, dqn, is_oracle, version, device=None, stochastic_latent=False, enable_amp=False, enable_quick_eval=True, enable_rule_based_agari_guard=False, name='NoName', boltzmann_epsilon=0, boltzmann_temp=1, top_p=1):
self.engine_type = 'mortal'
self.device = device or torch.device('cpu')
self.brain = brain.to(self.device).eval()
self.dqn = dqn.to(self.device).eval()
self.is_oracle, self.version, self.stochastic_latent = is_oracle, version, stochastic_latent
self.enable_amp, self.enable_quick_eval, self.enable_rule_based_agari_guard, self.name = enable_amp, enable_quick_eval, enable_rule_based_agari_guard, name
self.boltzmann_epsilon, self.boltzmann_temp, self.top_p = boltzmann_epsilon, boltzmann_temp, top_p
def react_batch(self, obs, masks, invisible_obs):
global ot_settings, is_online
if ot_settings['online']:
try:
list_obs, list_masks = [o.tolist() for o in obs], [m.tolist() for m in masks]
data = gzip.compress(json.dumps({'obs': list_obs, 'masks': list_masks}, separators=(',', ':')).encode('utf-8'))
headers = {'Authorization': ot_settings['api_key'], 'Content-Encoding': 'gzip'}
r = requests.post(f'{ot_settings["server"]}/react_batch_3p', headers=headers, data=data, timeout=OT_REQUEST_TIMEOUT)
assert r.status_code == 200
is_online = True
r_json = r.json()
return r_json['actions'], r_json['q_out'], r_json['masks'], r_json['is_greedy']
except:
is_online = False
try:
with torch.inference_mode():
return self._react_batch(obs, masks, invisible_obs)
except Exception as ex:
raise Exception(f'{ex}\n{traceback.format_exc()}')
def _react_batch(self, obs, masks, invisible_obs):
obs_tensor = torch.as_tensor(np.stack(obs, axis=0), dtype=torch.float32, device=self.device)
masks_tensor = torch.as_tensor(np.stack(masks, axis=0), dtype=torch.bool, device=self.device)
batch_size = obs_tensor.shape[0]
with torch.autocast(device_type=self.device.type, dtype=torch.bfloat16):
phi = self.brain(obs_tensor)
q_out = self.dqn(phi, masks_tensor)
if self.boltzmann_epsilon > 0:
is_greedy = torch.full((batch_size,), 1-self.boltzmann_epsilon, device=self.device).bernoulli().to(torch.bool)
logits = (q_out / self.boltzmann_temp).masked_fill(~masks_tensor, -torch.inf)
sampled = sample_top_p(logits, self.top_p)
actions = torch.where(is_greedy, q_out.argmax(-1), sampled)
else:
is_greedy = torch.ones(batch_size, dtype=torch.bool, device=self.device)
actions = q_out.argmax(-1)
return actions.tolist(), q_out.tolist(), masks_tensor.tolist(), is_greedy.tolist()
# ==========================================
# 4. 动态模型加载器
# ==========================================
def load_model(seat: int, model_file: str, arch: str) -> Bot:
""" 根据参数指定的架构,智能载入模型权重并构建对应的 Bot """
device = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')
print(f"Loading {model_file} (Arch: {arch}) on {device}")
state = torch.load(model_file, map_location=device)
version = state['config']['control']['version']
conv_channels = state['config']['resnet']['conv_channels']
num_blocks = state['config']['resnet']['num_blocks']
if arch.lower() == 'transformer':
brain = TransformerBrain(conv_channels=conv_channels, num_blocks=num_blocks, version=version).eval()
brain.transformer.enable_nested_tensor = False
dqn = TransformerDQN(version=version, d_model=conv_channels).eval()
brain_state = state.get('mortal', state.get('raw_brain', state))
dqn_state = state.get('current_dqn', state.get('raw_dqn', state))
engine = TransformerMortalEngine(
brain, dqn, is_oracle=False, version=version, device=device,
enable_amp=False, enable_quick_eval=False, enable_rule_based_agari_guard=True,
name='mortal-transformer', top_p=1,
)
elif arch.lower() == 'resnet':
brain = ResNetBrain(conv_channels=conv_channels, num_blocks=num_blocks, version=version).eval()
dqn = ResNetDQN(version=version).eval()
brain_state = state['mortal']
dqn_state = state['current_dqn']
engine = ResNetMortalEngine(
brain, dqn, is_oracle=False, version=version, device=device,
enable_amp=False, enable_quick_eval=False, enable_rule_based_agari_guard=True,
name='mortal-resnet', top_p=1,
)
else:
raise ValueError(f"Unknown architecture specified: {arch}")
brain.load_state_dict(brain_state)
dqn.load_state_dict(dqn_state)
return Bot(engine, seat)
# ==========================================
# 5. 高频对局逻辑
# ==========================================
def patch_event_fast(event_str):
if '"kita"' in event_str:
event_str = event_str.replace('"kita"', '"nukidora"')
if '"start_kyoku"' in event_str or '"deltas"' in event_str:
event = orjson.loads(event_str)
if event.get('type') == 'start_kyoku':
scores = event.setdefault('scores', [])
while len(scores) < 4: scores.append(0)
tehais = event.setdefault('tehais', [])
while len(tehais) < 4: tehais.append(["?" for _ in range(13)])
if 'deltas' in event:
deltas = event['deltas']
while len(deltas) < 4: deltas.append(0)
return orjson.dumps(event).decode('utf-8')
return event_str
def patch_resp_fast(resp_str):
if not resp_str: return resp_str
return resp_str.replace('"nukidora"', '"kita"')
_MODEL_CACHE = {}
def get_cached_model(player_id: int, model_file: str, arch: str):
key = (player_id, model_file, arch)
if key not in _MODEL_CACHE:
torch.set_num_threads(1)
_MODEL_CACHE[key] = load_model(player_id, model_file, arch)
return _MODEL_CACHE[key]
class MortalAgent:
def __init__(self, player_id: int, model_file: str, arch: str):
self.player_id = player_id
self.model = get_cached_model(player_id, model_file, arch)
def act(self, obs):
resp = None
for event in obs.new_events():
event_patched = patch_event_fast(event)
resp = patch_resp_fast(self.model.react(event_patched))
action = obs.select_action_from_mjai(resp)
assert action is not None, "Mortal must return a legal action"
return action
def play_three_games_with_seed(seed):
results = []
for test_seat in range(3):
env = RiichiEnv(game_mode="3p-red-half", rule=GameRule.default_tenhou())
agents = {}
for i in range(3):
# 将模型的架构标识也传入给 Agent 初始化进程
if i == test_seat:
agents[i] = MortalAgent(i, TEST_MODEL, TEST_ARCH)
else:
agents[i] = MortalAgent(i, EXAMINER_MODEL, EXAMINER_ARCH)
obs_dict = env.reset(seed=seed)
while not env.done():
actions = {pid: agents[pid].act(obs) for pid, obs in obs_dict.items()}
obs_dict = env.step(actions)
scores = env.scores()
ranks = env.ranks()
results.append((ranks[test_seat], scores[test_seat]))
return results
# ==========================================
# 6. 数据中心与后台独立评估线程
# ==========================================
def sync_models_from_hub():
if HF_TOKEN and "你的用户名" not in MODEL_REPO_ID:
print(f"☁️ 正在从模型仓库 [{MODEL_REPO_ID}] 拉取评估模型...")
try:
hf_hub_download(repo_id=MODEL_REPO_ID, filename=TEST_MODEL, repo_type="model", local_dir=".", token=HF_TOKEN)
hf_hub_download(repo_id=MODEL_REPO_ID, filename=EXAMINER_MODEL, repo_type="model", local_dir=".", token=HF_TOKEN)
print("🎉 模型环境准备完毕!")
except Exception as e:
print(f"❌ 拉取模型失败: {e}")
def sync_data_from_hub():
if HF_TOKEN and "你的用户名" not in DATA_REPO_ID:
try:
snapshot_download(repo_id=DATA_REPO_ID, repo_type="dataset", local_dir=".", allow_patterns=REPORT_FILE_PREFIX + "_*.txt", token=HF_TOKEN)
except Exception as e:
print(f"⚠️ 拉取历史战绩失败: {e}")
def sync_data_to_hub():
if HF_TOKEN and "你的用户名" not in DATA_REPO_ID:
try:
api.upload_file(path_or_fileobj=REPORT_FILE, path_in_repo=REPORT_FILE, repo_id=DATA_REPO_ID, repo_type="dataset", token=HF_TOKEN)
print(f"☁️ 节点 {WORKER_ID} 战绩已同步至 Hub: {time.strftime('%H:%M:%S')}")
except Exception as e:
print(f"❌ 同步失败: {e}")
def background_eval_loop():
sync_models_from_hub()
sync_data_from_hub()
NUM_WORKERS = 1
print(f"🚀 节点 [{WORKER_ID}] 后台对战线程已启动: {TEST_MODEL} ({TEST_ARCH}) 挑战双 {EXAMINER_MODEL} ({EXAMINER_ARCH})")
if not os.path.exists(REPORT_FILE): open(REPORT_FILE, 'w').close()
games_since_last_sync = 0
with concurrent.futures.ProcessPoolExecutor(max_workers=NUM_WORKERS) as executor:
futures = {executor.submit(play_three_games_with_seed, secrets.randbits(31)) for _ in range(NUM_WORKERS * 2)}
games_completed = 0
while EVAL_RUNNING and futures:
done, futures = concurrent.futures.wait(futures, return_when=concurrent.futures.FIRST_COMPLETED)
with open(REPORT_FILE, "a") as f:
for future in done:
try:
batch_results = future.result()
for rank, score in batch_results:
f.write(f"{rank} {score}\n")
games_completed += 1
games_since_last_sync += 1
print(f"[节点 {WORKER_ID}] 完成复式对局: 顺位 {rank}, 得点 {score}")
f.flush()
except Exception as e:
print(f"对局异常: {e}")
if EVAL_RUNNING:
futures.add(executor.submit(play_three_games_with_seed, secrets.randbits(31)))
if games_since_last_sync >= 60:
sync_data_to_hub()
sync_data_from_hub()
games_since_last_sync = 0
# ==========================================
# 7. 前端 Gradio 实时展示面板
# ==========================================
def read_and_analyze():
all_files = glob.glob(f"{REPORT_FILE_PREFIX}_*.txt")
if not all_files:
return f"⏳ 正在拉取模型,等待第一局完成...", None
ranks, scores = [], []
try:
for file in all_files:
with open(file, "r") as f:
lines = f.readlines()
for line in lines:
parts = line.strip().split()
if len(parts) == 2:
ranks.append(int(float(parts[0])))
scores.append(float(parts[1]))
total = len(ranks)
if total == 0:
return f"⏳ 模型已就绪,正在进行第一局对抗...", None
avg_rank = sum(ranks) / total
avg_score = sum(scores) / total
rank1_rate = ranks.count(1) / total * 100
rank2_rate = ranks.count(2) / total * 100
rank3_rate = ranks.count(3) / total * 100
last_update = time.strftime('%Y-%m-%d %H:%M:%S')
md_text = f"""
### 📊 对战简报
- ⚔️ **对抗阵容:** 1只 `{TEST_MODEL}` ({TEST_ARCH}) **VS** 2只 `{EXAMINER_MODEL}` ({EXAMINER_ARCH})
- 🧮 **总对局数:** {total} 局 (跨节点全局汇集)
- 🏆 **平均顺位:** {avg_rank:.3f}
- 💰 **平均得点:** {avg_score:.0f}
---
- 🥇 **一位率:** {rank1_rate:.1f}%
- 🥈 **二位率:** {rank2_rate:.1f}%
- 🥉 **三位率:** {rank3_rate:.1f}%
---
- 🌐 **当前节点 ID:** `{WORKER_ID}`
- 🕒 **刷新时间:** {last_update}
"""
fig = plt.figure(figsize=(10, 4))
ax1 = fig.add_subplot(121)
ax1.bar(['1st', '2nd', '3rd'], [rank1_rate, rank2_rate, rank3_rate], color=['#FFD700', '#C0C0C0', '#CD7F32'])
ax1.set_title(f'Rank Distribution for {TEST_MODEL}')
ax1.set_ylim(0, max(100, max([rank1_rate, rank2_rate, rank3_rate] + [0]) + 10))
for i, v in enumerate([rank1_rate, rank2_rate, rank3_rate]):
ax1.text(i, v + 2, f"{v:.1f}%", ha='center')
ax2 = fig.add_subplot(122)
df = pd.DataFrame({'score': scores})
df['ma'] = df['score'].rolling(window=min(10, max(1, len(df))), min_periods=1).mean()
ax2.plot(df['score'], alpha=0.3, color='gray', label='Raw Score')
ax2.plot(df['ma'], color='crimson', linewidth=2, label='Moving Avg (10)')
ax2.set_title('Score Trend')
ax2.legend()
plt.tight_layout()
return md_text, fig
except Exception as e:
return f"❌ 数据解析出错: {e}", None
with gr.Blocks() as demo:
gr.Markdown("# 🀄 Mahjong AI 基准评估舱")
gr.Markdown(f"当前正在评估: **{TEST_MODEL} ({TEST_ARCH})** 单挑两名 **{EXAMINER_MODEL} ({EXAMINER_ARCH})**。启动时会自动从 `Riichi-Model-Repo` 拉取权重。")
with gr.Row():
with gr.Column(scale=1):
stats_output = gr.Markdown("🚀 正在初始化基准环境并连接模型仓库...")
refresh_btn = gr.Button("🔄 手动刷新全局战绩")
with gr.Column(scale=2):
plot_output = gr.Plot()
demo.load(fn=read_and_analyze, inputs=None, outputs=[stats_output, plot_output])
timer = gr.Timer(15)
timer.tick(fn=read_and_analyze, inputs=None, outputs=[stats_output, plot_output])
refresh_btn.click(fn=read_and_analyze, inputs=None, outputs=[stats_output, plot_output])
if __name__ == "__main__":
t = threading.Thread(target=background_eval_loop, daemon=True)
t.start()
demo.queue().launch(server_name="0.0.0.0", server_port=7860, theme=gr.themes.Soft())