| 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 |
|
|
| |
| 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 |
|
|
| |
| TEST_MODEL = "4zv4_2.pth" |
| TEST_ARCH = "transformer" |
|
|
| EXAMINER_MODEL = "Elite4z9070.pth" |
| EXAMINER_ARCH = "resnet" |
|
|
| |
| |
| |
| try: |
| from libriichi3p.mjai import Bot |
| from libriichi3p.consts import obs_shape, oracle_obs_shape, ACTION_SPACE, GRP_SIZE |
| except: |
| |
| 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 模块,请确保环境已安装该依赖。") |
|
|
| |
| 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 |
|
|
| |
| |
| |
| 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 |
|
|
| |
| |
| |
| 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() |
|
|
| |
| |
| |
| 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() |
|
|
| |
| 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() |
|
|
|
|
| |
| |
| |
| 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) |
|
|
| |
| |
| |
| 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): |
| |
| 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 |
|
|
| |
| |
| |
| 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 |
|
|
| |
| |
| |
| 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()) |