ojaffe commited on
Commit
571442e
·
verified ·
1 Parent(s): c998913

Upload folder using huggingface_hub

Browse files
__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
@@ -211,11 +211,12 @@ def predict_next_frame(ens, context_frames: np.ndarray) -> np.ndarray:
211
  ctx_flip = context_flipped.clone()
212
  last_t = last_tensor.clone()
213
  last_f = last_flipped.clone()
 
214
  for step in range(PRED_FRAMES):
215
  ctx_in = ctx if noise_std == 0 else torch.clamp(ctx + torch.randn_like(ctx) * noise_std, 0, 1)
216
  ctx_flip_in = ctx_flip if noise_std == 0 else torch.clamp(ctx_flip + torch.randn_like(ctx_flip) * noise_std, 0, 1)
217
- ar_orig = _predict_ar_frame(ens.sonic_ar, ctx_in, last_t, residual_scale=1.08)
218
- ar_flip = _predict_ar_frame(ens.sonic_ar, ctx_flip_in, last_f, residual_scale=1.08)
219
  ar_flip_back = torch.flip(ar_flip, dims=[3])
220
  ar_frame = (ar_orig + ar_flip_back) / 2.0
221
  ar_preds_run.append(ar_frame)
 
211
  ctx_flip = context_flipped.clone()
212
  last_t = last_tensor.clone()
213
  last_f = last_flipped.clone()
214
+ sonic_scales = [1.04, 1.04, 1.04, 1.08, 1.08, 1.08, 1.12, 1.12]
215
  for step in range(PRED_FRAMES):
216
  ctx_in = ctx if noise_std == 0 else torch.clamp(ctx + torch.randn_like(ctx) * noise_std, 0, 1)
217
  ctx_flip_in = ctx_flip if noise_std == 0 else torch.clamp(ctx_flip + torch.randn_like(ctx_flip) * noise_std, 0, 1)
218
+ ar_orig = _predict_ar_frame(ens.sonic_ar, ctx_in, last_t, residual_scale=sonic_scales[step])
219
+ ar_flip = _predict_ar_frame(ens.sonic_ar, ctx_flip_in, last_f, residual_scale=sonic_scales[step])
220
  ar_flip_back = torch.flip(ar_flip, dims=[3])
221
  ar_frame = (ar_orig + ar_flip_back) / 2.0
222
  ar_preds_run.append(ar_frame)