Compare commits
4
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
0a0cdebaba | ||
|
|
cd9e7a8cac | ||
|
|
87e19dad5a | ||
|
|
79195ad18e |
@@ -40,6 +40,9 @@ dist/
|
||||
*.egg
|
||||
eggs/
|
||||
.eggs/
|
||||
node_modules/
|
||||
.vite
|
||||
vite.config.js.timestamp-*.mjs
|
||||
|
||||
# MkDocs documentation
|
||||
site/
|
||||
|
||||
@@ -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>
|
||||
Generated
+1901
File diff suppressed because it is too large
Load Diff
@@ -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"
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -0,0 +1,8 @@
|
||||
import App from './App.svelte'
|
||||
import './style.css'
|
||||
|
||||
const app = new App({
|
||||
target: document.getElementById('app')
|
||||
})
|
||||
|
||||
export default app
|
||||
@@ -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;
|
||||
}
|
||||
@@ -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',
|
||||
},
|
||||
},
|
||||
}
|
||||
})
|
||||
@@ -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
|
||||
@@ -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]
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user