Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
25 changes: 14 additions & 11 deletions flow_grpo/rewards.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,6 @@
"""

import os
import pickle
from io import BytesIO

import numpy as np
Expand All @@ -16,6 +15,8 @@
from PIL import Image
from requests.adapters import HTTPAdapter, Retry

from reward_server.protocol import encode_general_request, encode_reward_3d_request

REWARD_3D = "reward_3d"
REWARD_GENERAL = "reward_general"
REWARD_TOTAL = "reward_total"
Expand Down Expand Up @@ -117,13 +118,14 @@ def _fn(images, prompts, metadata):
for video_batch, prompt_batch, trajectory_batch in zip(
videos_batched, prompts_batched, camera_trajectories_batched
):
data = {
"videos": video_batch,
"prompts": prompt_batch,
"camera_trajectories": trajectory_batch,
}
response = sess.post(url, data=pickle.dumps(data), timeout=2000)
response_data = pickle.loads(response.content)
payload = encode_reward_3d_request(
video_batch,
prompt_batch,
trajectory_batch,
)
response = sess.post(url, json=payload, timeout=2000)
response.raise_for_status()
response_data = response.json()
all_scores += response_data["outputs"]
if response_data.get("details"):
all_reconstruction_scores += [float(item["gs_score"]) for item in response_data["details"]]
Expand Down Expand Up @@ -195,9 +197,10 @@ def _fn(images, prompts, metadata):

all_scores = []
for image_batch, prompt_batch in zip(images_batched, prompts_batched):
data = {"images": image_batch, "prompts": prompt_batch}
response = sess.post(url, data=pickle.dumps(data), timeout=1000)
response_data = pickle.loads(response.content)
payload = encode_general_request(image_batch, prompt_batch)
response = sess.post(url, json=payload, timeout=1000)
response.raise_for_status()
response_data = response.json()
all_scores += response_data["outputs"]

return all_scores, {}
Expand Down
115 changes: 115 additions & 0 deletions reward_server/protocol.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,115 @@
"""JSON protocol helpers shared by reward clients and servers."""

import base64
import binascii
from typing import Any


def _require_list(value: Any, field: str) -> list:
if not isinstance(value, list):
raise ValueError(f"'{field}' must be a list")
return value


def _encode_bytes(value: Any, field: str) -> str:
if not isinstance(value, (bytes, bytearray, memoryview)):
raise ValueError(f"'{field}' entries must be bytes")
return base64.b64encode(bytes(value)).decode("ascii")


def _decode_bytes(value: Any, field: str) -> bytes:
if not isinstance(value, str):
raise ValueError(f"'{field}' entries must be base64 strings")
try:
return base64.b64decode(value, validate=True)
except (binascii.Error, ValueError) as exc:
raise ValueError(f"'{field}' contains invalid base64 data") from exc


def _validate_prompts(prompts: Any, batch_size: int) -> list[str]:
prompts = _require_list(prompts, "prompts")
if len(prompts) != batch_size:
raise ValueError("media and prompt batch sizes must match")
if not all(isinstance(prompt, str) for prompt in prompts):
raise ValueError("'prompts' entries must be strings")
return prompts


def _to_json_compatible(value: Any) -> Any:
if value is None or isinstance(value, (bool, int, float, str)):
return value
if isinstance(value, dict):
if not all(isinstance(key, str) for key in value):
raise ValueError("camera trajectory keys must be strings")
return {key: _to_json_compatible(item) for key, item in value.items()}
if isinstance(value, (list, tuple)):
return [_to_json_compatible(item) for item in value]
if hasattr(value, "tolist"):
return _to_json_compatible(value.tolist())
if hasattr(value, "item"):
return _to_json_compatible(value.item())
raise ValueError(f"camera trajectory value is not JSON serializable: {type(value).__name__}")


def encode_general_request(images: list[bytes], prompts: list[str]) -> dict[str, Any]:
"""Encode a general-reward batch as a JSON-compatible object."""
images = _require_list(images, "images")
prompts = _validate_prompts(prompts, len(images))
return {
"images": [_encode_bytes(image, "images") for image in images],
"prompts": prompts,
}


def decode_general_request(payload: Any) -> tuple[list[bytes], list[str]]:
"""Validate and decode a general-reward request."""
if not isinstance(payload, dict):
raise ValueError("request payload must be a JSON object")
images = _require_list(payload.get("images"), "images")
prompts = _validate_prompts(payload.get("prompts"), len(images))
return [_decode_bytes(image, "images") for image in images], prompts


def encode_reward_3d_request(
videos: list[list[bytes]],
prompts: list[str],
camera_trajectories: list[Any] | None = None,
) -> dict[str, Any]:
"""Encode a 3D-reward batch as a JSON-compatible object."""
videos = _require_list(videos, "videos")
prompts = _validate_prompts(prompts, len(videos))
encoded_videos = []
for video in videos:
video = _require_list(video, "videos")
encoded_videos.append([_encode_bytes(frame, "videos") for frame in video])

payload = {"videos": encoded_videos, "prompts": prompts}
if camera_trajectories is not None:
camera_trajectories = _require_list(camera_trajectories, "camera_trajectories")
if len(camera_trajectories) != len(videos):
raise ValueError("video and camera trajectory batch sizes must match")
payload["camera_trajectories"] = _to_json_compatible(camera_trajectories)
return payload


def decode_reward_3d_request(
payload: Any,
) -> tuple[list[list[bytes]], list[str], list[Any] | None]:
"""Validate and decode a 3D-reward request."""
if not isinstance(payload, dict):
raise ValueError("request payload must be a JSON object")
videos = _require_list(payload.get("videos"), "videos")
prompts = _validate_prompts(payload.get("prompts"), len(videos))

decoded_videos = []
for video in videos:
video = _require_list(video, "videos")
decoded_videos.append([_decode_bytes(frame, "videos") for frame in video])

camera_trajectories = payload.get("camera_trajectories")
if camera_trajectories is not None:
camera_trajectories = _require_list(camera_trajectories, "camera_trajectories")
if len(camera_trajectories) != len(videos):
raise ValueError("video and camera trajectory batch sizes must match")

return decoded_videos, prompts, camera_trajectories
32 changes: 19 additions & 13 deletions scripts/serve_general_reward.py
Original file line number Diff line number Diff line change
@@ -1,44 +1,50 @@
#!/usr/bin/env python3
import argparse
import os
import pickle
import traceback

from flask import Blueprint, Flask, request
from flask import Blueprint, Flask, current_app, jsonify, request

from reward_server.general_reward import MultiGPUGeneralRewardManager
from reward_server.protocol import decode_general_request

root = Blueprint("root", __name__)
general_reward_manager = None


def create_app():
def create_app(manager=None):
global general_reward_manager
general_reward_manager = MultiGPUGeneralRewardManager()
general_reward_manager.initialize()
if manager is None:
from reward_server.general_reward import MultiGPUGeneralRewardManager

manager = MultiGPUGeneralRewardManager()
manager.initialize()
general_reward_manager = manager
app = Flask(__name__)
app.register_blueprint(root)
return app


@root.route("/", methods=["POST"])
def inference():
data = request.get_data()
if not request.is_json:
return jsonify({"error": "Request body must be JSON"}), 400

try:
payload = pickle.loads(data)
batch_images = payload["images"]
batch_prompts = payload["prompts"]
batch_images, batch_prompts = decode_general_request(request.get_json(silent=True))
batch_size = len(batch_images)
except ValueError as exc:
return jsonify({"error": str(exc)}), 400

try:
global general_reward_manager
if general_reward_manager is None:
outputs = [0.5] * batch_size
else:
outputs = general_reward_manager.compute_batch_scores(batch_images, batch_prompts)

return pickle.dumps({"outputs": outputs}), 200
return jsonify({"outputs": [float(output) for output in outputs]}), 200
except Exception:
return traceback.format_exc().encode("utf-8"), 500
current_app.logger.exception("General reward computation failed")
return jsonify({"error": "General reward computation failed"}), 500


HOST = "127.0.0.1"
Expand Down
56 changes: 27 additions & 29 deletions scripts/serve_reward_3d.py
Original file line number Diff line number Diff line change
@@ -1,10 +1,8 @@
import pickle
import os
import argparse
import traceback
import signal
import sys
from reward_server.reward_3d import MultiGPUReward3DManager
from reward_server.protocol import decode_reward_3d_request

# import debugpy
# try:
Expand All @@ -15,7 +13,7 @@
# except Exception as e:
# pass

from flask import Flask, request, Blueprint
from flask import Blueprint, Flask, current_app, jsonify, request

root = Blueprint("root", __name__)

Expand All @@ -29,15 +27,20 @@ def signal_handler(sig, frame):
reward_3d_manager.shutdown()
sys.exit(0)

def create_app(scorer_type='qwen', use_lpips=True):
def create_app(scorer_type='qwen', use_lpips=True, manager=None, install_signal_handlers=True):
global reward_3d_manager
print(f"Initializing multi-GPU 3D reward server (scorer: {scorer_type}, lpips: {use_lpips})...")
reward_3d_manager = MultiGPUReward3DManager(scorer_type=scorer_type, use_lpips=use_lpips)
reward_3d_manager.initialize()
if manager is None:
from reward_server.reward_3d import MultiGPUReward3DManager

print(f"Initializing multi-GPU 3D reward server (scorer: {scorer_type}, lpips: {use_lpips})...")
manager = MultiGPUReward3DManager(scorer_type=scorer_type, use_lpips=use_lpips)
manager.initialize()
reward_3d_manager = manager

# Register signal handlers for graceful shutdown
signal.signal(signal.SIGINT, signal_handler)
signal.signal(signal.SIGTERM, signal_handler)
if install_signal_handlers:
signal.signal(signal.SIGINT, signal_handler)
signal.signal(signal.SIGTERM, signal_handler)

app = Flask(__name__)
app.register_blueprint(root)
Expand All @@ -46,20 +49,22 @@ def create_app(scorer_type='qwen', use_lpips=True):
@root.route("/", methods=["POST"])
def inference():
print(f"received POST request from {request.remote_addr}")
data = request.get_data()
if not request.is_json:
return jsonify({"error": "Request body must be JSON"}), 400

try:
# expects a dict with "videos" and "prompts"
# videos: List[List[bytes]] - outer list is batch_size, inner list is frames per video
# prompts: List[str] - text prompts for each video
data = pickle.loads(data)

batch_videos = data["videos"] # List[List[bytes]]
batch_prompts = data["prompts"] # List[str]
batch_camera_trajectories = data.get("camera_trajectories")
batch_videos, batch_prompts, batch_camera_trajectories = decode_reward_3d_request(
request.get_json(silent=True)
)
batch_size = len(batch_videos)
print(f"Got batch of size {batch_size} for 3D reward evaluation")
except ValueError as exc:
return jsonify({"error": str(exc)}), 400

try:
global reward_3d_manager
if reward_3d_manager is None:
print("Error: 3D reward server is not initialized")
Expand All @@ -75,20 +80,13 @@ def inference():

print(f"3D reward batch processing results: {outputs}")

response = {"outputs": outputs, "details": details}

# returns: a dict with "outputs"
# outputs: List of scores (float values) with length = batch_size
response = pickle.dumps(response)

returncode = 200
return jsonify({
"outputs": [float(output) for output in outputs],
"details": details,
}), 200
except Exception:
response = traceback.format_exc()
print(response)
response = response.encode("utf-8")
returncode = 500

return response, returncode
current_app.logger.exception("3D reward computation failed")
return jsonify({"error": "3D reward computation failed"}), 500


HOST = "127.0.0.1"
Expand Down
Loading