Compare commits

...
Author SHA1 Message Date
kevin314 0a0cdebaba Add gpu selection 2026-02-05 04:02:49 +00:00
kevin314 cd9e7a8cac Add gpu provisioning 2026-02-04 07:39:31 +00:00
kevin314 87e19dad5a Add model resetting 2026-02-02 21:36:52 +00:00
kevin314 79195ad18e Add web app for matrix game 2026-02-01 00:12:07 +00:00
12 changed files with 3741 additions and 0 deletions
+3
View File
@@ -40,6 +40,9 @@ dist/
*.egg
eggs/
.eggs/
node_modules/
.vite
vite.config.js.timestamp-*.mjs
# MkDocs documentation
site/
+12
View File
@@ -0,0 +1,12 @@
<!DOCTYPE html>
<html lang="en">
<head>
<meta charset="UTF-8" />
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
<title>World Model Interface</title>
</head>
<body>
<div id="app"></div>
<script type="module" src="/src/main.js"></script>
</body>
</html>
File diff suppressed because it is too large Load Diff
+16
View File
@@ -0,0 +1,16 @@
{
"name": "wm-interface-frontend",
"private": true,
"version": "0.0.1",
"type": "module",
"scripts": {
"dev": "vite",
"build": "vite build",
"preview": "vite preview"
},
"devDependencies": {
"@sveltejs/vite-plugin-svelte": "^3.0.1",
"svelte": "^4.2.8",
"vite": "^5.0.11"
}
}
+18
View File
@@ -0,0 +1,18 @@
<svg width="252" height="105" viewBox="0 0 252 105" fill="none" xmlns="http://www.w3.org/2000/svg">
<path d="M89.4843 55.5457H101.361L87.7028 101H74.638L89.4843 55.5457Z" fill="#356CFF"/>
<path fill-rule="evenodd" clip-rule="evenodd" d="M96.0167 1.00057H112.645L118.583 48.273H104.924L103.737 39.7882H79.9827L67.5117 55.5457H85.3273L43.1638 101H28.3174L22.3789 55.5457H33.6621L38.4129 91.3031L58.604 68.2729H44.3515L96.0167 1.00057ZM100.768 13.1217L87.7028 29.4852H103.143L100.768 13.1217Z" fill="#356CFF"/>
<path d="M37.2252 1.00057L22.3789 48.273H36.0375L40.7884 30.6974L62.6727 30.6974L69.6727 21.0005L43.7576 21.0004L46.7269 11.9096L77.6727 11.9096L86 1.00057L37.2252 1.00057Z" fill="#356CFF"/>
<path fill-rule="evenodd" clip-rule="evenodd" d="M108.488 55.5457L94.2351 101C94.2351 101 105.518 101 120.959 101C136.399 101 144.078 93.0133 148.276 79.788C152.432 68.0157 153.027 55.5457 136.399 55.5457C119.771 55.5457 108.488 55.5457 108.488 55.5457ZM109.081 90.697L116.802 65.8487C116.802 65.8487 120.959 65.8487 132.242 65.8487C143.525 65.8487 137.586 78.5759 135.211 84.0304C133.307 88.4021 127.491 90.697 122.74 90.697C117.989 90.697 109.081 90.697 109.081 90.697Z" fill="#356CFF"/>
<path d="M173.188 1.00056L168.625 11.9096C168.625 11.9096 149.386 11.9092 142.525 11.9095C135.664 11.9098 136.586 20.3944 141.337 20.3944H159.747C168.654 20.3944 166.961 33.6899 163.904 38.5761C160.467 44.0675 157.371 48.273 148.463 48.273L125.188 48.273L124 37.97L147.87 37.97C153.808 37.97 156.184 29.4852 151.433 29.4852H131.836C120.142 29.4852 125.897 1.00043 141.337 1.00043L173.188 1.00056Z" fill="#356CFF"/>
<path d="M179.938 1.00056L175.688 11.9096L191.221 11.9096L179.938 48.273H192.409L203.692 11.9096L219.132 11.9095L223.289 1.00043L179.938 1.00056Z" fill="#356CFF"/>
<path d="M161.341 55.5457H202.845L198.5 65.8487H169.654L167.279 73.7268H188.5L184.749 82.8177H164.31L161.934 90.697H190.251L186.624 101H146.494L161.341 55.5457Z" fill="#356CFF"/>
<path fill-rule="evenodd" clip-rule="evenodd" d="M230.821 54.9391C255.169 54.9391 251.776 67.0602 249.231 77.9692C246.686 88.8783 240.917 101 217.757 101C194.596 101 195.606 88.8783 199.347 77.9692C203.089 67.0602 206.473 54.9391 230.821 54.9391ZM237.948 77.9692C239.984 70.6965 240.917 65.242 228.446 65.242C215.975 65.242 211.818 71.9087 210.037 77.9692C208.255 84.0298 208.255 91.3025 219.538 91.3025C230.821 91.3025 235.911 85.2419 237.948 77.9692Z" fill="#356CFF"/>
<path d="M173.188 1.00056L168.625 11.9096C168.625 11.9096 149.386 11.9092 142.525 11.9095C135.664 11.9098 136.586 20.3944 141.337 20.3944M173.188 1.00056C173.188 1.00056 156.777 1.00043 141.337 1.00043M173.188 1.00056L141.337 1.00043M141.337 20.3944C146.088 20.3944 150.839 20.3944 159.747 20.3944M141.337 20.3944H159.747M159.747 20.3944C168.654 20.3944 166.961 33.6899 163.904 38.5761C160.467 44.0675 157.371 48.273 148.463 48.273M148.463 48.273C139.556 48.273 125.188 48.273 125.188 48.273M148.463 48.273L125.188 48.273M125.188 48.273L124 37.97M124 37.97C124 37.97 141.931 37.97 147.87 37.97M124 37.97L147.87 37.97M147.87 37.97C153.808 37.97 156.184 29.4852 151.433 29.4852M151.433 29.4852C146.682 29.4852 138.962 29.4852 131.836 29.4852M151.433 29.4852H131.836M131.836 29.4852C120.142 29.4852 125.897 1.00043 141.337 1.00043M37.2252 1.00057L22.3789 48.273H36.0375L40.7884 30.6974L62.6727 30.6974L69.6727 21.0005L43.7576 21.0004L46.7269 11.9096L77.6727 11.9096L86 1.00057L37.2252 1.00057ZM96.0167 1.00057H112.645L118.583 48.273H104.924L103.737 39.7882H79.9827L67.5117 55.5457H85.3273L43.1638 101H28.3174L22.3789 55.5457H33.6621L38.4129 91.3031L58.604 68.2729H44.3515L96.0167 1.00057ZM87.7028 29.4852L100.768 13.1217L103.143 29.4852H87.7028ZM89.4843 55.5457H101.361L87.7028 101H74.638L89.4843 55.5457ZM108.488 55.5457L94.2351 101C94.2351 101 105.518 101 120.959 101C136.399 101 144.078 93.0133 148.276 79.788C152.432 68.0157 153.027 55.5457 136.399 55.5457C119.771 55.5457 108.488 55.5457 108.488 55.5457ZM116.802 65.8487L109.081 90.697C109.081 90.697 117.989 90.697 122.74 90.697C127.491 90.697 133.307 88.4021 135.211 84.0304C137.586 78.5759 143.525 65.8487 132.242 65.8487C120.959 65.8487 116.802 65.8487 116.802 65.8487ZM179.938 1.00056L175.688 11.9096L191.221 11.9096L179.938 48.273H192.409L203.692 11.9096L219.132 11.9095L223.289 1.00043L179.938 1.00056ZM161.341 55.5457H202.845L198.5 65.8487H169.654L167.279 73.7268H188.5L184.749 82.8177H164.31L161.934 90.697H190.251L186.624 101H146.494L161.341 55.5457ZM230.821 54.9391C255.169 54.9391 251.776 67.0602 249.231 77.9692C246.686 88.8783 240.917 101 217.757 101C194.596 101 195.606 88.8783 199.347 77.9692C203.089 67.0602 206.473 54.9391 230.821 54.9391ZM228.446 65.242C240.917 65.242 239.984 70.6965 237.948 77.9692C235.911 85.2419 230.821 91.3025 219.538 91.3025C208.255 91.3025 208.255 84.0298 210.037 77.9692C211.818 71.9087 215.975 65.242 228.446 65.242Z" stroke="#356CFF" stroke-width="1.18771"/>
<path d="M15.2524 55.5451L21.191 100.999L24.7541 100.999L18.8156 55.5451L15.2524 55.5451Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.18771"/>
<path d="M8.12646 55.5451L14.065 100.999L15.2527 100.999L9.31417 55.5451L8.12646 55.5451Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.18771"/>
<path d="M1 55.5451L6.93853 100.999L7.53239 100.999L1.59385 55.5451L1 55.5451Z" fill="#356CFF" stroke="#356CFF" stroke-width="0.593853"/>
<path d="M15.2524 48.2724L30.0988 1H33.6619L18.8156 48.2724H15.2524Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.18771"/>
<path d="M8.12646 48.2724L22.9728 1H24.1605L9.31417 48.2724H8.12646Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.18771"/>
<path d="M1 48.2724L15.8463 1H16.4402L1.59385 48.2724H1Z" fill="#356CFF" stroke="#356CFF" stroke-width="0.593853"/>
<path d="M85.3271 55.5457H67.5116L87 12.7363L44.3513 68.2729H58.6038L43.1636 101L85.3271 55.5457Z" fill="#FDC717" stroke="#FDC717" stroke-width="1.18771" stroke-miterlimit="16"/>
</svg>

