| """Optimized blend: Pong AR weight 0.85->0.65, Sonic unchanged 0.7->0.3.""" |
| import sys |
| import os |
| import numpy as np |
| import torch |
|
|
| sys.path.insert(0, "/home/coder/code") |
| from unet_model import UNet |
|
|
| CONTEXT_FRAMES = 8 |
| PRED_FRAMES = 8 |
| DEVICE = "cuda" if torch.cuda.is_available() else "cpu" |
|
|
|
|
| def detect_game(context_frames: np.ndarray) -> str: |
| first_8 = context_frames[:CONTEXT_FRAMES] |
| mean_val = first_8.mean() |
| std_val = first_8.std() |
| b_mean = first_8[:, :, :, 2].mean() |
| r_mean = first_8[:, :, :, 0].mean() |
| if mean_val > 100 and std_val < 80 and b_mean > r_mean * 1.5: |
| return "pole_position" |
| elif mean_val < 5 and 10 < std_val < 20: |
| return "pong" |
| else: |
| return "sonic" |
|
|
|
|
| def load_int8_state_dict(path, device): |
| """Load int8 quantized state dict and dequantize to float32.""" |
| quantized = torch.load(path, map_location='cpu', weights_only=False) |
| sd = {} |
| for k, v in quantized.items(): |
| if 'int8' in v: |
| sd[k] = (v['int8'].float() * v['scale']).to(device) |
| else: |
| sd[k] = v['float'].to(device) |
| return sd |
|
|
|
|
| class EnsembleModels: |
| def __init__(self): |
| self.models = {} |
| self.sonic_ar = None |
| self.sonic_direct = None |
| self.pong_direct = None |
| self.direct_cache = None |
| self.cache_step = 0 |
|
|
| def reset_cache(self): |
| self.direct_cache = None |
| self.cache_step = 0 |
|
|
|
|
| def load_model(model_dir: str): |
| ens = EnsembleModels() |
|
|
| |
| pong = UNet(in_channels=24, out_channels=3, |
| enc_channels=(32, 64, 128), bottleneck_channels=128, |
| upsample_mode="bilinear").to(DEVICE) |
| sd = torch.load(os.path.join(model_dir, "model_pong.pt"), |
| map_location=DEVICE, weights_only=True) |
| pong.load_state_dict({k: v.float() for k, v in sd.items()}) |
| pong.eval() |
| ens.models["pong"] = pong |
|
|
| |
| pong_direct = UNet(in_channels=24, out_channels=24, |
| enc_channels=(32, 64, 128), bottleneck_channels=128, |
| upsample_mode="bilinear").to(DEVICE) |
| sd = torch.load(os.path.join(model_dir, "model_pong_direct.pt"), |
| map_location=DEVICE, weights_only=True) |
| pong_direct.load_state_dict({k: v.float() for k, v in sd.items()}) |
| pong_direct.eval() |
| ens.pong_direct = pong_direct |
|
|
| |
| sonic_ar = UNet(in_channels=24, out_channels=3, |
| enc_channels=(48, 96, 192), bottleneck_channels=256, |
| upsample_mode="bilinear").to(DEVICE) |
| sd = torch.load(os.path.join(model_dir, "model_sonic_ar.pt"), |
| map_location=DEVICE, weights_only=True) |
| sonic_ar.load_state_dict({k: v.float() for k, v in sd.items()}) |
| sonic_ar.eval() |
| ens.sonic_ar = sonic_ar |
|
|
| |
| sonic_direct = UNet(in_channels=24, out_channels=24, |
| enc_channels=(48, 96, 192), bottleneck_channels=256, |
| upsample_mode="bilinear").to(DEVICE) |
| sd = load_int8_state_dict(os.path.join(model_dir, "model_sonic_direct.pt"), DEVICE) |
| sonic_direct.load_state_dict(sd) |
| sonic_direct.eval() |
| ens.sonic_direct = sonic_direct |
|
|
| |
| pp = UNet(in_channels=24, out_channels=24, |
| enc_channels=(24, 48, 96), bottleneck_channels=128, |
| upsample_mode="bilinear").to(DEVICE) |
| sd = torch.load(os.path.join(model_dir, "model_pole_position.pt"), |
| map_location=DEVICE, weights_only=True) |
| pp.load_state_dict({k: v.float() for k, v in sd.items()}) |
| pp.eval() |
| ens.models["pole_position"] = pp |
|
|
| return ens |
|
|
|
|
| def _predict_8frames_direct(model, context_tensor, last_tensor): |
| output = model(context_tensor) |
| residuals = output.reshape(1, PRED_FRAMES, 3, 64, 64) |
| last_expanded = last_tensor.unsqueeze(1).expand_as(residuals) |
| return torch.clamp(last_expanded + residuals, 0, 1) |
|
|
|
|
| def _predict_ar_frame(model, context_tensor, last_tensor): |
| residual = model(context_tensor) |
| return torch.clamp(last_tensor + residual, 0, 1) |
|
|
|
|
| def predict_next_frame(ens, context_frames: np.ndarray) -> np.ndarray: |
| game = detect_game(context_frames) |
| n = len(context_frames) |
|
|
| if n < CONTEXT_FRAMES: |
| padding = np.stack([context_frames[0]] * (CONTEXT_FRAMES - n), axis=0) |
| frames = np.concatenate([padding, context_frames], axis=0) |
| else: |
| frames = context_frames[-CONTEXT_FRAMES:] |
|
|
| frames_norm = frames.astype(np.float32) / 255.0 |
| frames_t = np.transpose(frames_norm, (0, 3, 1, 2)) |
| context = frames_t.reshape(1, -1, 64, 64) |
|
|
| last_frame = frames_norm[-1] |
| last_frame_t = np.transpose(last_frame, (2, 0, 1))[np.newaxis] |
|
|
| if game == "pong": |
| |
| if ens.direct_cache is not None and n > CONTEXT_FRAMES and ens.cache_step < PRED_FRAMES: |
| result = ens.direct_cache[ens.cache_step] |
| ens.cache_step += 1 |
| if ens.cache_step >= PRED_FRAMES: |
| ens.reset_cache() |
| return result |
|
|
| ens.reset_cache() |
| with torch.no_grad(): |
| context_tensor = torch.from_numpy(context).to(DEVICE) |
| last_tensor = torch.from_numpy(last_frame_t).to(DEVICE) |
|
|
| direct_pred = _predict_8frames_direct(ens.pong_direct, context_tensor, last_tensor) |
|
|
| ar_preds = [] |
| ctx = context_tensor.clone() |
| last_t = last_tensor.clone() |
| for step in range(PRED_FRAMES): |
| predicted = _predict_ar_frame(ens.models["pong"], ctx, last_t) |
| ar_preds.append(predicted) |
| ctx_frames = ctx.reshape(1, CONTEXT_FRAMES, 3, 64, 64) |
| ctx_frames = torch.cat([ctx_frames[:, 1:], predicted.unsqueeze(1)], dim=1) |
| ctx = ctx_frames.reshape(1, -1, 64, 64) |
| last_t = predicted |
|
|
| ar_pred = torch.stack(ar_preds, dim=1) |
|
|
| predicted = torch.zeros_like(direct_pred) |
| for step in range(PRED_FRAMES): |
| ar_weight = 0.85 - (step / (PRED_FRAMES - 1)) * 0.2 |
| direct_weight = 1.0 - ar_weight |
| predicted[:, step] = ar_weight * ar_pred[:, step] + direct_weight * direct_pred[:, step] |
|
|
| predicted_np = predicted[0].cpu().numpy() |
| ens.direct_cache = [] |
| for i in range(PRED_FRAMES): |
| frame = np.transpose(predicted_np[i], (1, 2, 0)) |
| frame = (frame * 255).clip(0, 255).astype(np.uint8) |
| ens.direct_cache.append(frame) |
|
|
| result = ens.direct_cache[ens.cache_step] |
| ens.cache_step += 1 |
| return result |
|
|
| elif game == "sonic": |
| |
| if ens.direct_cache is not None and n > CONTEXT_FRAMES and ens.cache_step < PRED_FRAMES: |
| result = ens.direct_cache[ens.cache_step] |
| ens.cache_step += 1 |
| if ens.cache_step >= PRED_FRAMES: |
| ens.reset_cache() |
| return result |
|
|
| ens.reset_cache() |
| with torch.no_grad(): |
| context_tensor = torch.from_numpy(context).to(DEVICE) |
| last_tensor = torch.from_numpy(last_frame_t).to(DEVICE) |
|
|
| direct_orig = _predict_8frames_direct(ens.sonic_direct, context_tensor, last_tensor) |
| context_flipped = torch.flip(context_tensor, dims=[3]) |
| last_flipped = torch.flip(last_tensor, dims=[3]) |
| direct_flipped = _predict_8frames_direct(ens.sonic_direct, context_flipped, last_flipped) |
| direct_flipped = torch.flip(direct_flipped, dims=[4]) |
| direct_pred = (direct_orig + direct_flipped) / 2.0 |
|
|
| ar_preds = [] |
| ctx = context_tensor.clone() |
| ctx_flip = context_flipped.clone() |
| last_t = last_tensor.clone() |
| last_f = last_flipped.clone() |
| for step in range(PRED_FRAMES): |
| ar_orig = _predict_ar_frame(ens.sonic_ar, ctx, last_t) |
| ar_flip = _predict_ar_frame(ens.sonic_ar, ctx_flip, last_f) |
| ar_flip_back = torch.flip(ar_flip, dims=[3]) |
| ar_frame = (ar_orig + ar_flip_back) / 2.0 |
| ar_preds.append(ar_frame) |
| ctx_frames = ctx.reshape(1, CONTEXT_FRAMES, 3, 64, 64) |
| ctx_frames = torch.cat([ctx_frames[:, 1:], ar_orig.unsqueeze(1)], dim=1) |
| ctx = ctx_frames.reshape(1, -1, 64, 64) |
| last_t = ar_orig |
| ctx_flip_frames = ctx_flip.reshape(1, CONTEXT_FRAMES, 3, 64, 64) |
| ctx_flip_frames = torch.cat([ctx_flip_frames[:, 1:], ar_flip.unsqueeze(1)], dim=1) |
| ctx_flip = ctx_flip_frames.reshape(1, -1, 64, 64) |
| last_f = ar_flip |
|
|
| ar_pred = torch.stack(ar_preds, dim=1) |
|
|
| predicted = torch.zeros_like(direct_pred) |
| for step in range(PRED_FRAMES): |
| ar_weight = 0.7 - (step / (PRED_FRAMES - 1)) * 0.4 |
| direct_weight = 1.0 - ar_weight |
| predicted[:, step] = ar_weight * ar_pred[:, step] + direct_weight * direct_pred[:, step] |
|
|
| predicted_np = predicted[0].cpu().numpy() |
| ens.direct_cache = [] |
| for i in range(PRED_FRAMES): |
| frame = np.transpose(predicted_np[i], (1, 2, 0)) |
| frame = (frame * 255).clip(0, 255).astype(np.uint8) |
| ens.direct_cache.append(frame) |
|
|
| result = ens.direct_cache[ens.cache_step] |
| ens.cache_step += 1 |
| return result |
|
|
| else: |
| |
| if ens.direct_cache is not None and n > CONTEXT_FRAMES and ens.cache_step < PRED_FRAMES: |
| result = ens.direct_cache[ens.cache_step] |
| ens.cache_step += 1 |
| if ens.cache_step >= PRED_FRAMES: |
| ens.reset_cache() |
| return result |
|
|
| ens.reset_cache() |
| with torch.no_grad(): |
| context_tensor = torch.from_numpy(context).to(DEVICE) |
| last_tensor = torch.from_numpy(last_frame_t).to(DEVICE) |
|
|
| predicted_orig = _predict_8frames_direct(ens.models["pole_position"], context_tensor, last_tensor) |
| context_flipped = torch.flip(context_tensor, dims=[3]) |
| last_flipped = torch.flip(last_tensor, dims=[3]) |
| predicted_flipped = _predict_8frames_direct(ens.models["pole_position"], context_flipped, last_flipped) |
| predicted_flipped = torch.flip(predicted_flipped, dims=[4]) |
| predicted = (predicted_orig + predicted_flipped) / 2.0 |
|
|
| predicted_np = predicted[0].cpu().numpy() |
| ens.direct_cache = [] |
| for i in range(PRED_FRAMES): |
| frame = np.transpose(predicted_np[i], (1, 2, 0)) |
| frame = (frame * 255).clip(0, 255).astype(np.uint8) |
| ens.direct_cache.append(frame) |
|
|
| result = ens.direct_cache[ens.cache_step] |
| ens.cache_step += 1 |
| return result |
|
|