| 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 |
| from model3pNEW import load_model as load_model_student |
|
|
| |
| try: |
| from libriichiSanma import state as sanma_state |
| except ImportError: |
| import libriichi as sanma_state |
|
|
| |
| |
| |
| DATASET_REPO = "AstraNASA/tenhou-scc" |
| 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') |
|
|
| |
| 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 |
|
|
| |
| |
| |
| 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() |
| |
| |
| 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}", |
| 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}") |
|
|
| |
| |
| |
| def play_dagger_game(encoder: DAggerEncoder, student_bots: dict, teacher_bots: dict): |
| env = RiichiEnv(game_mode="3p-red-half", rule=GameRule.default_tenhou()) |
| |
| |
| python_states = {i: sanma_state.PlayerState(i) for i in range(3)} |
| |
| |
| |
| 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) |
| |
| |
| cans = python_states[pid].update(event_patched) |
| |
| |
| teacher_resp_str = teacher_bots[pid].react(event_patched) |
| |
| |
| 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 |
| |
| |
| 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"}') |
| |
| |
| 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)} |
| |
| |
| 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..." |
| |
| |
| 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() |
|
|
| |
| |
| |
| 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) |