ojaffe commited on
Commit
eb84bdb
·
verified ·
1 Parent(s): 2ec6f44

Upload folder using huggingface_hub

Browse files
Files changed (2) hide show
  1. __pycache__/predict.cpython-311.pyc +0 -0
  2. predict.py +16 -19
__pycache__/predict.cpython-311.pyc CHANGED
Binary files a/__pycache__/predict.cpython-311.pyc and b/__pycache__/predict.cpython-311.pyc differ
 
predict.py CHANGED
@@ -151,23 +151,18 @@ def predict_next_frame(ens, context_frames: np.ndarray) -> np.ndarray:
151
 
152
  direct_pred = _predict_8frames_direct(ens.pong_direct, context_tensor, last_tensor)
153
 
154
- # Multi-run Pong AR with tiny noise diversity
155
- all_pong_ar_runs = []
156
- for noise_std in [0.0, 0.25/255.0, 0.5/255.0]:
157
- ar_preds_run = []
158
- ctx = context_tensor.clone()
159
- last_t = last_tensor.clone()
160
- for step in range(PRED_FRAMES):
161
- ctx_in = ctx if noise_std == 0 else torch.clamp(ctx + torch.randn_like(ctx) * noise_std, 0, 1)
162
- predicted = _predict_ar_frame(ens.models["pong"], ctx_in, last_t)
163
- ar_preds_run.append(predicted)
164
- ctx_frames = ctx.reshape(1, CONTEXT_FRAMES, 3, 64, 64)
165
- ctx_frames = torch.cat([ctx_frames[:, 1:], predicted.unsqueeze(1)], dim=1)
166
- ctx = ctx_frames.reshape(1, -1, 64, 64)
167
- last_t = predicted
168
- all_pong_ar_runs.append(torch.stack(ar_preds_run, dim=1))
169
 
170
- ar_pred = sum(all_pong_ar_runs) / len(all_pong_ar_runs)
171
 
172
  predicted = torch.zeros_like(direct_pred)
173
  for step in range(PRED_FRAMES):
@@ -236,10 +231,12 @@ def predict_next_frame(ens, context_frames: np.ndarray) -> np.ndarray:
236
  ar_pred = sum(all_ar_runs) / len(all_ar_runs)
237
 
238
  predicted = torch.zeros_like(direct_pred)
 
239
  for step in range(PRED_FRAMES):
240
- ar_weight = 0.65 - (step / (PRED_FRAMES - 1)) * 0.3
241
- direct_weight = 1.0 - ar_weight
242
- predicted[:, step] = ar_weight * ar_pred[:, step] + direct_weight * direct_pred[:, step]
 
243
 
244
  predicted_np = predicted[0].cpu().numpy()
245
  ens.direct_cache = []
 
151
 
152
  direct_pred = _predict_8frames_direct(ens.pong_direct, context_tensor, last_tensor)
153
 
154
+ ar_preds = []
155
+ ctx = context_tensor.clone()
156
+ last_t = last_tensor.clone()
157
+ for step in range(PRED_FRAMES):
158
+ predicted = _predict_ar_frame(ens.models["pong"], ctx, last_t)
159
+ ar_preds.append(predicted)
160
+ ctx_frames = ctx.reshape(1, CONTEXT_FRAMES, 3, 64, 64)
161
+ ctx_frames = torch.cat([ctx_frames[:, 1:], predicted.unsqueeze(1)], dim=1)
162
+ ctx = ctx_frames.reshape(1, -1, 64, 64)
163
+ last_t = predicted
 
 
 
 
 
164
 
165
+ ar_pred = torch.stack(ar_preds, dim=1)
166
 
167
  predicted = torch.zeros_like(direct_pred)
168
  for step in range(PRED_FRAMES):
 
231
  ar_pred = sum(all_ar_runs) / len(all_ar_runs)
232
 
233
  predicted = torch.zeros_like(direct_pred)
234
+ eps = 1e-6
235
  for step in range(PRED_FRAMES):
236
+ w = 0.65 - (step / (PRED_FRAMES - 1)) * 0.3
237
+ ar_safe = torch.clamp(ar_pred[:, step], eps, 1.0)
238
+ direct_safe = torch.clamp(direct_pred[:, step], eps, 1.0)
239
+ predicted[:, step] = torch.clamp(ar_safe.pow(w) * direct_safe.pow(1.0 - w), 0, 1)
240
 
241
  predicted_np = predicted[0].cpu().numpy()
242
  ens.direct_cache = []