dennny123's picture
Fix Superstar Saga Space build and simplify port
ee44f77
Raw
History Blame Contribute Delete
14.5 kB
import gradio as gr
from pathlib import Path
import cv2
import numpy as np
from PIL import Image
import time
import tempfile
import requests
import zipfile
from huggingface_hub import hf_hub_download
import mgba.image
from pygba import PyGBA
from pygba.utils import KEY_MAP
import spaces
from nitrogen.inference_session import InferenceSession
from nitrogen.shared import BUTTON_ACTION_TOKENS
ROM_URL = "https://github.com/vbaemulator/GBA-Roms/raw/main/Mario%20%26%20Luigi%20-%20Superstar%20Saga%20(USA).zip"
def preprocess_img(frame):
"""Convert GBA frame to 256x256 RGB PIL Image for model input"""
if isinstance(frame, Image.Image):
frame = np.array(frame)
if len(frame.shape) == 2:
frame = cv2.cvtColor(frame, cv2.COLOR_GRAY2RGB)
elif frame.shape[2] == 4:
frame = cv2.cvtColor(frame, cv2.COLOR_RGBA2RGB)
frame_resized = cv2.resize(frame, (256, 256), interpolation=cv2.INTER_AREA)
return Image.fromarray(frame_resized)
def framebuffer_to_array(framebuffer):
"""Convert the mGBA framebuffer to a standard HxWx3 RGB array."""
return np.array(framebuffer.to_pil().convert("RGB")).transpose(1, 0, 2)
def gamepad_to_gba_buttons(pred, button_threshold=0.5, joystick_threshold=0.3):
"""Convert model's gamepad prediction to GBA button presses"""
j_left, j_right, buttons = pred["j_left"], pred["j_right"], pred["buttons"]
pressed_buttons = []
if len(buttons) == 0:
return pressed_buttons
button_vals = buttons[0]
if len(button_vals) < len(BUTTON_ACTION_TOKENS):
return pressed_buttons
# D-Pad mapping (indices 1-4)
if button_vals[1] > button_threshold:
pressed_buttons.append("down")
if button_vals[2] > button_threshold:
pressed_buttons.append("left")
if button_vals[3] > button_threshold:
pressed_buttons.append("right")
if button_vals[4] > button_threshold:
pressed_buttons.append("up")
# Joystick fallback if no D-pad pressed
if not any(b in pressed_buttons for b in ["up", "down", "left", "right"]):
if len(j_left) > 0:
xl, yl = j_left[0]
if abs(xl) > joystick_threshold or abs(yl) > joystick_threshold:
if abs(xl) > abs(yl):
if xl > joystick_threshold:
pressed_buttons.append("right")
elif xl < -joystick_threshold:
pressed_buttons.append("left")
else:
if yl > joystick_threshold:
pressed_buttons.append("down")
elif yl < -joystick_threshold:
pressed_buttons.append("up")
# Action buttons
if button_vals[18] > button_threshold: # SOUTH -> A
pressed_buttons.append("A")
if button_vals[5] > button_threshold: # EAST -> B
pressed_buttons.append("B")
if button_vals[19] > button_threshold: # START
pressed_buttons.append("start")
if button_vals[0] > button_threshold: # BACK -> SELECT
pressed_buttons.append("select")
# Shoulder buttons
if button_vals[7] > button_threshold: # LEFT_SHOULDER -> L
pressed_buttons.append("L")
if button_vals[14] > button_threshold: # RIGHT_SHOULDER -> R
pressed_buttons.append("R")
# Alternative mappings
if button_vals[10] > button_threshold and "A" not in pressed_buttons: # NORTH -> A
pressed_buttons.append("A")
if button_vals[20] > button_threshold and "B" not in pressed_buttons: # WEST -> B
pressed_buttons.append("B")
return pressed_buttons
def run_action_frames(gba, buttons_to_press, frame_skip, button_hold_frames):
"""Press buttons briefly, then release for the remaining frames."""
actions = [KEY_MAP[button] for button in buttons_to_press if button in KEY_MAP]
if actions:
hold_frames = min(frame_skip, button_hold_frames)
gba.core.set_keys(*actions)
for _ in range(hold_frames):
gba.core.run_frame()
gba.core.clear_keys(*actions)
else:
hold_frames = 0
remaining_frames = max(frame_skip - hold_frames, 0)
for _ in range(remaining_frames):
gba.core.run_frame()
@spaces.GPU
def play_superstar_saga(
cfg_scale: float,
context_length: int,
max_steps: int,
frame_skip: int,
button_threshold: float,
display_every: int,
update_delay: float
):
"""Generator that yields frames while playing Mario & Luigi: Superstar Saga"""
# Download ROM from URL
yield None, "⏳ Downloading ROM file...", "Fetching from GitHub..."
try:
temp_dir = Path(tempfile.gettempdir())
rom_path = temp_dir / "SuperstarSaga.gba"
# Download and extract ROM if not already cached
if not rom_path.exists():
zip_path = temp_dir / "SuperstarSaga.zip"
# Download ZIP file
response = requests.get(ROM_URL, stream=True)
response.raise_for_status()
with open(zip_path, 'wb') as f:
for chunk in response.iter_content(chunk_size=8192):
f.write(chunk)
yield None, "⏳ Extracting ROM from ZIP...", "Unpacking game file..."
# Extract the GBA file from the ZIP
with zipfile.ZipFile(zip_path, 'r') as zip_ref:
# Find the .gba file in the ZIP
gba_files = [f for f in zip_ref.namelist() if f.endswith('.gba')]
if gba_files:
# Extract the first .gba file found
zip_ref.extract(gba_files[0], temp_dir)
extracted_path = temp_dir / gba_files[0]
# Rename to our expected path
extracted_path.rename(rom_path)
else:
yield None, "❌ No .gba file found in ZIP archive", None
return
# Clean up ZIP file
zip_path.unlink()
yield None, "✅ ROM downloaded and extracted successfully", None
time.sleep(0.5)
else:
yield None, "✅ Using cached ROM", None
time.sleep(0.3)
except Exception as e:
yield None, f"❌ Error downloading ROM: {str(e)}", None
return
# Download checkpoint from HuggingFace using HfFileSystem
yield None, "⏳ Downloading checkpoint from nvidia/NitroGen...", None
try:
ckpt_path = Path(
hf_hub_download(
repo_id="nvidia/NitroGen",
filename="ng.pt",
)
)
yield None, "✅ Checkpoint ready", None
time.sleep(0.3)
except Exception as e:
yield None, f"❌ Error downloading checkpoint: {str(e)}", None
return
# Initialize inference session
yield None, "⏳ Initializing inference session...", None
session = InferenceSession.from_ckpt(
str(ckpt_path),
cfg_scale=cfg_scale,
context_length=context_length
)
session.reset()
# Initialize PyGBA with PyGBAEnv wrapper (proper approach)
yield None, "⏳ Loading ROM into PyGBA...", None
try:
gba = PyGBA.load(str(rom_path))
framebuffer = mgba.image.Image(*gba.core.desired_video_dimensions())
gba.core.set_video_buffer(framebuffer)
gba.core.run_frame()
observation = framebuffer_to_array(framebuffer)
yield None, "✅ Game initialized", f"Screen: {observation.shape}"
time.sleep(0.3)
except Exception as e:
yield None, f"❌ Error initializing PyGBA: {str(e)}", None
return
# Display settings
width, height = 720, 480 # GBA 240x160 scaled 3x
step_count = 0
# Button timing
button_hold_frames = 4
yield None, f"🎮 Starting gameplay with {max_steps} max steps", None
time.sleep(1)
try:
while step_count < max_steps:
obs_processed = preprocess_img(observation)
pred = session.predict(obs_processed)
buttons_to_press = gamepad_to_gba_buttons(pred, button_threshold)
run_action_frames(gba, buttons_to_press, frame_skip, button_hold_frames)
observation = framebuffer_to_array(framebuffer)
# Yield display update at specified frequency
if step_count % display_every == 0:
frame_display = cv2.resize(
observation,
(width, height),
interpolation=cv2.INTER_NEAREST
)
# Create action info
action_info = f"**Step {step_count}/{max_steps}**\n\n"
action_info += f"🎮 **Buttons:** {', '.join(buttons_to_press) if buttons_to_press else 'None'}\n\n"
action_info += f"⚡ **Speed:** {frame_skip}x frame skip\n\n"
action_info += f"📊 **Progress:** {step_count/max_steps*100:.1f}%"
# Create stats info
stats_info = f"**Inference Details**\n\n"
if len(pred.get("buttons", [])) > 0:
button_vals = pred["buttons"][0]
active_buttons = [
f"{BUTTON_ACTION_TOKENS[i]}: {button_vals[i]:.2f}"
for i in range(min(len(button_vals), len(BUTTON_ACTION_TOKENS)))
if button_vals[i] > button_threshold
]
if active_buttons:
stats_info += "**Active Predictions:**\n"
stats_info += "\n".join(f"- {btn}" for btn in active_buttons[:5])
else:
stats_info += "No buttons above threshold"
# Yield frame and info
yield frame_display, action_info, stats_info
time.sleep(update_delay)
step_count += 1
except Exception as e:
yield None, f"⏹️ Stopped at step {step_count}", f"Reason: {str(e)}"
finally:
del session
# Create Gradio interface
with gr.Blocks(title="NitroGen Superstar Saga Player") as app:
gr.Markdown("# 🎮 NitroGen Mario & Luigi: Superstar Saga Player")
gr.Markdown("Stream Superstar Saga gameplay powered by NitroGen AI model")
with gr.Row():
with gr.Column(scale=1):
gr.Markdown("### 🤖 Model Settings")
gr.Markdown("**Model:** nvidia/NitroGen (ng.pt) - downloaded from Hugging Face Hub")
gr.Markdown("**ROM:** Automatically downloaded from configured URL")
cfg_input = gr.Slider(
label="CFG Scale",
minimum=0.0,
maximum=3.0,
value=1.0,
step=0.1,
info="Classifier-free guidance scale"
)
ctx_input = gr.Slider(
label="Context Length",
minimum=1,
maximum=32,
value=1,
step=1,
info="Number of past frames to use"
)
gr.Markdown("### ⚙️ Playback Settings")
max_steps_input = gr.Slider(
label="Max Steps",
minimum=100,
maximum=10000,
value=1000,
step=100,
info="Maximum inference steps"
)
frame_skip_input = gr.Slider(
label="Frame Skip",
minimum=1,
maximum=64,
value=16,
step=1,
info="Emulator frames per inference"
)
button_threshold_input = gr.Slider(
label="Button Threshold",
minimum=0.0,
maximum=1.0,
value=0.5,
step=0.05,
info="Threshold for button activation"
)
display_every_input = gr.Slider(
label="Display Every N Steps",
minimum=1,
maximum=10,
value=1,
step=1,
info="Update display frequency (1=every step, higher=faster but less frequent)"
)
update_delay_input = gr.Slider(
label="Update Delay (seconds)",
minimum=0.1,
maximum=3.0,
value=1.0,
step=0.1,
info="Wait time after each display update (higher=more time for the image to load)"
)
start_btn = gr.Button("🚀 Start Playing", variant="primary", size="lg")
with gr.Column(scale=2):
image_output = gr.Image(
label="Game Stream",
height=600,
interactive=False
)
with gr.Row():
with gr.Column():
action_output = gr.Markdown(
label="Actions",
value="**Waiting to start...**"
)
with gr.Column():
stats_output = gr.Markdown(
label="Statistics",
value="**No data yet**"
)
gr.Markdown("""
### 📝 Instructions
1. Adjust playback settings as needed
2. Click "🚀 Start Playing" to begin streaming
3. Game frames update in real-time with actions
**Automatic Setup:**
- **Model**: nvidia/NitroGen checkpoint (ng.pt) from HuggingFace Hub
- **ROM**: Mario & Luigi: Superstar Saga downloaded from configured URL
- Model and ROM are cached automatically for faster subsequent runs
**Tips:**
- **Display Every N Steps**: 1 = update every step, higher = faster but less frequent
- **Update Delay**: 1s default gives images time to load, reduce for faster updates
- **Frame Skip**: 16 = game runs 16 frames per inference (faster gameplay)
""")
start_btn.click(
fn=play_superstar_saga,
inputs=[
cfg_input,
ctx_input,
max_steps_input,
frame_skip_input,
button_threshold_input,
display_every_input,
update_delay_input
],
outputs=[image_output, action_output, stats_output],
show_progress="full"
)
if __name__ == "__main__":
app.launch(
server_name="0.0.0.0",
server_port=7860,
share=False
)