ojaffe commited on
Commit
ab7c749
·
verified ·
1 Parent(s): 5e7fdcb

Upload folder using huggingface_hub

Browse files
Files changed (2) hide show
  1. __pycache__/predict.cpython-311.pyc +0 -0
  2. predict.py +3 -17
__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
@@ -174,7 +174,7 @@ def predict_next_frame(ens, context_frames: np.ndarray) -> np.ndarray:
174
  ens.direct_cache = []
175
  for i in range(PRED_FRAMES):
176
  frame = np.transpose(predicted_np[i], (1, 2, 0))
177
- frame = (frame * 255).clip(0, 255).astype(np.uint8)
178
  ens.direct_cache.append(frame)
179
 
180
  result = ens.direct_cache[ens.cache_step]
@@ -191,20 +191,6 @@ def predict_next_frame(ens, context_frames: np.ndarray) -> np.ndarray:
191
  return result
192
 
193
  ens.reset_cache()
194
-
195
- # Medium static detection: all frames vs first frame, max diff < 2/255
196
- max_diff = 0
197
- first_frame = frames[0]
198
- for i in range(1, len(frames)):
199
- diff = np.abs(frames[i].astype(np.float32) - first_frame.astype(np.float32)).mean()
200
- max_diff = max(max_diff, diff)
201
- if max_diff < 2.0 / 255.0:
202
- last_ctx_uint8 = (last_frame * 255).clip(0, 255).astype(np.uint8)
203
- ens.direct_cache = [last_ctx_uint8.copy() for _ in range(PRED_FRAMES)]
204
- result = ens.direct_cache[ens.cache_step]
205
- ens.cache_step += 1
206
- return result
207
-
208
  with torch.no_grad():
209
  context_tensor = torch.from_numpy(context).to(DEVICE)
210
  last_tensor = torch.from_numpy(last_frame_t).to(DEVICE)
@@ -254,7 +240,7 @@ def predict_next_frame(ens, context_frames: np.ndarray) -> np.ndarray:
254
  ens.direct_cache = []
255
  for i in range(PRED_FRAMES):
256
  frame = np.transpose(predicted_np[i], (1, 2, 0))
257
- frame = (frame * 255).clip(0, 255).astype(np.uint8)
258
  ens.direct_cache.append(frame)
259
 
260
  result = ens.direct_cache[ens.cache_step]
@@ -286,7 +272,7 @@ def predict_next_frame(ens, context_frames: np.ndarray) -> np.ndarray:
286
  ens.direct_cache = []
287
  for i in range(PRED_FRAMES):
288
  frame = np.transpose(predicted_np[i], (1, 2, 0))
289
- frame = (frame * 255).clip(0, 255).astype(np.uint8)
290
  ens.direct_cache.append(frame)
291
 
292
  result = ens.direct_cache[ens.cache_step]
 
174
  ens.direct_cache = []
175
  for i in range(PRED_FRAMES):
176
  frame = np.transpose(predicted_np[i], (1, 2, 0))
177
+ frame = np.round(frame * 255).clip(0, 255).astype(np.uint8)
178
  ens.direct_cache.append(frame)
179
 
180
  result = ens.direct_cache[ens.cache_step]
 
191
  return result
192
 
193
  ens.reset_cache()
 
 
 
 
 
 
 
 
 
 
 
 
 
 
194
  with torch.no_grad():
195
  context_tensor = torch.from_numpy(context).to(DEVICE)
196
  last_tensor = torch.from_numpy(last_frame_t).to(DEVICE)
 
240
  ens.direct_cache = []
241
  for i in range(PRED_FRAMES):
242
  frame = np.transpose(predicted_np[i], (1, 2, 0))
243
+ frame = np.round(frame * 255).clip(0, 255).astype(np.uint8)
244
  ens.direct_cache.append(frame)
245
 
246
  result = ens.direct_cache[ens.cache_step]
 
272
  ens.direct_cache = []
273
  for i in range(PRED_FRAMES):
274
  frame = np.transpose(predicted_np[i], (1, 2, 0))
275
+ frame = np.round(frame * 255).clip(0, 255).astype(np.uint8)
276
  ens.direct_cache.append(frame)
277
 
278
  result = ens.direct_cache[ens.cache_step]