After

Width:  |  Height:  |  Size: 5.7 KiB

File diff suppressed because it is too large Load Diff
+8
View File
@@ -0,0 +1,8 @@
import App from './App.svelte'
import './style.css'
const app = new App({
target: document.getElementById('app')
})
export default app
+20
View File
@@ -0,0 +1,20 @@
* {
margin: 0;
padding: 0;
box-sizing: border-box;
}
body {
font-family: -apple-system, BlinkMacSystemFont, 'Segoe UI', Roboto, Oxygen, Ubuntu, Cantarell, sans-serif;
background: #1a1a1a;
color: white;
min-height: 100vh;
overflow-y: auto;
}
#app {
width: 100%;
min-height: 100vh;
display: flex;
flex-direction: column;
}
+18
View File
@@ -0,0 +1,18 @@
import { defineConfig } from 'vite'
import { svelte } from '@sveltejs/vite-plugin-svelte'
export default defineConfig({
plugins: [svelte()],
server: {
port: 5173,
proxy: {
'/ws': {
target: 'http://localhost:8000',
ws: true,
},
'/status': {
target: 'http://localhost:8000',
},
},
}
})
+35
View File
@@ -0,0 +1,35 @@
# Model configuration
MODEL_CONFIG = {
"model_path": "FastVideo/Matrix-Game-2.0-Base-Diffusers",
"keyboard_dim": 4,
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/universal/0000.png",
}
# Keyboard mappings (WASD)
KEYBOARD_MAP = {
"w": [1, 0, 0, 0],
"s": [0, 1, 0, 0],
"a": [0, 0, 1, 0],
"d": [0, 0, 0, 1],
}
# Camera mappings (Arrow keys)
CAM_VALUE = 0.1
CAMERA_MAP = {
"ArrowUp": [CAM_VALUE, 0],
"ArrowDown": [-CAM_VALUE, 0],
"ArrowLeft": [0, -CAM_VALUE],
"ArrowRight": [0, CAM_VALUE],
}
# Generation limits
MAX_BLOCKS = 50
SESSION_TIMEOUT_SECONDS = 90
# Frame settings
NUM_FRAMES = 597
FRAME_HEIGHT = 352
FRAME_WIDTH = 640
NUM_INFERENCE_STEPS = 10
JPEG_QUALITY = 85
BATCH_SIZE = 12
+477
View File
@@ -0,0 +1,477 @@
import asyncio
import multiprocessing as mp
import os
import subprocess
import traceback
from dataclasses import dataclass
from enum import Enum
from multiprocessing import Process, Queue
from typing import Optional
import numpy as np
import torch
from config import (
MODEL_CONFIG, NUM_FRAMES, FRAME_HEIGHT, FRAME_WIDTH, NUM_INFERENCE_STEPS
)
class CommandType(Enum):
"""Commands sent from main process to GPU worker."""
INIT = "init"
STEP = "step"
RESET = "reset"
SHUTDOWN = "shutdown"
@dataclass
class Command:
"""Command sent to GPU worker subprocess."""
type: CommandType
data: dict = None
def __post_init__(self):
if self.data is None:
self.data = {}
@dataclass
class Response:
"""Response from GPU worker subprocess."""
success: bool
frames: Optional[list[np.ndarray]] = None
error: Optional[str] = None
def gpu_worker_process(
gpu_id: int,
cuda_device: str,
command_queue: Queue,
response_queue: Queue,
):
"""
Worker process that runs on a single GPU.
This function runs in a subprocess with CUDA_VISIBLE_DEVICES set to a single GPU.
"""
# Set CUDA_VISIBLE_DEVICES BEFORE importing torch or any CUDA code
os.environ["CUDA_VISIBLE_DEVICES"] = cuda_device
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "FLASH_ATTN"
# Now import the generator (this will initialize CUDA with the single visible GPU)
from fastvideo.entrypoints.streaming_generator import StreamingVideoGenerator
from fastvideo.models.dits.matrixgame.utils import expand_action_to_frames
from fastvideo.configs.sample.base import SamplingParam
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.utils import align_to, shallow_asdict
generator = None
def initialize_generator():
nonlocal generator
print(f"[GPU {gpu_id}] Loading model...")
generator = StreamingVideoGenerator.from_pretrained(
MODEL_CONFIG["model_path"],
num_gpus=1,
use_fsdp_inference=True,
dit_cpu_offload=True,
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
pin_cpu_memory=True,
)
# Initial reset
actions = {
"keyboard": torch.zeros((NUM_FRAMES, MODEL_CONFIG["keyboard_dim"])),
"mouse": torch.zeros((NUM_FRAMES, 2))
}
generator.reset(
prompt="",
image_path=MODEL_CONFIG["image_url"],
mouse_cond=actions["mouse"].unsqueeze(0),
keyboard_cond=actions["keyboard"].unsqueeze(0),
grid_sizes=torch.tensor([150, 44, 80]),
num_frames=NUM_FRAMES,
height=FRAME_HEIGHT,
width=FRAME_WIDTH,
num_inference_steps=NUM_INFERENCE_STEPS,
)
print(f"[GPU {gpu_id}] Model loaded and ready")
def do_step(keyboard_vector: list, mouse_vector: list) -> list[np.ndarray]:
"""Execute a generation step."""
action = {
"keyboard": torch.tensor(keyboard_vector).cuda(),
"mouse": torch.tensor(mouse_vector).cuda()
}
keyboard_cond, mouse_cond = expand_action_to_frames(action, 12)
frames, _ = generator.step(keyboard_cond, mouse_cond)
return frames if frames is not None else []
def do_reset() -> list[np.ndarray]:
"""Reset the generator and return initial frames."""
# Clear local state
generator.accumulated_frames = []
generator.block_idx = 0
generator.block_dir = None
if generator.writer:
generator.writer.close()
generator.writer = None
# Use queue mode to clear
if generator._use_queue_mode and generator.executor._streaming_enabled:
generator.executor.submit_clear()
clear_result = generator.executor.wait_result()
if clear_result.error:
raise clear_result.error
# Set up sampling parameters
actions = {
"keyboard": torch.zeros((NUM_FRAMES, MODEL_CONFIG["keyboard_dim"])),
"mouse": torch.zeros((NUM_FRAMES, 2))
}
if generator.sampling_param is None:
generator.sampling_param = SamplingParam.from_pretrained(
generator.fastvideo_args.model_path)
generator.sampling_param.update({
"prompt": "",
"image_path": MODEL_CONFIG["image_url"],
"mouse_cond": actions["mouse"].unsqueeze(0),
"keyboard_cond": actions["keyboard"].unsqueeze(0),
"grid_sizes": torch.tensor([150, 44, 80]),
"num_frames": NUM_FRAMES,
"height": FRAME_HEIGHT,
"width": FRAME_WIDTH,
"num_inference_steps": NUM_INFERENCE_STEPS,
})
generator.sampling_param.height = align_to(generator.sampling_param.height, 16)
generator.sampling_param.width = align_to(generator.sampling_param.width, 16)
latents_size = [
(generator.sampling_param.num_frames - 1) // 4 + 1,
generator.sampling_param.height // 8,
generator.sampling_param.width // 8
]
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
generator.sampling_param.return_frames = True
generator.sampling_param.save_video = False
# Create new batch
generator.batch = ForwardBatch(
**shallow_asdict(generator.sampling_param),
eta=0.0,
n_tokens=n_tokens,
VSA_sparsity=generator.fastvideo_args.VSA_sparsity,
)
# Use queue mode to reset
if generator._use_queue_mode:
generator.executor.submit_reset(generator.batch, generator.fastvideo_args)
result = generator.executor.wait_result()
if result.error:
raise result.error
else:
generator.executor.execute_streaming_reset(generator.batch, generator.fastvideo_args)
# Generate initial frame
return do_step([0, 0, 0, 0], [0, 0])
# Main worker loop
print(f"[GPU {gpu_id}] Worker process starting...")
try:
while True:
cmd: Command = command_queue.get()
if cmd.type == CommandType.SHUTDOWN:
print(f"[GPU {gpu_id}] Shutting down...")
if generator:
generator.shutdown()
response_queue.put(Response(success=True))
break
elif cmd.type == CommandType.INIT:
try:
initialize_generator()
response_queue.put(Response(success=True))
except Exception as e:
print(f"[GPU {gpu_id}] Init error: {e}")
traceback.print_exc()
response_queue.put(Response(success=False, error=str(e)))
elif cmd.type == CommandType.STEP:
try:
keyboard = cmd.data.get("keyboard", [0, 0, 0, 0])
mouse = cmd.data.get("mouse", [0, 0])
frames = do_step(keyboard, mouse)
response_queue.put(Response(success=True, frames=frames))
except Exception as e:
print(f"[GPU {gpu_id}] Step error: {e}")
traceback.print_exc()
response_queue.put(Response(success=False, error=str(e)))
elif cmd.type == CommandType.RESET:
try:
frames = do_reset()
response_queue.put(Response(success=True, frames=frames))
except Exception as e:
print(f"[GPU {gpu_id}] Reset error: {e}")
traceback.print_exc()
response_queue.put(Response(success=False, error=str(e)))
except Exception as e:
print(f"[GPU {gpu_id}] Worker crashed: {e}")
traceback.print_exc()
print(f"[GPU {gpu_id}] Worker process exiting")
class GPUSlot:
"""Manages a single GPU worker subprocess."""
def __init__(self, gpu_id: int, cuda_device: str):
self.gpu_id = gpu_id
self.cuda_device = cuda_device
self.client_id: Optional[str] = None
self.process: Optional[Process] = None
self.command_queue: Optional[Queue] = None
self.response_queue: Optional[Queue] = None
self.ready: bool = False
self._lock = asyncio.Lock()
@property
def is_available(self) -> bool:
return self.client_id is None and self.ready and self.process is not None and self.process.is_alive()
async def start(self):
"""Start the GPU worker subprocess."""
ctx = mp.get_context("spawn")
self.command_queue = ctx.Queue()
self.response_queue = ctx.Queue()
self.process = ctx.Process(
target=gpu_worker_process,
args=(self.gpu_id, self.cuda_device, self.command_queue, self.response_queue),
daemon=False, # Must be False so executor can spawn child workers
)
loop = asyncio.get_event_loop()
await loop.run_in_executor(None, self.process.start)
# Send init command and wait for response
response = await self._send_command(Command(CommandType.INIT), timeout=600.0)
if not response.success:
raise RuntimeError(f"GPU {self.gpu_id} failed to initialize: {response.error}")
self.ready = True
async def _send_command(self, cmd: Command, timeout: float = 300.0) -> Response:
"""Send a command to the worker and wait for response."""
loop = asyncio.get_event_loop()
# Send command in thread pool to not block
await loop.run_in_executor(None, self.command_queue.put, cmd)
# Wait for response with timeout
def get_response():
return self.response_queue.get(timeout=timeout)
response = await loop.run_in_executor(None, get_response)
return response
async def step(self, keyboard: list, mouse: list) -> list[np.ndarray]:
"""Execute a generation step."""
async with self._lock:
response = await self._send_command(
Command(CommandType.STEP, {"keyboard": keyboard, "mouse": mouse})
)
if not response.success:
raise RuntimeError(f"Step failed: {response.error}")
return response.frames or []
async def reset(self) -> list[np.ndarray]:
"""Reset the generator."""
async with self._lock:
response = await self._send_command(Command(CommandType.RESET))
if not response.success:
raise RuntimeError(f"Reset failed: {response.error}")
return response.frames or []
async def shutdown(self):
"""Shutdown the worker subprocess."""
if self.process and self.process.is_alive():
try:
await self._send_command(Command(CommandType.SHUTDOWN), timeout=30.0)
except Exception:
pass
self.process.terminate()
self.process.join(timeout=5)
if self.process.is_alive():
self.process.kill()
class GPUPool:
"""Manages multiple GPU worker subprocesses."""
def __init__(self, gpu_ids: list[int]):
self.gpu_ids = gpu_ids
self.slots: dict[int, GPUSlot] = {
gpu_id: GPUSlot(gpu_id, str(i)) for i, gpu_id in enumerate(gpu_ids)
}
self.waiting_list: list[tuple[str, asyncio.Event, "WebSocket"]] = []
self.client_gpu_map: dict[str, int] = {}
self._pool_lock = asyncio.Lock()
async def initialize(self):
"""Start initializing all GPU workers in the background.
Returns immediately so the server can start accepting connections.
GPUs become available as each one finishes loading.
"""
print(f"Initializing GPU pool with {len(self.gpu_ids)} GPUs: {self.gpu_ids}")
for gpu_id in self.gpu_ids:
asyncio.create_task(self._init_gpu(gpu_id))
async def _init_gpu(self, gpu_id: int):
"""Initialize a single GPU and assign any waiting clients."""
try:
await self.slots[gpu_id].start()
except Exception as e:
print(f"GPU {gpu_id} failed to initialize: {e}")
return
print(f"GPU pool: {gpu_id} ready ({sum(1 for s in self.slots.values() if s.ready)}/{len(self.gpu_ids)})")
# Check if anyone is waiting for a GPU
async with self._pool_lock:
slot = self.slots[gpu_id]
if slot.is_available and self.waiting_list:
waiting_client_id, ready_event, _ = self.waiting_list.pop(0)
slot.client_id = waiting_client_id
self.client_gpu_map[waiting_client_id] = gpu_id
print(f"Client {waiting_client_id[:8]} assigned GPU {gpu_id} from queue")
ready_event.set()
# Notify remaining clients of updated positions
await self._send_queue_updates()
async def acquire(self, client_id: str, websocket=None) -> tuple[int, GPUSlot]:
"""
Acquire a GPU slot for a client.
Returns (gpu_id, slot) when available.
"""
async with self._pool_lock:
for gpu_id, slot in self.slots.items():
if slot.is_available:
slot.client_id = client_id
self.client_gpu_map[client_id] = gpu_id
print(f"Client {client_id[:8]} acquired GPU {gpu_id}")
return gpu_id, slot
# No slot available, wait in queue
print(f"Client {client_id[:8]} waiting in queue (all {len(self.gpu_ids)} GPUs busy)")
ready_event = asyncio.Event()
async with self._pool_lock:
self.waiting_list.append((client_id, ready_event, websocket))
await ready_event.wait()
gpu_id = self.client_gpu_map.get(client_id)
if gpu_id is None:
raise RuntimeError(f"Client {client_id} was signaled but has no GPU assigned")
return gpu_id, self.slots[gpu_id]
async def release(self, client_id: str):
"""Release a GPU slot when a client disconnects."""
async with self._pool_lock:
gpu_id = self.client_gpu_map.pop(client_id, None)
if gpu_id is None:
return
slot = self.slots[gpu_id]
print(f"Client {client_id[:8]} released GPU {gpu_id}")
# Reset the generator so it's fresh for the next client
try:
await slot.reset()
print(f"[GPU {gpu_id}] Reset complete, ready for next client")
except Exception as e:
print(f"[GPU {gpu_id}] Reset on release failed: {e}")
slot.client_id = None
# Assign to next waiting client
if self.waiting_list:
waiting_client_id, ready_event, _ = self.waiting_list.pop(0)
slot.client_id = waiting_client_id
self.client_gpu_map[waiting_client_id] = gpu_id
print(f"Client {waiting_client_id[:8]} assigned GPU {gpu_id} from queue")
ready_event.set()
# Notify remaining clients of updated positions
await self._send_queue_updates()
async def _send_queue_updates(self):
"""Send updated queue positions to all waiting clients. Must be called with _pool_lock held."""
for i, (cid, _, ws) in enumerate(self.waiting_list):
if ws is not None:
try:
await ws.send_json({
"type": "queue_status",
"position": i + 1,
"total_gpus": len(self.gpu_ids),
"available_gpus": 0,
})
except Exception:
pass # Client may have disconnected
async def shutdown(self):
"""Shutdown all GPU workers."""
print("Shutting down GPU pool...")
tasks = [slot.shutdown() for slot in self.slots.values()]
await asyncio.gather(*tasks, return_exceptions=True)
print("GPU pool shutdown complete")
def get_status(self) -> dict:
"""Get the current status of the GPU pool."""
return {
"total_gpus": len(self.gpu_ids),
"available_gpus": sum(1 for slot in self.slots.values() if slot.is_available),
"queue_size": len(self.waiting_list),
"gpu_status": {
gpu_id: {
"available": slot.is_available,
"client_id": slot.client_id[:8] if slot.client_id else None,
"process_alive": slot.process.is_alive() if slot.process else False,
}
for gpu_id, slot in self.slots.items()
}
}
def get_available_gpus() -> list[int]:
"""Get list of available GPU IDs from environment or auto-detect."""
cuda_visible = os.environ.get("CUDA_VISIBLE_DEVICES", "")
if cuda_visible:
return [int(x.strip()) for x in cuda_visible.split(",") if x.strip()]
# Auto-detect available GPUs
try:
result = subprocess.run(
["nvidia-smi", "--query-gpu=index", "--format=csv,noheader"],
capture_output=True, text=True
)
if result.returncode == 0:
return [int(x.strip()) for x in result.stdout.strip().split("\n") if x.strip()]
except Exception:
pass
return [0]
+216
View File
@@ -0,0 +1,216 @@
from fastapi import FastAPI, WebSocket, WebSocketDisconnect
from fastapi.middleware.cors import CORSMiddleware
from fastapi.staticfiles import StaticFiles
from contextlib import asynccontextmanager
import asyncio
import base64
import cv2
import os
import time
import uuid
from config import (
KEYBOARD_MAP, CAMERA_MAP, MAX_BLOCKS, JPEG_QUALITY, BATCH_SIZE,
SESSION_TIMEOUT_SECONDS, MODEL_CONFIG
)
from gpu_pool import GPUPool, GPUSlot, get_available_gpus
# Global GPU pool
gpu_pool: GPUPool = None
@asynccontextmanager
async def lifespan(app: FastAPI):
"""Application lifespan manager."""
global gpu_pool
print("Starting server...")
# Get available GPUs
gpu_ids = get_available_gpus()
print(f"Detected GPUs: {gpu_ids}")
# Initialize GPU pool (spawns subprocess per GPU)
gpu_pool = GPUPool(gpu_ids)
await gpu_pool.initialize()
print("Server ready")
yield
print("Shutting down server...")
await gpu_pool.shutdown()
app = FastAPI(lifespan=lifespan)
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
@app.get("/status")
async def get_status():
"""Get the current status of the GPU pool."""
if gpu_pool is None:
return {"error": "GPU pool not initialized"}
return gpu_pool.get_status()
def encode_frames(frames: list) -> list[str]:
"""Encode frames to base64 JPEG."""
encoded = []
for frame in frames:
frame_rgb = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
_, buffer = cv2.imencode('.jpg', frame_rgb, [cv2.IMWRITE_JPEG_QUALITY, JPEG_QUALITY])
encoded.append(base64.b64encode(buffer.tobytes()).decode("utf-8"))
return encoded
async def send_frames(websocket: WebSocket, frames: list):
"""Encode and send frames in batches."""
if not frames:
return
encoded_frames = encode_frames(frames)
for i in range(0, len(encoded_frames), BATCH_SIZE):
batch = encoded_frames[i:i + BATCH_SIZE]
await websocket.send_json({"type": "frame_batch", "frames": batch})
@app.websocket("/ws")
async def websocket_endpoint(websocket: WebSocket):
await websocket.accept()
client_id = str(uuid.uuid4())
print(f"Client {client_id[:8]} connected")
# Send initial queue status
status = gpu_pool.get_status()
await websocket.send_json({
"type": "queue_status",
"position": status["queue_size"] + 1 if status["available_gpus"] == 0 else 0,
"total_gpus": status["total_gpus"],
"available_gpus": status["available_gpus"]
})
gpu_id = None
slot: GPUSlot = None
timeout_task: asyncio.Task = None
async def session_timeout():
"""Close the session after timeout."""
await asyncio.sleep(SESSION_TIMEOUT_SECONDS)
print(f"[GPU {gpu_id}] Session timeout for client {client_id[:8]}")
try:
await websocket.send_json({
"type": "session_timeout",
"message": f"Session expired after {SESSION_TIMEOUT_SECONDS} seconds"
})
await websocket.close(code=1000, reason="Session timeout")
except Exception:
pass # WebSocket may already be closed
try:
# Acquire a GPU slot (may wait in queue)
gpu_id, slot = await gpu_pool.acquire(client_id, websocket)
# Start session timeout
timeout_task = asyncio.create_task(session_timeout())
# Notify client they're connected to a GPU
await websocket.send_json({
"type": "gpu_assigned",
"gpu_id": gpu_id,
"session_timeout": SESSION_TIMEOUT_SECONDS,
"image_url": MODEL_CONFIG.get("image_url")
})
# Generate and send initial frame (GPU is already reset and ready)
block_count = 0
await websocket.send_json({"type": "block_count", "count": block_count, "max": MAX_BLOCKS})
try:
frames = await slot.step([0, 0, 0, 0], [0, 0])
block_count = 1
await websocket.send_json({"type": "block_count", "count": block_count, "max": MAX_BLOCKS})
await send_frames(websocket, frames)
except Exception as e:
print(f"[GPU {gpu_id}] Initial frame error: {e}")
# Main event loop
while True:
data = await websocket.receive_json()
message_type = data.get("type", "key")
if message_type == "reset":
print(f"[GPU {gpu_id}] Reset requested by client {client_id[:8]}")
await websocket.send_json({"type": "reset_started"})
try:
frames = await slot.reset()
block_count = 1
await websocket.send_json({"type": "block_count", "count": block_count, "max": MAX_BLOCKS})
await send_frames(websocket, frames)
await websocket.send_json({"type": "reset_complete"})
print(f"[GPU {gpu_id}] Reset complete")
except Exception as e:
print(f"[GPU {gpu_id}] Reset error: {e}")
await websocket.send_json({"type": "reset_complete"})
continue
# Handle key press
key = data.get("key")
if key and block_count < MAX_BLOCKS:
if key in CAMERA_MAP:
keyboard_vector = [0, 0, 0, 0]
mouse_vector = CAMERA_MAP[key]
else:
keyboard_vector = KEYBOARD_MAP.get(key, [0, 0, 0, 0])
mouse_vector = [0, 0]
print(f"[GPU {gpu_id}] Key '{key}' pressed, generating block {block_count + 1}...")
t_start = time.time()
frames = await slot.step(keyboard_vector, mouse_vector)
t_generation = time.time() - t_start
block_count += 1
await websocket.send_json({"type": "block_count", "count": block_count, "max": MAX_BLOCKS})
if frames:
t_encode_start = time.time()
await send_frames(websocket, frames)
t_total = time.time() - t_encode_start
print(f"[GPU {gpu_id}] Gen: {t_generation:.3f}s, Encode+Send: {t_total:.3f}s")
except WebSocketDisconnect:
print(f"Client {client_id[:8]} disconnected")
except Exception as e:
print(f"Client {client_id[:8]} error: {e}")
import traceback
traceback.print_exc()
finally:
# Cancel the timeout task if still running
if timeout_task and not timeout_task.done():
timeout_task.cancel()
try:
await timeout_task
except asyncio.CancelledError:
pass
if client_id:
await gpu_pool.release(client_id)
# Serve built frontend (must be after API/WebSocket routes)
static_dir = os.path.join(os.path.dirname(__file__), "..", "client", "dist")
if os.path.isdir(static_dir):
app.mount("/", StaticFiles(directory=static_dir, html=True), name="static")
if __name__ == "__main__":
import uvicorn
uvicorn.run(app, host="0.0.0.0", port=8000)