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())