world-model / predict.py
ojaffe's picture
Upload folder using huggingface_hub
5a6434e verified
Raw
History Blame
2.38 kB
"""Prediction interface for Multi-Scale Flow-Warp-Mask U-Net v10 with TTA."""
import sys
import os
import numpy as np
import torch
sys.path.insert(0, "/home/coder/code")
from multiscale_flow_model import MultiScaleFlowUNet
from flownet_model import differentiable_warp
CONTEXT_LEN = 4
CHANNELS = [56, 112, 224]
def load_model(model_dir: str):
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = MultiScaleFlowUNet(in_channels=12, channels=CHANNELS)
model_path = os.path.join(model_dir, "model.pt")
state_dict = torch.load(model_path, map_location=device, weights_only=True)
state_dict = {k: v.float() for k, v in state_dict.items()}
model.load_state_dict(state_dict)
model.to(device)
model.eval()
return {"model": model, "device": device}
def _prepare_input(context_frames):
if len(context_frames) >= CONTEXT_LEN:
frames = context_frames[-CONTEXT_LEN:]
else:
pad_count = CONTEXT_LEN - len(context_frames)
padding = np.stack([context_frames[0]] * pad_count, axis=0)
frames = np.concatenate([padding, context_frames], axis=0)
frames_t = torch.from_numpy(frames.astype(np.float32) / 255.0)
frames_t = frames_t.permute(0, 3, 1, 2)
return frames_t
def _run_model(model, frames_t, device):
last_frame = frames_t[-1].unsqueeze(0)
inp = frames_t.reshape(1, -1, 64, 64)
inp = inp.to(device)
last_frame = last_frame.to(device)
flows, mask, gen_frame = model(inp)
# Use finest flow (last element)
flow = flows[-1]
warped = differentiable_warp(last_frame, flow)
pred = mask * warped + (1 - mask) * gen_frame
pred = torch.clamp(pred, 0, 1)
return pred
def predict_next_frame(model_dict, context_frames: np.ndarray) -> np.ndarray:
model = model_dict["model"]
device = model_dict["device"]
frames_t = _prepare_input(context_frames)
with torch.no_grad():
# Original prediction
pred1 = _run_model(model, frames_t, device)
# TTA: horizontally flipped prediction
frames_flipped = frames_t.flip(-1)
pred2_flipped = _run_model(model, frames_flipped, device)
pred2 = pred2_flipped.flip(-1)
# Average
pred = (pred1 + pred2) / 2.0
pred = pred[0].cpu().permute(1, 2, 0).numpy()
pred = (pred * 255).clip(0, 255).astype(np.uint8)
return pred