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

Upload folder using huggingface_hub

Browse files
Files changed (2) hide show
  1. __pycache__/predict.cpython-311.pyc +0 -0
  2. predict.py +20 -22
__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,18 +151,23 @@ 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
- 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):
@@ -195,19 +200,12 @@ def predict_next_frame(ens, context_frames: np.ndarray) -> np.ndarray:
195
  context_tensor = torch.from_numpy(context).to(DEVICE)
196
  last_tensor = torch.from_numpy(last_frame_t).to(DEVICE)
197
 
 
198
  context_flipped = torch.flip(context_tensor, dims=[3])
199
  last_flipped = torch.flip(last_tensor, dims=[3])
200
-
201
- # Multi-run direct with noise diversity
202
- all_direct_runs = []
203
- for noise_std in [0.0, 0.5/255.0, 1.0/255.0]:
204
- ctx_in = context_tensor if noise_std == 0 else torch.clamp(context_tensor + torch.randn_like(context_tensor) * noise_std, 0, 1)
205
- ctx_flip_in = context_flipped if noise_std == 0 else torch.clamp(context_flipped + torch.randn_like(context_flipped) * noise_std, 0, 1)
206
- direct_orig = _predict_8frames_direct(ens.sonic_direct, ctx_in, last_tensor)
207
- direct_flipped = _predict_8frames_direct(ens.sonic_direct, ctx_flip_in, last_flipped)
208
- direct_flipped = torch.flip(direct_flipped, dims=[4])
209
- all_direct_runs.append((direct_orig + direct_flipped) / 2.0)
210
- direct_pred = sum(all_direct_runs) / len(all_direct_runs)
211
 
212
  # Multi-run AR with noise diversity
213
  all_ar_runs = []
 
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):
 
200
  context_tensor = torch.from_numpy(context).to(DEVICE)
201
  last_tensor = torch.from_numpy(last_frame_t).to(DEVICE)
202
 
203
+ direct_orig = _predict_8frames_direct(ens.sonic_direct, context_tensor, last_tensor)
204
  context_flipped = torch.flip(context_tensor, dims=[3])
205
  last_flipped = torch.flip(last_tensor, dims=[3])
206
+ direct_flipped = _predict_8frames_direct(ens.sonic_direct, context_flipped, last_flipped)
207
+ direct_flipped = torch.flip(direct_flipped, dims=[4])
208
+ direct_pred = (direct_orig + direct_flipped) / 2.0
 
 
 
 
 
 
 
 
209
 
210
  # Multi-run AR with noise diversity
211
  all_ar_runs = []