ffzeroHua commited on
Commit
9a37da3
·
verified ·
1 Parent(s): a8ecf76

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +5 -8
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
- print(f"⚡ [Worker {worker_idx}] 加载完毕,全速运转!")
 
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: