Upload folder using huggingface_hub
Browse files- __pycache__/predict.cpython-311.pyc +0 -0
- predict.py +6 -5
- sweep.py +42 -0
__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
|
@@ -155,7 +155,7 @@ def predict_next_frame(ens, context_frames: np.ndarray) -> np.ndarray:
|
|
| 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, residual_scale=1.
|
| 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)
|
|
@@ -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 = np.round(frame * 255 + 0.
|
| 178 |
ens.direct_cache.append(frame)
|
| 179 |
|
| 180 |
result = ens.direct_cache[ens.cache_step]
|
|
@@ -202,8 +202,9 @@ def predict_next_frame(ens, context_frames: np.ndarray) -> np.ndarray:
|
|
| 202 |
direct_flipped = torch.flip(direct_flipped, dims=[4])
|
| 203 |
direct_pred = (direct_orig + direct_flipped) / 2.0
|
| 204 |
|
| 205 |
-
# Multi-run AR with noise diversity
|
| 206 |
all_ar_runs = []
|
|
|
|
| 207 |
for noise_std in [0.0, 1.0/255.0, 2.0/255.0]:
|
| 208 |
ar_preds_run = []
|
| 209 |
ctx = context_tensor.clone()
|
|
@@ -240,7 +241,7 @@ def predict_next_frame(ens, context_frames: np.ndarray) -> np.ndarray:
|
|
| 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 + 0.
|
| 244 |
ens.direct_cache.append(frame)
|
| 245 |
|
| 246 |
result = ens.direct_cache[ens.cache_step]
|
|
@@ -272,7 +273,7 @@ def predict_next_frame(ens, context_frames: np.ndarray) -> np.ndarray:
|
|
| 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]
|
|
|
|
| 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, residual_scale=1.02)
|
| 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)
|
|
|
|
| 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 + 0.1).clip(0, 255).astype(np.uint8)
|
| 178 |
ens.direct_cache.append(frame)
|
| 179 |
|
| 180 |
result = ens.direct_cache[ens.cache_step]
|
|
|
|
| 202 |
direct_flipped = torch.flip(direct_flipped, dims=[4])
|
| 203 |
direct_pred = (direct_orig + direct_flipped) / 2.0
|
| 204 |
|
| 205 |
+
# Multi-run AR with noise diversity (fixed seed for reproducibility)
|
| 206 |
all_ar_runs = []
|
| 207 |
+
torch.manual_seed(2)
|
| 208 |
for noise_std in [0.0, 1.0/255.0, 2.0/255.0]:
|
| 209 |
ar_preds_run = []
|
| 210 |
ctx = context_tensor.clone()
|
|
|
|
| 241 |
ens.direct_cache = []
|
| 242 |
for i in range(PRED_FRAMES):
|
| 243 |
frame = np.transpose(predicted_np[i], (1, 2, 0))
|
| 244 |
+
frame = np.round(frame * 255 + 0.1).clip(0, 255).astype(np.uint8)
|
| 245 |
ens.direct_cache.append(frame)
|
| 246 |
|
| 247 |
result = ens.direct_cache[ens.cache_step]
|
|
|
|
| 273 |
ens.direct_cache = []
|
| 274 |
for i in range(PRED_FRAMES):
|
| 275 |
frame = np.transpose(predicted_np[i], (1, 2, 0))
|
| 276 |
+
frame = np.round(frame * 255 + 0.1).clip(0, 255).astype(np.uint8)
|
| 277 |
ens.direct_cache.append(frame)
|
| 278 |
|
| 279 |
result = ens.direct_cache[ens.cache_step]
|
sweep.py
ADDED
|
@@ -0,0 +1,42 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Sweep Pong AR residual scale."""
|
| 2 |
+
import subprocess
|
| 3 |
+
import json
|
| 4 |
+
import re
|
| 5 |
+
|
| 6 |
+
predict_path = "/home/coder/experiments/2026-04-12-330000-pong-amp-sweep/predict.py"
|
| 7 |
+
|
| 8 |
+
results = {}
|
| 9 |
+
for scale in [1.01, 1.02, 1.03, 1.04, 1.05, 1.06, 1.07]:
|
| 10 |
+
with open(predict_path, 'r') as f:
|
| 11 |
+
content = f.read()
|
| 12 |
+
|
| 13 |
+
# Replace Pong AR residual_scale (only in Pong section, identified by "pong" model)
|
| 14 |
+
content = re.sub(
|
| 15 |
+
r'_predict_ar_frame\(ens\.models\["pong"\], ctx, last_t, residual_scale=[\d.]+\)',
|
| 16 |
+
f'_predict_ar_frame(ens.models["pong"], ctx, last_t, residual_scale={scale})',
|
| 17 |
+
content
|
| 18 |
+
)
|
| 19 |
+
|
| 20 |
+
with open(predict_path, 'w') as f:
|
| 21 |
+
f.write(content)
|
| 22 |
+
|
| 23 |
+
result = subprocess.run(
|
| 24 |
+
['python', 'task/score.py', '--model_path', '/home/coder/experiments/2026-04-12-330000-pong-amp-sweep'],
|
| 25 |
+
capture_output=True, text=True, cwd='/home/coder'
|
| 26 |
+
)
|
| 27 |
+
|
| 28 |
+
for line in result.stdout.strip().split('\n'):
|
| 29 |
+
if '"score"' in line:
|
| 30 |
+
data = json.loads(line)
|
| 31 |
+
results[scale] = {
|
| 32 |
+
'score': data['score'],
|
| 33 |
+
'pong': data['per_game']['pong']['ssim'],
|
| 34 |
+
'sonic': data['per_game']['sonic']['ssim'],
|
| 35 |
+
'pp': data['per_game']['pole_position']['ssim']
|
| 36 |
+
}
|
| 37 |
+
print(f"Scale {scale}: overall={data['score']:.4f} pong={data['per_game']['pong']['ssim']:.4f}")
|
| 38 |
+
break
|
| 39 |
+
|
| 40 |
+
print("\n=== Summary ===")
|
| 41 |
+
best_scale = max(results.keys(), key=lambda s: results[s]['pong'])
|
| 42 |
+
print(f"Best Pong scale: {best_scale} with pong={results[best_scale]['pong']:.4f}, overall={results[best_scale]['score']:.4f}")
|