link / linux_play.py
Boroborob's picture
Upload linux_play.py
6f249d0 verified
Raw
History Blame Contribute Delete
27 kB
import os
import sys
import time
import json
import argparse
from pathlib import Path
from collections import OrderedDict
import cv2
import numpy as np
from PIL import Image
from nitrogen.shared import BUTTON_ACTION_TOKENS, PATH_REPO
# Do not import nitrogen.game_env at module import time.
# game_env imports pyautogui/dxcam, which fails on headless Linux or non-Windows.
try:
from nitrogen.inference_viz import create_viz, VideoRecorder
_VIZ_IMPORT_ERROR = None
except Exception as exc: # Allows --no-video debugging even if video deps fail.
create_viz = None
VideoRecorder = None
_VIZ_IMPORT_ERROR = exc
class _NoOpGamepad:
def update(self):
pass
class _NoOpGamepadEmulator:
def __init__(self):
self.gamepad = _NoOpGamepad()
def press_button(self, button):
pass
def release_button(self, button):
pass
class MockGamepadEnv:
"""Linux/debug replacement for GamepadEnv.
It has the same small API surface used by this script:
reset/pause/unpause/step/close and gamepad_emulator.
step() returns a random RGB PIL image as obs.
"""
def __init__(
self,
game,
game_speed=1.0,
env_fps=60,
async_mode=True,
width=1280,
height=720,
seed=0,
):
self.game = game
self.game_speed = game_speed
self.env_fps = env_fps
self.async_mode = async_mode
self.width = width
self.height = height
self.rng = np.random.default_rng(seed)
self.step_count = 0
self.gamepad_emulator = _NoOpGamepadEmulator()
def _random_obs(self):
arr = self.rng.integers(
0,
256,
size=(self.height, self.width, 3),
dtype=np.uint8,
)
return Image.fromarray(arr, mode="RGB")
def reset(self):
self.step_count = 0
return self._random_obs()
def pause(self):
pass
def unpause(self):
pass
def close(self):
pass
def step(self, action):
self.step_count += 1
obs = self._random_obs()
reward = 0.0
terminated = False
truncated = False
info = {
"mock_env": True,
"step_count": self.step_count,
"game": self.game,
}
return obs, reward, terminated, truncated, info
class VisualDebugGamepadEnv:
"""Linux/debug environment with meaningful RGB observations.
This env does not need a display server. It renders a simple 2D scene with OpenCV:
- a controllable player dot
- a target dot
- recent button states
- left/right joystick values
Therefore action outputs visibly change the next observation/video.
"""
def __init__(
self,
game,
game_speed=1.0,
env_fps=60,
async_mode=True,
width=1280,
height=720,
seed=0,
):
self.game = game
self.game_speed = game_speed
self.env_fps = env_fps
self.async_mode = async_mode
self.width = width
self.height = height
self.rng = np.random.default_rng(seed)
self.step_count = 0
self.gamepad_emulator = _NoOpGamepadEmulator()
self.player = np.array([width * 0.25, height * 0.5], dtype=np.float32)
self.velocity = np.zeros(2, dtype=np.float32)
self.target = np.array([width * 0.78, height * 0.5], dtype=np.float32)
self.last_action = make_zero_action()
self.trail = []
@staticmethod
def _scalar(v):
if isinstance(v, np.ndarray):
return float(np.asarray(v).reshape(-1)[0])
return float(v)
def reset(self):
self.step_count = 0
self.player = np.array([self.width * 0.25, self.height * 0.5], dtype=np.float32)
self.velocity[:] = 0
self.target = np.array([
self.rng.uniform(self.width * 0.55, self.width * 0.9),
self.rng.uniform(self.height * 0.2, self.height * 0.8),
], dtype=np.float32)
self.trail = []
self.last_action = make_zero_action()
return self._render()
def pause(self):
pass
def unpause(self):
pass
def close(self):
pass
def step(self, action):
self.step_count += 1
self.last_action = action
lx = self._scalar(action.get("AXIS_LEFTX", 0)) / 32767.0
ly = self._scalar(action.get("AXIS_LEFTY", 0)) / 32767.0
rx = self._scalar(action.get("AXIS_RIGHTX", 0)) / 32767.0
rt = self._scalar(action.get("RIGHT_TRIGGER", 0)) / 255.0
lt = self._scalar(action.get("LEFT_TRIGGER", 0)) / 255.0
# D-pad also moves the player, making binary button mapping visible.
if action.get("DPAD_LEFT", 0):
lx -= 1.0
if action.get("DPAD_RIGHT", 0):
lx += 1.0
if action.get("DPAD_UP", 0):
ly -= 1.0
if action.get("DPAD_DOWN", 0):
ly += 1.0
speed = 10.0 + 18.0 * max(rt, 0.0)
damping = 0.78 if not action.get("SOUTH", 0) else 0.55
accel = np.array([lx, ly], dtype=np.float32) * speed
self.velocity = self.velocity * damping + accel
if action.get("EAST", 0):
# A visible jump/boost.
self.velocity += np.array([18.0, -12.0], dtype=np.float32)
if action.get("WEST", 0):
self.velocity *= 0.35
if action.get("NORTH", 0):
# Move target so button effects are visible even without player motion.
self.target = np.array([
self.rng.uniform(self.width * 0.1, self.width * 0.9),
self.rng.uniform(self.height * 0.15, self.height * 0.85),
], dtype=np.float32)
self.player += self.velocity
margin = 35
self.player[0] = np.clip(self.player[0], margin, self.width - margin)
self.player[1] = np.clip(self.player[1], margin, self.height - margin)
self.trail.append(tuple(self.player.astype(int)))
self.trail = self.trail[-80:]
dist = float(np.linalg.norm(self.player - self.target))
reward = -dist / max(self.width, self.height)
if dist < 45:
reward = 1.0
self.target = np.array([
self.rng.uniform(self.width * 0.1, self.width * 0.9),
self.rng.uniform(self.height * 0.15, self.height * 0.85),
], dtype=np.float32)
obs = self._render(lx=lx, ly=ly, rx=rx, lt=lt, rt=rt, reward=reward)
terminated = False
truncated = False
info = {
"debug_env": "visual",
"step_count": self.step_count,
"distance_to_target": dist,
"game": self.game,
}
return obs, reward, terminated, truncated, info
def _render(self, lx=0.0, ly=0.0, rx=0.0, lt=0.0, rt=0.0, reward=0.0):
img = np.zeros((self.height, self.width, 3), dtype=np.uint8)
img[:] = (28, 30, 36)
# Background grid.
for x in range(0, self.width, 80):
cv2.line(img, (x, 0), (x, self.height), (45, 48, 56), 1)
for y in range(0, self.height, 80):
cv2.line(img, (0, y), (self.width, y), (45, 48, 56), 1)
# A few fixed obstacles for spatial reference.
cv2.rectangle(img, (self.width // 2 - 80, 130), (self.width // 2 + 80, 210), (80, 80, 95), -1)
cv2.rectangle(img, (self.width // 2 - 120, self.height - 230), (self.width // 2 + 120, self.height - 150), (80, 80, 95), -1)
# Trail.
if len(self.trail) >= 2:
for i in range(1, len(self.trail)):
thickness = 1 + i // 25
cv2.line(img, self.trail[i - 1], self.trail[i], (90, 180, 255), thickness)
# Target and player.
target_xy = tuple(self.target.astype(int))
player_xy = tuple(self.player.astype(int))
cv2.circle(img, target_xy, 34, (70, 220, 100), -1)
cv2.circle(img, target_xy, 45, (70, 220, 100), 2)
cv2.circle(img, player_xy, 28, (70, 130, 255), -1)
cv2.circle(img, player_xy, 36, (230, 235, 245), 2)
# Direction vector from joystick.
tip = (int(self.player[0] + lx * 90), int(self.player[1] + ly * 90))
cv2.arrowedLine(img, player_xy, tip, (255, 220, 120), 5, tipLength=0.3)
# Text panel.
panel_x, panel_y = 30, 35
line_h = 32
cv2.putText(img, "Linux VisualDebugGamepadEnv", (panel_x, panel_y), cv2.FONT_HERSHEY_SIMPLEX, 0.9, (245, 245, 245), 2)
cv2.putText(img, f"step={self.step_count} reward={reward:.3f}", (panel_x, panel_y + line_h), cv2.FONT_HERSHEY_SIMPLEX, 0.8, (230, 230, 230), 2)
cv2.putText(img, f"left=({lx:+.2f},{ly:+.2f}) right_x={rx:+.2f} LT={lt:.2f} RT={rt:.2f}", (panel_x, panel_y + 2 * line_h), cv2.FONT_HERSHEY_SIMPLEX, 0.75, (230, 230, 230), 2)
pressed = [k for k, v in self.last_action.items() if not isinstance(v, np.ndarray) and bool(v)]
pressed_text = "pressed: " + (", ".join(pressed[:10]) if pressed else "none")
cv2.putText(img, pressed_text, (panel_x, panel_y + 3 * line_h), cv2.FONT_HERSHEY_SIMPLEX, 0.7, (230, 230, 230), 2)
# Button lamps.
button_names = ["WEST", "SOUTH", "EAST", "NORTH", "START", "BACK"]
bx = panel_x
by = panel_y + 5 * line_h
for idx, name in enumerate(button_names):
on = bool(self.last_action.get(name, 0))
color = (80, 220, 120) if on else (75, 75, 85)
x0 = bx + idx * 125
cv2.rectangle(img, (x0, by), (x0 + 105, by + 42), color, -1)
cv2.putText(img, name, (x0 + 8, by + 28), cv2.FONT_HERSHEY_SIMPLEX, 0.55, (10, 10, 10) if on else (210, 210, 210), 2)
return Image.fromarray(cv2.cvtColor(img, cv2.COLOR_BGR2RGB))
class GymRGBGamepadEnv:
"""Optional Gymnasium RGB-array environment for Linux.
Example:
pip install "gymnasium[classic-control]" pygame
python play_linux_debug_visual.py --debug-env gym --gym-id CartPole-v1 --mock-policy --max-steps 10
"""
def __init__(self, game, game_speed=1.0, env_fps=60, async_mode=True, gym_id="CartPole-v1", seed=0, **kwargs):
os.environ.setdefault("SDL_VIDEODRIVER", "dummy")
import gymnasium as gym
self.game = game
self.env_fps = env_fps
self.gym_id = gym_id
self.seed = seed
self.step_count = 0
self.gamepad_emulator = _NoOpGamepadEmulator()
self.env = gym.make(gym_id, render_mode="rgb_array")
self.discrete_n = getattr(self.env.action_space, "n", None)
@staticmethod
def _scalar(v):
if isinstance(v, np.ndarray):
return float(np.asarray(v).reshape(-1)[0])
return float(v)
def _render(self):
arr = self.env.render()
return Image.fromarray(np.asarray(arr, dtype=np.uint8), mode="RGB")
def reset(self):
self.step_count = 0
self.env.reset(seed=self.seed)
return self._render()
def pause(self):
pass
def unpause(self):
pass
def close(self):
self.env.close()
def _map_gamepad_to_gym_action(self, action):
lx = self._scalar(action.get("AXIS_LEFTX", 0)) / 32767.0
if self.discrete_n is not None:
if self.discrete_n == 2:
return 1 if lx > 0 else 0
pressed = [name for name in BUTTON_ACTION_TOKENS if action.get(name, 0) and name in action]
if pressed:
return min(BUTTON_ACTION_TOKENS.index(pressed[0]), self.discrete_n - 1)
return int(np.clip(round((lx + 1.0) * 0.5 * (self.discrete_n - 1)), 0, self.discrete_n - 1))
# Continuous fallback: repeat left stick into action shape.
shape = self.env.action_space.shape
low, high = self.env.action_space.low, self.env.action_space.high
raw = np.zeros(shape, dtype=np.float32)
raw.flat[0] = lx
return np.clip(raw, low, high)
def step(self, action):
self.step_count += 1
gym_action = self._map_gamepad_to_gym_action(action)
_, reward, terminated, truncated, info = self.env.step(gym_action)
if terminated or truncated:
self.env.reset()
info = dict(info)
info.update({"debug_env": "gym", "gym_id": self.gym_id, "gym_action": gym_action, "step_count": self.step_count})
return self._render(), float(reward), bool(terminated), bool(truncated), info
class MockModelClient:
"""Optional model-server replacement for quickly testing script plumbing."""
def __init__(self, action_len=4, seed=0):
self.action_len = action_len
self.rng = np.random.default_rng(seed)
def reset(self):
pass
def info(self):
return {
"action_downsample_ratio": 1,
"ckpt_path": "mock_linux_debug.ckpt",
}
def predict(self, obs):
n = self.action_len
return {
"j_left": self.rng.uniform(-1.0, 1.0, size=(n, 2)),
"j_right": self.rng.uniform(-1.0, 1.0, size=(n, 2)),
"buttons": self.rng.random(size=(n, len(BUTTON_ACTION_TOKENS))),
}
class NullVideoRecorder:
def __init__(self, *args, **kwargs):
pass
def __enter__(self):
return self
def __exit__(self, exc_type, exc, tb):
return False
def add_frame(self, frame):
pass
def parse_args():
parser = argparse.ArgumentParser(description="VLM Inference")
parser.add_argument("--process", type=str, default="celeste.exe", help="Game to play")
parser.add_argument("--allow-menu", action="store_true", help="Allow menu actions (Disabled by default)")
parser.add_argument("--host", type=str, default="10.245.92.193", help="Host/IP of model server")
parser.add_argument("--port", type=int, default=5555, help="Port for model server")
# Linux/debug helpers.
parser.add_argument(
"--mock-env",
action="store_true",
help="Use the Linux debug environment instead of the real GamepadEnv.",
)
parser.add_argument(
"--debug-env",
type=str,
default="gym",
choices=["visual", "random", "gym"],
help="Linux/debug env backend. visual is deterministic and action-responsive; random is pure noise; gym uses Gymnasium.",
)
parser.add_argument(
"--gym-id",
type=str,
default="CartPole-v1",
help="Gymnasium environment id used when --debug-env gym.",
)
parser.add_argument(
"--real-env",
action="store_true",
help="Force real GamepadEnv even on Linux. Useful only if you really configured GUI/game backend.",
)
parser.add_argument(
"--mock-policy",
type=bool,
default=False,
help="Use random policy outputs instead of connecting to the model server.",
)
parser.add_argument(
"--mock-seed",
type=int,
default=0,
help="Random seed for mock env/policy.",
)
parser.add_argument(
"--mock-width",
type=int,
default=1280,
help="Mock observation width.",
)
parser.add_argument(
"--mock-height",
type=int,
default=720,
help="Mock observation height.",
)
parser.add_argument(
"--mock-action-len",
type=int,
default=4,
help="Number of actions predicted by MockModelClient per step.",
)
parser.add_argument(
"--max-steps",
type=int,
default=100,
help="Stop after N outer-loop steps. 0 means run forever, matching the original script.",
)
parser.add_argument(
"--no-video",
action="store_true",
help="Do not write mp4 files. Useful when debugging without ffmpeg/video deps.",
)
return parser.parse_args()
def preprocess_img(main_image):
main_cv = cv2.cvtColor(np.array(main_image), cv2.COLOR_RGB2BGR)
final_image = cv2.resize(main_cv, (256, 256), interpolation=cv2.INTER_AREA)
return Image.fromarray(cv2.cvtColor(final_image, cv2.COLOR_BGR2RGB))
def make_zero_action():
return OrderedDict(
[
("WEST", 0),
("SOUTH", 0),
("BACK", 0),
("DPAD_DOWN", 0),
("DPAD_LEFT", 0),
("DPAD_RIGHT", 0),
("DPAD_UP", 0),
("GUIDE", 0),
("AXIS_LEFTX", np.array([0], dtype=np.int64)),
("AXIS_LEFTY", np.array([0], dtype=np.int64)),
("LEFT_SHOULDER", 0),
("LEFT_TRIGGER", np.array([0], dtype=np.int64)),
("AXIS_RIGHTX", np.array([0], dtype=np.int64)),
("AXIS_RIGHTY", np.array([0], dtype=np.int64)),
("LEFT_THUMB", 0),
("RIGHT_THUMB", 0),
("RIGHT_SHOULDER", 0),
("RIGHT_TRIGGER", np.array([0], dtype=np.int64)),
("START", 0),
("EAST", 0),
("NORTH", 0),
]
)
def load_policy(args):
if args.mock_policy:
policy = MockModelClient(action_len=args.mock_action_len, seed=args.mock_seed)
else:
from nitrogen.inference_client import ModelClient
policy = ModelClient(host=args.host, port=args.port)
# policy = ModelClient(port=args.port)
policy.reset()
policy_info = policy.info()
return policy, policy_info
def make_env(args, use_mock_env):
if use_mock_env:
if args.debug_env == "visual":
return VisualDebugGamepadEnv(
game=args.process,
game_speed=1.0,
env_fps=60,
async_mode=True,
width=args.mock_width,
height=args.mock_height,
seed=args.mock_seed,
)
if args.debug_env == "gym":
return GymRGBGamepadEnv(
game=args.process,
game_speed=1.0,
env_fps=60,
async_mode=True,
gym_id=args.gym_id,
seed=args.mock_seed,
)
return MockGamepadEnv(
game=args.process,
game_speed=1.0,
env_fps=60,
async_mode=True,
width=args.mock_width,
height=args.mock_height,
seed=args.mock_seed,
)
# Import only when the real environment is requested.
# On Linux this can fail because game_env imports pyautogui/dxcam.
from nitrogen.game_env import GamepadEnv
return GamepadEnv(
game=args.process,
game_speed=1.0,
env_fps=60,
async_mode=True,
)
def maybe_init_game_menu(args, env, use_mock_env):
if use_mock_env:
return
# These games require opening a menu to initialize the controller.
if args.process not in {"isaac-ng.exe", "Cuphead.exe"}:
return
print(f"GamepadEnv ready for {args.process} at {env.env_fps} FPS")
input("Press enter to create a virtual controller and start rollouts...")
for i in range(3):
print(f"{3 - i}...")
time.sleep(1)
def press(button):
env.gamepad_emulator.press_button(button)
env.gamepad_emulator.gamepad.update()
time.sleep(0.05)
env.gamepad_emulator.release_button(button)
env.gamepad_emulator.gamepad.update()
press("SOUTH")
for _ in range(5):
press("EAST")
time.sleep(0.3)
def main():
args = parse_args()
use_mock_env = args.mock_env or (sys.platform != "win32" and not args.real_env)
if use_mock_env:
print(f"Using Linux debug env backend: {args.debug_env}")
else:
print("Using real GamepadEnv.")
if args.no_video:
Recorder = NullVideoRecorder
else:
if VideoRecorder is None or create_viz is None:
raise RuntimeError(
"Failed to import nitrogen.inference_viz. Re-run with --no-video for lightweight debugging."
) from _VIZ_IMPORT_ERROR
Recorder = VideoRecorder
policy, policy_info = load_policy(args)
action_downsample_ratio = int(policy_info.get("action_downsample_ratio", 1))
ckpt_path = policy_info.get("ckpt_path", "unknown_ckpt")
CKPT_NAME = Path(ckpt_path).stem
NO_MENU = not args.allow_menu
PATH_DEBUG = PATH_REPO / "debug"
PATH_DEBUG.mkdir(parents=True, exist_ok=True)
PATH_OUT = (PATH_REPO / "out" / CKPT_NAME).resolve()
PATH_OUT.mkdir(parents=True, exist_ok=True)
BUTTON_PRESS_THRES = 0.5
TOKEN_SET = BUTTON_ACTION_TOKENS
class RolloutFinished(Exception):
pass
video_files = sorted(PATH_OUT.glob("*_DEBUG.mp4"))
if video_files:
existing_numbers = [f.name.split("_")[0] for f in video_files]
existing_numbers = [int(n) for n in existing_numbers if n.isdigit()]
next_number = max(existing_numbers) + 1 if existing_numbers else 1
else:
next_number = 1
PATH_MP4_DEBUG = PATH_OUT / f"{next_number:04d}_DEBUG.mp4"
PATH_MP4_CLEAN = PATH_OUT / f"{next_number:04d}_CLEAN.mp4"
PATH_ACTIONS = PATH_OUT / f"{next_number:04d}_ACTIONS.json"
zero_action = make_zero_action()
print("Model loaded, starting environment...")
if not use_mock_env:
for i in range(3):
print(f"{3 - i}...")
time.sleep(1)
env = make_env(args, use_mock_env=use_mock_env)
maybe_init_game_menu(args, env, use_mock_env=use_mock_env)
env.reset()
env.pause()
# Initial call to get state. In debug mode this returns a PIL RGB observation.
obs, reward, terminated, truncated, info = env.step(action=zero_action)
step_count = 0
with Recorder(str(PATH_MP4_DEBUG), fps=60, crf=32, preset="medium") as debug_recorder:
with Recorder(str(PATH_MP4_CLEAN), fps=60, crf=28, preset="medium") as clean_recorder:
try:
while True:
obs = preprocess_img(obs)
obs.save(PATH_DEBUG / f"{step_count:05d}.png")
pred = policy.predict(obs)
j_left = np.asarray(pred["j_left"])
j_right = np.asarray(pred["j_right"])
buttons = np.asarray(pred["buttons"])
n = len(buttons)
assert n == len(j_left) == len(j_right), "Mismatch in action lengths"
env_actions = []
for i in range(n):
move_action = zero_action.copy()
xl, yl = j_left[i]
xr, yr = j_right[i]
move_action["AXIS_LEFTX"] = np.array([int(xl * 32767)], dtype=np.int64)
move_action["AXIS_LEFTY"] = np.array([int(yl * 32767)], dtype=np.int64)
move_action["AXIS_RIGHTX"] = np.array([int(xr * 32767)], dtype=np.int64)
move_action["AXIS_RIGHTY"] = np.array([int(yr * 32767)], dtype=np.int64)
button_vector = buttons[i]
assert len(button_vector) == len(TOKEN_SET), (
"Button vector length does not match token set length"
)
for name, value in zip(TOKEN_SET, button_vector):
if "TRIGGER" in name:
move_action[name] = np.array([int(value * 255)], dtype=np.int64)
else:
move_action[name] = 1 if value > BUTTON_PRESS_THRES else 0
env_actions.append(move_action)
print(
f"Executing {len(env_actions)} actions, "
f"each action will be repeated {action_downsample_ratio} times"
)
for i, a in enumerate(env_actions):
if NO_MENU:
if a["START"]:
print("Model predicted start, disabling this action")
a["GUIDE"] = 0
a["START"] = 0
a["BACK"] = 0
for _ in range(action_downsample_ratio):
obs, reward, terminated, truncated, info = env.step(action=a)
# if terminated or truncated:
# if terminated:
# print(f"[STOP] terminated=True, balance failed at step={step_count}, substep={i}")
# if truncated:
# print(f"[STOP] truncated=True, episode time limit reached at step={step_count}, substep={i}")
# raise RolloutFinished
if not args.no_video:
obs_viz = np.array(obs).copy()
clean_viz = cv2.resize(
obs_viz,
(1920, 1080),
interpolation=cv2.INTER_AREA,
)
debug_viz = create_viz(
cv2.resize(
obs_viz,
(1280, 720),
interpolation=cv2.INTER_AREA,
),
i,
j_left,
j_right,
buttons,
token_set=TOKEN_SET,
)
debug_recorder.add_frame(debug_viz)
clean_recorder.add_frame(clean_viz)
with open(PATH_ACTIONS, "a") as f:
for i, a in enumerate(env_actions):
serializable_action = {}
for k, v in a.items():
if isinstance(v, np.ndarray):
serializable_action[k] = v.tolist()
else:
serializable_action[k] = v
serializable_action["step"] = step_count
serializable_action["substep"] = i
json.dump(serializable_action, f)
f.write("\n")
step_count += 1
if args.max_steps > 0 and step_count >= args.max_steps:
print(f"Reached --max-steps={args.max_steps}, exiting.")
break
except RolloutFinished:
print("Rollout finished. Closing environment...")
finally:
env.unpause()
env.close()
if __name__ == "__main__":
main()