Update app.py
Browse files
app.py
CHANGED
|
@@ -154,19 +154,16 @@ class DAggerEncoder:
|
|
| 154 |
print(f"⚠️ 上传失败: {e}")
|
| 155 |
|
| 156 |
# ==========================================
|
| 157 |
-
# [单局 DAgger 引擎]
|
| 158 |
# ==========================================
|
| 159 |
-
def play_dagger_game(encoder: DAggerEncoder):
|
| 160 |
env = RiichiEnv(game_mode="3p-red-half", rule=GameRule.default_tenhou())
|
| 161 |
|
| 162 |
-
# 每
|
| 163 |
-
# Student 负责打牌,Teacher 负责指点
|
| 164 |
-
student_bots = {i: load_model_student(i, STUDENT_MODEL) for i in range(3)}
|
| 165 |
-
teacher_bots = {i: load_model_teacher(i, TEACHER_MODEL) for i in range(3)}
|
| 166 |
-
|
| 167 |
-
# 纯 Python 状态机,负责截获 Student 视角的 1012维 Obs
|
| 168 |
python_states = {i: sanma_state.PlayerState(i) for i in range(3)}
|
| 169 |
|
|
|
|
|
|
|
| 170 |
obs_dict = env.reset()
|
| 171 |
|
| 172 |
while not env.done():
|
|
@@ -181,10 +178,10 @@ def play_dagger_game(encoder: DAggerEncoder):
|
|
| 181 |
# 1. 更新 Python 状态机以生成 1012 维特征
|
| 182 |
cans = python_states[pid].update(event_patched)
|
| 183 |
|
| 184 |
-
# 2. 幽灵教师思考 (
|
| 185 |
teacher_resp_str = teacher_bots[pid].react(event_patched)
|
| 186 |
|
| 187 |
-
# 3.
|
| 188 |
if cans.can_act:
|
| 189 |
obs_s, mask_s = python_states[pid].encode_obs(4, False)
|
| 190 |
|
|
@@ -193,16 +190,15 @@ def play_dagger_game(encoder: DAggerEncoder):
|
|
| 193 |
action_dict = json.loads(teacher_resp) if teacher_resp else None
|
| 194 |
label_idx = action_to_label(pid, action_dict)
|
| 195 |
|
| 196 |
-
# 只有当存在有效决策空间时,才记录
|
| 197 |
if int(np.count_nonzero(mask_s)) > 1:
|
| 198 |
encoder.append(obs_s, mask_s, label_idx)
|
| 199 |
except Exception:
|
| 200 |
pass
|
| 201 |
|
| 202 |
-
# 4. 学生思考 (
|
| 203 |
student_resp_str = student_bots[pid].react(event_patched)
|
| 204 |
|
| 205 |
-
# 提取学生
|
| 206 |
if student_resp_str and '"type":"none"' not in student_resp_str.replace(' ', ''):
|
| 207 |
student_final_action = env_obs.select_action_from_mjai(patch_resp_fast(student_resp_str))
|
| 208 |
|
|
@@ -211,27 +207,36 @@ def play_dagger_game(encoder: DAggerEncoder):
|
|
| 211 |
|
| 212 |
actions_to_env[pid] = student_final_action
|
| 213 |
|
| 214 |
-
# 环境
|
| 215 |
obs_dict = env.step(actions_to_env)
|
| 216 |
|
| 217 |
worker_status["games_played"] += 1
|
| 218 |
|
| 219 |
# ==========================================
|
| 220 |
-
# [后台循环]
|
| 221 |
# ==========================================
|
| 222 |
def background_dagger_loop():
|
| 223 |
sync_models_from_hub()
|
| 224 |
encoder = DAggerEncoder(chunk_size=2048)
|
| 225 |
-
worker_status["status"] = "Generating DAgger Data..."
|
| 226 |
|
| 227 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 228 |
while EVAL_RUNNING:
|
| 229 |
try:
|
| 230 |
-
|
|
|
|
| 231 |
except Exception as e:
|
| 232 |
print(f"Game crashed: {e}")
|
| 233 |
traceback.print_exc()
|
| 234 |
-
time.sleep(1)
|
| 235 |
|
| 236 |
# ==========================================
|
| 237 |
# [前端 UI]
|
|
|
|
| 154 |
print(f"⚠️ 上传失败: {e}")
|
| 155 |
|
| 156 |
# ==========================================
|
| 157 |
+
# [单局 DAgger 引擎] (已剥离模型加载)
|
| 158 |
# ==========================================
|
| 159 |
+
def play_dagger_game(encoder: DAggerEncoder, student_bots: dict, teacher_bots: dict):
|
| 160 |
env = RiichiEnv(game_mode="3p-red-half", rule=GameRule.default_tenhou())
|
| 161 |
|
| 162 |
+
# ⚠️ Python 层的特征状态机极其轻量,为了绝对安全,每局仍然重新实例化
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 163 |
python_states = {i: sanma_state.PlayerState(i) for i in range(3)}
|
| 164 |
|
| 165 |
+
# env.reset() 会生成包含 "start_game" 的初始事件流
|
| 166 |
+
# 底层的 Rust Bot 收到 "start_game" 后会自动清空上一局的残余记忆
|
| 167 |
obs_dict = env.reset()
|
| 168 |
|
| 169 |
while not env.done():
|
|
|
|
| 178 |
# 1. 更新 Python 状态机以生成 1012 维特征
|
| 179 |
cans = python_states[pid].update(event_patched)
|
| 180 |
|
| 181 |
+
# 2. 幽灵教师思考 (常驻内存的 Bot 接收事件)
|
| 182 |
teacher_resp_str = teacher_bots[pid].react(event_patched)
|
| 183 |
|
| 184 |
+
# 3. 抓取 DAgger 数据
|
| 185 |
if cans.can_act:
|
| 186 |
obs_s, mask_s = python_states[pid].encode_obs(4, False)
|
| 187 |
|
|
|
|
| 190 |
action_dict = json.loads(teacher_resp) if teacher_resp else None
|
| 191 |
label_idx = action_to_label(pid, action_dict)
|
| 192 |
|
|
|
|
| 193 |
if int(np.count_nonzero(mask_s)) > 1:
|
| 194 |
encoder.append(obs_s, mask_s, label_idx)
|
| 195 |
except Exception:
|
| 196 |
pass
|
| 197 |
|
| 198 |
+
# 4. 学生思考 (常驻内存的 Bot 接收事件)
|
| 199 |
student_resp_str = student_bots[pid].react(event_patched)
|
| 200 |
|
| 201 |
+
# 提取学生动作
|
| 202 |
if student_resp_str and '"type":"none"' not in student_resp_str.replace(' ', ''):
|
| 203 |
student_final_action = env_obs.select_action_from_mjai(patch_resp_fast(student_resp_str))
|
| 204 |
|
|
|
|
| 207 |
|
| 208 |
actions_to_env[pid] = student_final_action
|
| 209 |
|
| 210 |
+
# 环境步进
|
| 211 |
obs_dict = env.step(actions_to_env)
|
| 212 |
|
| 213 |
worker_status["games_played"] += 1
|
| 214 |
|
| 215 |
# ==========================================
|
| 216 |
+
# [后台循环] (常驻内存优化版)
|
| 217 |
# ==========================================
|
| 218 |
def background_dagger_loop():
|
| 219 |
sync_models_from_hub()
|
| 220 |
encoder = DAggerEncoder(chunk_size=2048)
|
|
|
|
| 221 |
|
| 222 |
+
worker_status["status"] = "Loading Models into Memory (One-time)..."
|
| 223 |
+
print("🧠 正在将 Student 和 Teacher 模型加载进内存 (常驻)...")
|
| 224 |
+
|
| 225 |
+
# 🚀 提速核心:在这里只加载一次,然后反复复用!
|
| 226 |
+
# 只要保证每局开头环境喂给它们 "start_game",Rust 层就会自动初始化
|
| 227 |
+
student_bots = {i: load_model_student(i, STUDENT_MODEL) for i in range(3)}
|
| 228 |
+
teacher_bots = {i: load_model_teacher(i, TEACHER_MODEL) for i in range(3)}
|
| 229 |
+
|
| 230 |
+
worker_status["status"] = "Generating DAgger Data at High Speed..."
|
| 231 |
+
print("⚡ 模型加载完毕,数据工厂开始全速运转!")
|
| 232 |
+
|
| 233 |
while EVAL_RUNNING:
|
| 234 |
try:
|
| 235 |
+
# 将内存中的 Bot 引用传给打牌函数
|
| 236 |
+
play_dagger_game(encoder, student_bots, teacher_bots)
|
| 237 |
except Exception as e:
|
| 238 |
print(f"Game crashed: {e}")
|
| 239 |
traceback.print_exc()
|
|
|
|
| 240 |
|
| 241 |
# ==========================================
|
| 242 |
# [前端 UI]
|