File size: 11,923 Bytes
ed0d2b6 719249d ed0d2b6 719249d ed0d2b6 719249d ed0d2b6 7b328e5 719249d 7b328e5 719249d 7b328e5 719249d ed0d2b6 719249d ed0d2b6 719249d ed0d2b6 719249d 64150db 719249d ed0d2b6 719249d ed0d2b6 719249d ed0d2b6 719249d ed0d2b6 719249d ed0d2b6 719249d ed0d2b6 719249d 64150db 719249d ed0d2b6 719249d ed0d2b6 719249d ed0d2b6 719249d ed0d2b6 719249d ed0d2b6 719249d 9a37da3 ed0d2b6 719249d ed0d2b6 719249d ed0d2b6 719249d 9a37da3 ed0d2b6 719249d ed0d2b6 719249d ed0d2b6 719249d ed0d2b6 719249d 2174b85 719249d 2174b85 719249d 2174b85 719249d 2174b85 719249d ed0d2b6 719249d ed0d2b6 719249d ed0d2b6 2174b85 719249d ed0d2b6 2174b85 719249d ed0d2b6 2174b85 719249d ed0d2b6 2174b85 719249d ed0d2b6 719249d 1a1396f 719249d 2174b85 719249d ed0d2b6 2174b85 ed0d2b6 05dcf78 a8ecf76 05dcf78 2174b85 9a37da3 2174b85 719249d ed0d2b6 2174b85 ed0d2b6 05dcf78 ed0d2b6 719249d ed0d2b6 719249d ed0d2b6 719249d | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 | import os
import json
import orjson
import pickle
import threading
import time
import uuid
import glob
import random
import traceback
import numpy as np
from datetime import datetime
import gradio as gr
from huggingface_hub import HfApi, hf_hub_download
# ==========================================
# [依赖导入]
# ==========================================
from riichienv import RiichiEnv, GameRule
# 导入两种不同架构的加载函数
from model3pLOCAL import load_model as load_model_teacher # 775 维 (Teacher: 9070)
from model3pNEW import load_model as load_model_student # 1012 维 (Student: 57000)
# 导入纯 Python 版的状态机,用于提取 1012 维特征
try:
from libriichiSanma import state as sanma_state
except ImportError:
import libriichi as sanma_state
# ==========================================
# [全局配置]
# ==========================================
DATASET_REPO = "AstraNASA/tenhou-scc" # 存放打包好的 DAgger 数据的仓库
MODEL_REPO_ID = "ffzeroHua/Riichi-Model-Repo"
HF_TOKEN = os.getenv("HF_TOKEN")
HF_TOKEN_2 = os.getenv("HF_TOKEN_2")
WORKER_ID = os.getenv("WORKER_ID", str(uuid.uuid4())[:6])
# 🚀 对战双方权重
STUDENT_MODEL = "StudentSanma_Distilled_Step57000.pth"
TEACHER_MODEL = "Elite4z9070.pth"
# 动作空间映射
MASK_3P = [
"1m", "2m", "3m", "4m", "5m", "6m", "7m", "8m", "9m",
"1p", "2p", "3p", "4p", "5p", "6p", "7p", "8p", "9p",
"1s", "2s", "3s", "4s", "5s", "6s", "7s", "8s", "9s",
"E", "S", "W", "N", "P", "F", "C",
'5mr', '5pr', '5sr',
'reach', 'pon', 'kan', 'nukidora', 'hora', 'ryukyoku', 'none'
]
NONE_CODE = MASK_3P.index('none')
KAN_CODE = MASK_3P.index('kan')
# UI 状态监控
worker_status = {
"games_played": 0,
"records_extracted": 0,
"chunks_uploaded": 0,
"status": "Starting DAgger Factory..."
}
api = HfApi(token=HF_TOKEN) if HF_TOKEN else None
EVAL_RUNNING = True
# ==========================================
# [辅助函数]
# ==========================================
def sync_models_from_hub():
if HF_TOKEN and "你的用户名" not in MODEL_REPO_ID:
print(f"☁️ 正在从 Hub 拉取模型...")
hf_hub_download(repo_id=MODEL_REPO_ID, filename=STUDENT_MODEL, repo_type="model", local_dir=".", token=HF_TOKEN_2)
hf_hub_download(repo_id=MODEL_REPO_ID, filename=TEACHER_MODEL, repo_type="model", local_dir=".", token=HF_TOKEN_2)
print("✅ 模型就绪!")
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"')
def action_to_label(who, action_dict):
if action_dict is None: return NONE_CODE
if action_dict.get('actor') != who or action_dict.get('type') == 'tsumo': return NONE_CODE
t = action_dict['type']
if t == 'dahai': return MASK_3P.index(action_dict['pai'])
if t in ('daiminkan', 'ankan', 'kakan'): return KAN_CODE
if t in MASK_3P: return MASK_3P.index(t)
return NONE_CODE
# ==========================================
# [DAgger 特征打包器]
# ==========================================
class DAggerEncoder:
def __init__(self, worker_idx, chunk_size=2048):
self.worker_idx = worker_idx # 记住自己的身份
self.chunk_size = chunk_size
self.inputs, self.outputs = [], []
self.local_pool_dir = f"dagger_pool_{WORKER_ID}"
os.makedirs(self.local_pool_dir, exist_ok=True)
def append(self, obs_s, mask_s, label):
self.inputs.append({
"obs_student": obs_s,
"mask_student": mask_s
})
self.outputs.append(label)
worker_status["records_extracted"] += 1
if len(self.inputs) >= self.chunk_size:
self.save_and_upload()
def save_and_upload(self):
filename = f"chunk_dagger_{datetime.now().strftime('%Y%m%d_%H%M%S')}_{WORKER_ID}_W{self.worker_idx}.pkl"
filepath = os.path.join(self.local_pool_dir, filename)
# 保存本地
with open(filepath, 'wb') as f:
pickle.dump({'inputs': self.inputs, 'outputs': self.outputs}, f)
self.inputs.clear()
self.outputs.clear()
# 上传到 Hub (使用独立的线程防止阻塞打牌)
if api and DATASET_REPO:
threading.Thread(target=self._upload_task, args=(filepath, filename)).start()
def _upload_task(self, filepath, filename):
try:
api.upload_file(
path_or_fileobj=filepath,
path_in_repo=f"dagger_chunks/{filename}", # 传到 dagger_chunks 文件夹下
repo_id=DATASET_REPO,
repo_type="dataset"
)
worker_status["chunks_uploaded"] += 1
os.remove(filepath)
print(f"☁️ [DAgger] 成功上传数据块: {filename}")
except Exception as e:
print(f"⚠️ 上传失败: {e}")
# ==========================================
# [单局 DAgger 引擎] (已剥离模型加载)
# ==========================================
def play_dagger_game(encoder: DAggerEncoder, student_bots: dict, teacher_bots: dict):
env = RiichiEnv(game_mode="3p-red-half", rule=GameRule.default_tenhou())
# ⚠️ Python 层的特征状态机极其轻量,为了绝对安全,每局仍然重新实例化
python_states = {i: sanma_state.PlayerState(i) for i in range(3)}
# env.reset() 会生成包含 "start_game" 的初始事件流
# 底层的 Rust Bot 收到 "start_game" 后会自动清空上一局的残余记忆
obs_dict = env.reset()
while not env.done():
actions_to_env = {}
for pid, env_obs in obs_dict.items():
student_final_action = None
for event_str in env_obs.new_events():
event_patched = patch_event_fast(event_str)
# 1. 更新 Python 状态机以生成 1012 维特征
cans = python_states[pid].update(event_patched)
# 2. 幽灵教师思考 (常驻内存的 Bot 接收事件)
teacher_resp_str = teacher_bots[pid].react(event_patched)
# 3. 抓取 DAgger 数据
if cans.can_act:
obs_s, mask_s = python_states[pid].encode_obs(4, False)
try:
teacher_resp = patch_resp_fast(teacher_resp_str)
action_dict = json.loads(teacher_resp) if teacher_resp else None
label_idx = action_to_label(pid, action_dict)
if int(np.count_nonzero(mask_s)) > 1:
encoder.append(obs_s, mask_s, label_idx)
except Exception:
pass
# 4. 学生思考 (常驻内存的 Bot 接收事件)
student_resp_str = student_bots[pid].react(event_patched)
# 提取学生动作
if student_resp_str and '"type":"none"' not in student_resp_str.replace(' ', ''):
student_final_action = env_obs.select_action_from_mjai(patch_resp_fast(student_resp_str))
if student_final_action is None:
student_final_action = env_obs.select_action_from_mjai('{"type": "none"}')
# 兜底 2: 如果不能 Pass(比如自己摸牌后必须切牌),强制选取合法动作
if student_final_action is None:
try:
# 获取当前环境允许的所有合法动作
legal_actions = env_obs.legal_actions
if callable(legal_actions):
legal_actions = legal_actions()
if legal_actions and len(legal_actions) > 0:
student_final_action = legal_actions[0] # 强行拿第一个合法动作兜底
except Exception as e:
pass
actions_to_env[pid] = student_final_action
# 环境步进
obs_dict = env.step(actions_to_env)
worker_status["games_played"] += 1
# ==========================================
# [后台循环] (常驻内存优化版)
# ==========================================
import multiprocessing
# 专门为子进程准备的常驻死循环
def dagger_worker_process(worker_idx):
import os
import time
import torch
torch.set_num_threads(1)
os.environ["OMP_NUM_THREADS"] = "1"
student_bots = {i: load_model_student(i, STUDENT_MODEL) for i in range(3)}
teacher_bots = {i: load_model_teacher(i, TEACHER_MODEL) for i in range(3)}
# 修改点 3:把 worker_idx 传进去
encoder = DAggerEncoder(worker_idx=worker_idx, chunk_size=2048)
while EVAL_RUNNING:
try:
play_dagger_game(encoder, student_bots, teacher_bots)
except Exception as e:
print(f"[Worker {worker_idx}] 牌局崩溃已兜底: {e}")
time.sleep(1)
# ==========================================
# [后台循环] (多进程压榨版)
# ==========================================
def background_dagger_loop():
sync_models_from_hub()
worker_status["status"] = "Starting Multiprocessing Factory..."
# HF Free Space 提供 2 个 vCPU,我们开 2 个进程刚好榨干它
NUM_WORKERS = 2
processes = []
# 启动多进程
for i in range(NUM_WORKERS):
p = multiprocessing.Process(target=dagger_worker_process, args=(i,))
p.daemon = True # 主进程死,子进程跟着死
p.start()
processes.append(p)
worker_status["status"] = f"Running with {NUM_WORKERS} concurrent workers."
# 守护进程
for p in processes:
p.join()
# ==========================================
# [前端 UI]
# ==========================================
def get_stats():
md = f"""
### ⚔️ DAgger 数据工厂运行中...
- **🏭 工厂状态:** {worker_status['status']}
- **🤖 掌控对局:** 3只 `{STUDENT_MODEL}` (1012 维)
- **👻 幽灵教练:** 3只 `{TEACHER_MODEL}` (9070 黑盒)
---
- **🀄 已完成对局:** {worker_status['games_played']} 局
- **🧠 提取的高质量决策:** {worker_status['records_extracted']} 条
- **📦 已上传数据块 (Chunks):** {worker_status['chunks_uploaded']} 个
- **🌐 节点 ID:** `{WORKER_ID}`
*提示:在你的 Colab 微调程序中,只需将 target_prefix 设为 `"dagger_chunks"` 即可开始训练!*
"""
return md
with gr.Blocks() as demo:
gr.Markdown("# 🀄 Mahjong DAgger 数据提纯引擎")
stats_output = gr.Markdown("⏳ 正在初始化 DAgger 引擎并拉取模型...")
demo.load(fn=get_stats, inputs=None, outputs=stats_output)
gr.Timer(5).tick(fn=get_stats, inputs=None, outputs=stats_output)
if __name__ == "__main__":
t = threading.Thread(target=background_dagger_loop, daemon=True)
t.start()
demo.queue().launch(server_name="0.0.0.0", server_port=7860) |