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)