Update app.py
Browse files
app.py
CHANGED
|
@@ -107,7 +107,8 @@ def action_to_label(who, action_dict):
|
|
| 107 |
# [DAgger 特征打包器]
|
| 108 |
# ==========================================
|
| 109 |
class DAggerEncoder:
|
| 110 |
-
def __init__(self, chunk_size=2048):
|
|
|
|
| 111 |
self.chunk_size = chunk_size
|
| 112 |
self.inputs, self.outputs = [], []
|
| 113 |
self.local_pool_dir = f"dagger_pool_{WORKER_ID}"
|
|
@@ -126,7 +127,7 @@ class DAggerEncoder:
|
|
| 126 |
self.save_and_upload()
|
| 127 |
|
| 128 |
def save_and_upload(self):
|
| 129 |
-
filename = f"chunk_dagger_{datetime.now().strftime('%Y%m%d_%H%M%S')}_{WORKER_ID}.pkl"
|
| 130 |
filepath = os.path.join(self.local_pool_dir, filename)
|
| 131 |
|
| 132 |
# 保存本地
|
|
@@ -233,22 +234,18 @@ import multiprocessing
|
|
| 233 |
|
| 234 |
# 专门为子进程准备的常驻死循环
|
| 235 |
def dagger_worker_process(worker_idx):
|
| 236 |
-
# 🚀 核心修复:在子进程内部重新 import,确保它一定能找到环境!
|
| 237 |
import os
|
| 238 |
import time
|
| 239 |
import torch
|
| 240 |
|
| 241 |
-
# 强制当前进程只使用 1 个线程,防止多进程抢占 CPU 导致死锁
|
| 242 |
torch.set_num_threads(1)
|
| 243 |
os.environ["OMP_NUM_THREADS"] = "1"
|
| 244 |
|
| 245 |
-
print(f"👷 [Worker {worker_idx}] 启动中,正在加载专属常驻模型...")
|
| 246 |
-
# 子进程独立加载模型,防止多进程共享 PyTorch 对象的显存/内存冲突
|
| 247 |
student_bots = {i: load_model_student(i, STUDENT_MODEL) for i in range(3)}
|
| 248 |
teacher_bots = {i: load_model_teacher(i, TEACHER_MODEL) for i in range(3)}
|
| 249 |
-
encoder = DAggerEncoder(chunk_size=2048)
|
| 250 |
|
| 251 |
-
|
|
|
|
| 252 |
|
| 253 |
while EVAL_RUNNING:
|
| 254 |
try:
|
|
|
|
| 107 |
# [DAgger 特征打包器]
|
| 108 |
# ==========================================
|
| 109 |
class DAggerEncoder:
|
| 110 |
+
def __init__(self, worker_idx, chunk_size=2048):
|
| 111 |
+
self.worker_idx = worker_idx # 记住自己的身份
|
| 112 |
self.chunk_size = chunk_size
|
| 113 |
self.inputs, self.outputs = [], []
|
| 114 |
self.local_pool_dir = f"dagger_pool_{WORKER_ID}"
|
|
|
|
| 127 |
self.save_and_upload()
|
| 128 |
|
| 129 |
def save_and_upload(self):
|
| 130 |
+
filename = f"chunk_dagger_{datetime.now().strftime('%Y%m%d_%H%M%S')}_{WORKER_ID}_W{self.worker_idx}.pkl"
|
| 131 |
filepath = os.path.join(self.local_pool_dir, filename)
|
| 132 |
|
| 133 |
# 保存本地
|
|
|
|
| 234 |
|
| 235 |
# 专门为子进程准备的常驻死循环
|
| 236 |
def dagger_worker_process(worker_idx):
|
|
|
|
| 237 |
import os
|
| 238 |
import time
|
| 239 |
import torch
|
| 240 |
|
|
|
|
| 241 |
torch.set_num_threads(1)
|
| 242 |
os.environ["OMP_NUM_THREADS"] = "1"
|
| 243 |
|
|
|
|
|
|
|
| 244 |
student_bots = {i: load_model_student(i, STUDENT_MODEL) for i in range(3)}
|
| 245 |
teacher_bots = {i: load_model_teacher(i, TEACHER_MODEL) for i in range(3)}
|
|
|
|
| 246 |
|
| 247 |
+
# 修改点 3:把 worker_idx 传进去
|
| 248 |
+
encoder = DAggerEncoder(worker_idx=worker_idx, chunk_size=2048)
|
| 249 |
|
| 250 |
while EVAL_RUNNING:
|
| 251 |
try:
|