Compare commits
14
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a6d4d815f8 | ||
|
|
42da276006 | ||
|
|
7b3c4fbdd1 | ||
|
|
f041bd3f49 | ||
|
|
fbb9a73ab4 | ||
|
|
daea023383 | ||
|
|
0422d18377 | ||
|
|
f46c6d923f | ||
|
|
79a31b0b45 | ||
|
|
ef98769b46 | ||
|
|
869c5f7370 | ||
|
|
74529f22a4 | ||
|
|
2a2d67e792 | ||
|
|
bb3367c9ac |
@@ -1,6 +1,6 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
# from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples"
|
||||
def main():
|
||||
@@ -8,30 +8,35 @@ def main():
|
||||
# model.
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
model_name = "Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
model_name,
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
dit_cpu_offload=False,
|
||||
dit_cpu_offload=True,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
|
||||
pin_cpu_memory=False,
|
||||
pin_cpu_memory=True,
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
# sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
|
||||
# sampling_param.num_frames = 45
|
||||
# sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
|
||||
sampling_param = SamplingParam.from_pretrained(model_name)
|
||||
sampling_param.image_path = "test.jpg"
|
||||
# sampling_param.num_inference_steps = 0
|
||||
# Generate videos with the same simple API, regardless of GPU count
|
||||
prompt = (
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True)
|
||||
i2v_prompt = "An astronaut hatching from an egg, on the surface of the moon, the darkness and depth of space realised in the background. High quality, ultrarealistic detail and breath-taking movie-like camera shot."
|
||||
i2v_prompt = "A little girl is packing a suitcase and the contents starts flying out of the suitcase everywhere."
|
||||
prompt = i2v_prompt
|
||||
# prompt = (
|
||||
# "A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
# "wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
# "natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
# )
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, sampling_param=sampling_param)
|
||||
# video = generator.generate_video(prompt, sampling_param=sampling_param, output_path="wan_t2v_videos/")
|
||||
return
|
||||
|
||||
# Generate another video with a different prompt, without reloading the
|
||||
# model!
|
||||
|
||||
@@ -31,6 +31,7 @@ def main():
|
||||
prompt = (
|
||||
"A neon-lit alley in futuristic Tokyo during a heavy rainstorm at night. The puddles reflect glowing signs in kanji, advertising ramen, karaoke, and VR arcades. A woman in a translucent raincoat walks briskly with an LED umbrella. Steam rises from a street food cart, and a cat darts across the screen. Raindrops are visible on the camera lens, creating a cinematic bokeh effect."
|
||||
)
|
||||
prompt = "A vintage train snakes through the mountains, its plume of white steam rising dramatically against the jagged peaks. The cars glint in the late afternoon sun, their deep crimson and gold accents lending a touch of elegance. The tracks carve a precarious path along the cliffside, revealing glimpses of a roaring river far below. Inside, passengers peer out the large windows, their faces lit with awe as the landscape unfolds."
|
||||
start_time = time.perf_counter()
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, sampling_param=sampling_param)
|
||||
end_time = time.perf_counter()
|
||||
@@ -45,7 +46,7 @@ def main():
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
start_time = time.perf_counter()
|
||||
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True)
|
||||
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=False)
|
||||
end_time = time.perf_counter()
|
||||
gen_time2 = end_time - start_time
|
||||
|
||||
|
||||
@@ -0,0 +1,38 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples"
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=2,
|
||||
use_fsdp_inference=True,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=False,
|
||||
image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-I2V-14B-480P-Diffusers")
|
||||
sampling_param.num_frames = 61
|
||||
sampling_param.num_inference_steps = 40
|
||||
sampling_param.guidance_scale = 5.0
|
||||
sampling_param.height = 448
|
||||
sampling_param.width = 832
|
||||
sampling_param.seed = 1024
|
||||
sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
|
||||
# Generate videos with the same simple API, regardless of GPU count
|
||||
prompt = (
|
||||
"An astronaut hatching from an egg, on the surface of the moon, the darkness and depth of space realised in the background. High quality, ultrarealistic detail and breath-taking movie-like camera shot."
|
||||
)
|
||||
video = generator.generate_video(prompt, sampling_param=sampling_param, output_path=OUTPUT_PATH, save_video=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,182 @@
|
||||
# FastVideo Dual Model Setup (T2V + I2V)
|
||||
|
||||
This document describes the dual model functionality that has been added to the FastVideo Gradio app, supporting both Text-to-Video (T2V) and Image-to-Video (I2V) generation modes with specialized models.
|
||||
|
||||
## Overview
|
||||
|
||||
The Gradio app now supports both Text-to-Video (T2V) and Image-to-Video (I2V) generation modes using specialized models:
|
||||
- **T2V Model**: `FastVideo/FastWan2.1-T2V-1.3B-Diffusers` for text-to-video generation
|
||||
- **I2V Model**: `Wan-AI/Wan2.1-I2V-14B-480P-Diffusers` for image-to-video generation
|
||||
|
||||
Users can switch between modes using the tabbed interface and upload images for I2V generation.
|
||||
|
||||
## Model Configuration
|
||||
|
||||
### T2V Model
|
||||
- **Model**: `FastVideo/FastWan2.1-T2V-1.3B-Diffusers`
|
||||
- **Purpose**: Text-to-video generation
|
||||
- **Default Parameters**: Optimized for text prompts
|
||||
|
||||
### I2V Model
|
||||
- **Model**: `Wan-AI/Wan2.1-I2V-14B-480P-Diffusers`
|
||||
- **Purpose**: Image-to-video generation
|
||||
- **Default Parameters**: Optimized for image animation
|
||||
|
||||
## Changes Made
|
||||
|
||||
### Backend Changes (`ray_serve_backend.py`)
|
||||
|
||||
1. **Added dual model support**:
|
||||
- Separate model paths for T2V and I2V
|
||||
- Automatic model selection based on request type
|
||||
- Independent model initialization
|
||||
|
||||
2. **Updated `VideoGenerationRequest`**:
|
||||
- Added `model_type` field ("t2v" or "i2v")
|
||||
- Added `image_path` field for I2V input
|
||||
|
||||
3. **Enhanced `FastVideoAPI` class**:
|
||||
- Dual model initialization (`t2v_generator` and `i2v_generator`)
|
||||
- Separate default parameters for each model
|
||||
- Automatic model selection in `generate_video` method
|
||||
|
||||
### Frontend Changes (`gradio_frontend.py`)
|
||||
|
||||
1. **Tabbed interface**:
|
||||
- "Text-to-Video" tab for T2V generation
|
||||
- "Image-to-Video" tab for I2V generation
|
||||
|
||||
2. **Automatic model selection**:
|
||||
- T2V tab uses T2V model automatically
|
||||
- I2V tab uses I2V model automatically
|
||||
- Model type sent in API requests
|
||||
|
||||
3. **Separate event handlers**:
|
||||
- `handle_t2v_generation` for text-to-video
|
||||
- `handle_i2v_generation` for image-to-video
|
||||
|
||||
## Usage
|
||||
|
||||
### Starting the Application
|
||||
|
||||
1. **Using the combined startup script (recommended)**:
|
||||
```bash
|
||||
python start_ray_serve_app.py
|
||||
```
|
||||
|
||||
2. **Manual startup**:
|
||||
```bash
|
||||
# Start backend
|
||||
python ray_serve_backend.py \
|
||||
--t2v_model_path "FastVideo/FastWan2.1-T2V-1.3B-Diffusers" \
|
||||
--i2v_model_path "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
|
||||
|
||||
# Start frontend
|
||||
python gradio_frontend.py --backend_url "http://localhost:8000"
|
||||
```
|
||||
|
||||
### Using T2V Mode
|
||||
|
||||
1. Navigate to the "Text-to-Video" tab
|
||||
2. Enter a text prompt describing the video you want to generate
|
||||
3. Adjust advanced parameters if needed
|
||||
4. Click "Run" to generate the video
|
||||
|
||||
### Using I2V Mode
|
||||
|
||||
1. Navigate to the "Image-to-Video" tab
|
||||
2. Upload an image using the image upload component
|
||||
3. Enter a prompt describing how the image should animate
|
||||
4. Adjust advanced parameters if needed
|
||||
5. Click "Run" to generate the video
|
||||
|
||||
### Example Prompts
|
||||
|
||||
**T2V Examples**:
|
||||
- "A hand enters the frame, pulling a sheet of plastic wrap over three balls of dough placed on a wooden surface."
|
||||
- "A vintage train snakes through the mountains, its plume of white steam rising dramatically against the jagged peaks."
|
||||
|
||||
**I2V Examples**:
|
||||
- "The image comes to life with subtle movement, the scene gently animating while maintaining the original composition and mood."
|
||||
- "The static image transforms into a dynamic scene with natural motion, preserving the original lighting and atmosphere."
|
||||
|
||||
## Testing
|
||||
|
||||
A comprehensive test script is provided to verify both T2V and I2V functionality:
|
||||
|
||||
```bash
|
||||
python test_i2v.py
|
||||
```
|
||||
|
||||
This script:
|
||||
- Tests backend health
|
||||
- Tests T2V functionality with text prompts
|
||||
- Tests I2V functionality with image uploads
|
||||
- Verifies response formats for both modes
|
||||
- Cleans up test files
|
||||
|
||||
## Technical Details
|
||||
|
||||
### Backend API Changes
|
||||
|
||||
The `/generate_video` endpoint now accepts:
|
||||
|
||||
```json
|
||||
{
|
||||
"prompt": "Animation description",
|
||||
"model_type": "t2v", // or "i2v"
|
||||
"image_path": "/path/to/input/image.png", // for I2V
|
||||
// ... other parameters
|
||||
}
|
||||
```
|
||||
|
||||
### Model Selection Logic
|
||||
|
||||
- **T2V Mode**: Uses `FastVideo/FastWan2.1-T2V-1.3B-Diffusers`
|
||||
- **I2V Mode**: Uses `Wan-AI/Wan2.1-I2V-14B-480P-Diffusers`
|
||||
- **Automatic Selection**: Based on presence of `image_path` and `model_type`
|
||||
|
||||
### Memory Management
|
||||
|
||||
- Both models are loaded independently
|
||||
- Automatic cleanup prevents memory leaks
|
||||
- Temporary files are cleaned up after processing
|
||||
|
||||
## Configuration Options
|
||||
|
||||
### Command Line Arguments
|
||||
|
||||
**Backend**:
|
||||
- `--t2v_model_path`: Path to T2V model
|
||||
- `--i2v_model_path`: Path to I2V model
|
||||
- `--output_path`: Output directory
|
||||
- `--host`, `--port`: Server configuration
|
||||
|
||||
**Frontend**:
|
||||
- `--backend_url`: Backend API URL
|
||||
- `--t2v_model_path`, `--i2v_model_path`: Model paths (for reference)
|
||||
- `--host`, `--port`: Server configuration
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
1. **Model Loading Issues**: Ensure both models are accessible
|
||||
2. **Memory Issues**: The backend includes automatic cleanup
|
||||
3. **Image Upload Failures**: Check image format and size
|
||||
4. **Generation Failures**: Check backend logs for detailed errors
|
||||
|
||||
## Performance Considerations
|
||||
|
||||
- **Model Loading**: Both models are loaded at startup
|
||||
- **Memory Usage**: Higher memory requirements due to dual models
|
||||
- **Generation Time**: I2V may take longer due to larger model size
|
||||
- **GPU Requirements**: Ensure sufficient VRAM for both models
|
||||
|
||||
## Future Enhancements
|
||||
|
||||
Potential improvements:
|
||||
- Model switching without restart
|
||||
- Batch processing for both modes
|
||||
- Advanced image preprocessing
|
||||
- Progress indicators
|
||||
- Result caching and history
|
||||
- Model-specific parameter optimization
|
||||
@@ -0,0 +1,222 @@
|
||||
# FastVideo with Ray Serve Backend and Gradio Frontend
|
||||
|
||||
This setup provides a scalable web application for FastVideo inference using Ray Serve as the backend and Gradio as the frontend.
|
||||
|
||||
## Architecture
|
||||
|
||||
- **Backend**: Ray Serve handles video generation requests with GPU acceleration
|
||||
- **Frontend**: Gradio provides a user-friendly web interface
|
||||
- **Communication**: HTTP REST API between frontend and backend
|
||||
|
||||
## Features
|
||||
|
||||
- ✅ Scalable backend with Ray Serve
|
||||
- ✅ GPU-accelerated video generation
|
||||
- ✅ User-friendly Gradio interface
|
||||
- ✅ Health monitoring and error handling
|
||||
- ✅ All original functionality preserved
|
||||
- ✅ Easy deployment and management
|
||||
|
||||
## Installation
|
||||
|
||||
1. Install the additional dependencies:
|
||||
|
||||
```bash
|
||||
pip install -r requirements_ray_serve.txt
|
||||
```
|
||||
|
||||
2. Ensure you have the FastVideo model available (the default is `FastVideo/FastHunyuan-diffusers`)
|
||||
|
||||
## Usage
|
||||
|
||||
### Option 1: Start Both Services Together (Recommended)
|
||||
|
||||
Use the startup script to launch both backend and frontend:
|
||||
|
||||
```bash
|
||||
python start_ray_serve_app.py
|
||||
```
|
||||
|
||||
This will:
|
||||
- Start the Ray Serve backend on port 8000
|
||||
- Start the Gradio frontend on port 7860
|
||||
- Monitor both services and provide unified logging
|
||||
- Handle graceful shutdown with Ctrl+C
|
||||
|
||||
### Option 2: Start Services Separately
|
||||
|
||||
#### Start Backend Only
|
||||
|
||||
```bash
|
||||
python ray_serve_backend.py --model_path FastVideo/FastHunyuan-diffusers --output_path outputs
|
||||
```
|
||||
|
||||
#### Start Frontend Only
|
||||
|
||||
```bash
|
||||
python gradio_frontend.py --backend_url http://localhost:8000 --model_path FastVideo/FastHunyuan-diffusers
|
||||
```
|
||||
|
||||
## Configuration
|
||||
|
||||
### Command Line Arguments
|
||||
|
||||
#### Startup Script (`start_ray_serve_app.py`)
|
||||
|
||||
- `--model_path`: Path to the FastVideo model (default: `FastVideo/FastHunyuan-diffusers`)
|
||||
- `--output_path`: Directory to save generated videos (default: `outputs`)
|
||||
- `--backend_host`: Backend host to bind to (default: `0.0.0.0`)
|
||||
- `--backend_port`: Backend port (default: `8000`)
|
||||
- `--frontend_host`: Frontend host to bind to (default: `0.0.0.0`)
|
||||
- `--frontend_port`: Frontend port (default: `7860`)
|
||||
- `--skip_backend_check`: Skip backend health check
|
||||
|
||||
#### Backend (`ray_serve_backend.py`)
|
||||
|
||||
- `--model_path`: Path to the FastVideo model
|
||||
- `--output_path`: Directory to save generated videos
|
||||
- `--host`: Host to bind to
|
||||
- `--port`: Port to bind to
|
||||
|
||||
#### Frontend (`gradio_frontend.py`)
|
||||
|
||||
- `--backend_url`: URL of the Ray Serve backend
|
||||
- `--model_path`: Path to the model (for default parameters)
|
||||
- `--host`: Host to bind to
|
||||
- `--port`: Port to bind to
|
||||
|
||||
### Environment Variables
|
||||
|
||||
You can also set these environment variables:
|
||||
|
||||
- `FASTVIDEO_MODEL_PATH`: Path to the FastVideo model
|
||||
- `FASTVIDEO_OUTPUT_PATH`: Directory to save generated videos
|
||||
- `RAY_SERVE_HOST`: Backend host
|
||||
- `RAY_SERVE_PORT`: Backend port
|
||||
- `GRADIO_HOST`: Frontend host
|
||||
- `GRADIO_PORT`: Frontend port
|
||||
|
||||
## API Endpoints
|
||||
|
||||
### Backend API (Ray Serve)
|
||||
|
||||
- `GET /health`: Health check endpoint
|
||||
- `POST /generate_video`: Video generation endpoint
|
||||
|
||||
#### Video Generation Request
|
||||
|
||||
```json
|
||||
{
|
||||
"prompt": "A beautiful sunset over the ocean",
|
||||
"negative_prompt": "blurry, low quality",
|
||||
"use_negative_prompt": true,
|
||||
"seed": 42,
|
||||
"guidance_scale": 7.5,
|
||||
"num_frames": 21,
|
||||
"height": 512,
|
||||
"width": 512,
|
||||
"num_inference_steps": 20,
|
||||
"randomize_seed": false
|
||||
}
|
||||
```
|
||||
|
||||
#### Video Generation Response
|
||||
|
||||
```json
|
||||
{
|
||||
"output_path": "/path/to/generated/video.mp4",
|
||||
"seed": 42,
|
||||
"success": true,
|
||||
"error_message": null
|
||||
}
|
||||
```
|
||||
|
||||
## Deployment
|
||||
|
||||
### Local Development
|
||||
|
||||
1. Start the application:
|
||||
```bash
|
||||
python start_ray_serve_app.py
|
||||
```
|
||||
|
||||
2. Access the frontend at: `http://localhost:7860`
|
||||
3. Access the backend API at: `http://localhost:8000`
|
||||
|
||||
### Production Deployment
|
||||
|
||||
For production deployment, consider:
|
||||
|
||||
1. **Load Balancing**: Use a reverse proxy (nginx, traefik) in front of the services
|
||||
2. **Monitoring**: Add monitoring and logging (Prometheus, Grafana)
|
||||
3. **Scaling**: Configure Ray Serve for horizontal scaling
|
||||
4. **Security**: Add authentication and rate limiting
|
||||
5. **Storage**: Use shared storage for video outputs
|
||||
|
||||
### Docker Deployment
|
||||
|
||||
Create a Dockerfile for containerized deployment:
|
||||
|
||||
```dockerfile
|
||||
FROM python:3.9-slim
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
# Install dependencies
|
||||
COPY requirements_ray_serve.txt .
|
||||
RUN pip install -r requirements_ray_serve.txt
|
||||
|
||||
# Copy application files
|
||||
COPY . .
|
||||
|
||||
# Expose ports
|
||||
EXPOSE 8000 7860
|
||||
|
||||
# Start the application
|
||||
CMD ["python", "start_ray_serve_app.py"]
|
||||
```
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
### Common Issues
|
||||
|
||||
1. **Backend not starting**: Check GPU availability and Ray installation
|
||||
2. **Frontend can't connect**: Verify backend URL and network connectivity
|
||||
3. **Video generation fails**: Check model path and GPU memory
|
||||
4. **Port conflicts**: Change ports using command line arguments
|
||||
|
||||
### Logs
|
||||
|
||||
- Backend logs are prefixed with `[BACKEND]`
|
||||
- Frontend logs are prefixed with `[FRONTEND]`
|
||||
- Use `--skip_backend_check` if you need to debug startup issues
|
||||
|
||||
### Performance Tuning
|
||||
|
||||
- Adjust `num_replicas` in the Ray Serve deployment for scaling
|
||||
- Configure `max_concurrent_queries` based on GPU memory
|
||||
- Use multiple GPUs by modifying `ray_actor_options`
|
||||
|
||||
## Migration from Original Gradio Demo
|
||||
|
||||
The new setup maintains full compatibility with the original functionality:
|
||||
|
||||
1. All parameters and options are preserved
|
||||
2. The same example prompts are included
|
||||
3. The UI layout and behavior are identical
|
||||
4. Video generation quality is the same
|
||||
|
||||
The main differences are:
|
||||
- Backend processing is now handled by Ray Serve
|
||||
- Better error handling and status monitoring
|
||||
- Scalable architecture for production use
|
||||
- Separation of concerns between frontend and backend
|
||||
|
||||
## Contributing
|
||||
|
||||
To extend this setup:
|
||||
|
||||
1. Add new endpoints to `ray_serve_backend.py`
|
||||
2. Update the frontend in `gradio_frontend.py`
|
||||
3. Modify the startup script if needed
|
||||
4. Update this README with new features
|
||||
@@ -0,0 +1,501 @@
|
||||
import argparse
|
||||
import os
|
||||
import time
|
||||
import json
|
||||
import statistics
|
||||
import asyncio
|
||||
import aiohttp
|
||||
from copy import deepcopy
|
||||
from typing import List, Dict, Any
|
||||
import threading
|
||||
|
||||
import torch
|
||||
|
||||
# All the prompts for stress testing
|
||||
STRESS_TEST_PROMPTS = [
|
||||
"A person reading a book with words that float off the pages and form pictures.",
|
||||
"A person diving into a pool of liquid crystal, creating ripples of light.",
|
||||
"A handheld shot chasing after a group of friends laughing and playing on the beach at sunset.",
|
||||
"A mysterious ancient temple hidden in the jungle.",
|
||||
"A high-speed train navigating a steep descent.",
|
||||
"a toy robot wearing blue jeans and a white t shirt taking a pleasant stroll in Antarctica during a winter storm",
|
||||
"A cheetah accelerating to full speed while chasing its prey.",
|
||||
"A serene orchard is in full bloom, with trees heavy with blossoms and bees buzzing around, darting from flower to flower in a display of natural harmony.",
|
||||
"A little child let out a big yawn",
|
||||
"Subtle reflections of a woman on the window of a train moving at hyper-speed in a Japanese city.",
|
||||
"A truck left along the edge of a cliff, revealing the stunning coastal landscape below with waves crashing against the rocks.",
|
||||
"A red bird transforms into a flag",
|
||||
"A zoom-out from a single leaf on a tree to reveal the entire forest, showcasing the vastness and diversity of the woodland.",
|
||||
"A slow-motion video of a liquid droplet bouncing on a water-repellent surface.",
|
||||
"Static camera shot. A dinasour running near some lions and chasing them away.",
|
||||
"an adorable kangaroo wearing purple overalls and cowboy boots taking a pleasant stroll in Mumbai India during a beautiful sunset",
|
||||
"A zoom-in on an artist's brush touching the canvas, highlighting the texture of the paint and the strokes being made.",
|
||||
"an old man wearing blue jeans and a white t shirt taking a pleasant stroll in Mumbai India during a colorful festival",
|
||||
"A woman is ascending to the sky from the ground",
|
||||
"View out a window of a giant strange creature walking in rundown city at night, one single street lamp dimly lighting the area.",
|
||||
"An arc shot around a lone tree in a vast, foggy field at dawn, revealing the changing light and shadows.",
|
||||
"A person sculpting a statue out of a waterfall, the water solidifying under their touch.",
|
||||
"The person's forehead creased with concentration as she worked on a challenging puzzle.",
|
||||
"The person's cheeks flushed with pleasure as she savored a delicious meal.",
|
||||
"Hand-drawn simple line art, a young kid looking up into space with a wondrous expression on his face.",
|
||||
"A crab made of different jewlery is walking on the beach. As it walks, it drops different jewelry pieces like diamonds, pearls, etc",
|
||||
"Gold coins are falling out when elevator door opens",
|
||||
"the scene transitions from huge waves into a snowy mountain at sunset",
|
||||
"a giant cathedral is completely filled with cats. there are cats everywhere you look. a man enters the cathedral and bows before the giant cat king sitting on a throne.",
|
||||
"A mother dog gently picks up a piece of meat and carefully places it in her puppy's bowl, her eyes filled with warmth and care as she watches her little one eat.",
|
||||
"A soap bubble floating in the air, displaying iridescent colors that shift and change as it moves through different angles of light.",
|
||||
"A truck left alongside a train moving through the countryside, matching its speed and revealing the changing landscape.",
|
||||
"An astronaut walking between stone buildings.",
|
||||
"A close-up shot of the person's face reveals his fear and desperation as he navigates the ship through the storm.",
|
||||
"A frozen lake slowly cracking and thawing as spring arrives, with sheets of ice breaking apart and drifting across the surface.",
|
||||
"A FPV shot zooming through a tunnel into a vibrant underwater space.",
|
||||
"a toy robot wearing blue jeans and a white t shirt taking a pleasant stroll in Mumbai India during a colorful festival",
|
||||
"A person sips on a smoothie, the cool and fruity flavors refreshing her mouth.",
|
||||
"In a vibrant theater, a magician in dazzling attire stands center stage, pulling a comically oversized rubber chicken from an ornate, old-fashioned box. His costume shimmers under the stage lights, adding to the spectacle. The crowd erupts in laughter and applause, their faces filled with joy and amazement. The magician's expression hints at mischievous delight as he holds up the rubber chicken, his performance bringing cheer to the audience.",
|
||||
"A hamster running on a spinning wheel.",
|
||||
"A quaint village nestled in a valley is surrounded by blooming cherry blossoms, with petals drifting through the air as villagers go about their daily activities, adding life to the scene.",
|
||||
"In a tranquil forest clearing, a sparkling waterfall cascades down into a clear pool, surrounded by lush greenery and flowers, with occasional birds fluttering by.",
|
||||
"A woman beamed with pride as she watched her child perform on stage.",
|
||||
"an adorable kangaroo wearing blue jeans and a white t shirt taking a pleasant stroll in Mumbai India during a winter storm",
|
||||
"A man is eating salad",
|
||||
"An Asian girl wearing a bright yellow T-shirt and white pants is Hip-Hop dancing",
|
||||
"nighttime footage of a hermit crab using an incandescent lightbulb as its shell",
|
||||
"a toy robot wearing a green dress and a sun hat taking a pleasant stroll in Antarctica during a beautiful sunset",
|
||||
"A goat operating a food truck, serving gourmet grilled cheese sandwiches to a line of animals.",
|
||||
"Macro shot. Man in an antique scuba helmet with dark glass walking out of a flower",
|
||||
"A bustling train station in the heart of a vibrant city.",
|
||||
"Light filtering through a canopy of autumn leaves, casting warm, dappled patterns of yellow, orange, and red onto the ground.",
|
||||
"Chimneys in the setting sun",
|
||||
"A longboarder accelerating downhill, carving through turns.",
|
||||
"A couple runs through a sudden downpour, laughing and splashing in puddles as they try to find shelter.",
|
||||
"A glass of iced coffee condensing water on the outside, with droplets forming and sliding down the glass in slow motion.",
|
||||
"macro shot of a leaf showing tiny trains moving through its veins",
|
||||
"A corgi wearing sunglasses walks on the beach of a tropical island",
|
||||
"Borneo wildlife on the Kinabatangan River",
|
||||
"A beautiful silhouette animation shows a wolf howling at the moon, feeling lonely, until it finds its pack.",
|
||||
"an adorable kangaroo wearing blue jeans and a white t shirt taking a pleasant stroll in Johannesburg South Africa during a colorful festival",
|
||||
"A green monster made of plants walks through an airport.",
|
||||
"A close up view of a glass sphere that has a zen garden within it. There is a small dwarf in the sphere who is raking the zen garden and creating patterns in the sand.",
|
||||
"A person on a scooter colliding with a park bench, the scooter tipping over.",
|
||||
"A tilt-up from a city street, ascending to show the skyline with its mix of modern and historic architecture.",
|
||||
"A chef tossing a pancake into the air and catching it.",
|
||||
"A woman whispering a secret into a friend's ear.",
|
||||
"A vulture circling high in the sky.",
|
||||
"A medieval castle overlooking a bustling renaissance fair.",
|
||||
"a toy robot wearing purple overalls and cowboy boots taking a pleasant stroll in Mumbai India during a beautiful sunset",
|
||||
"A man standing in front of a burning building giving the 'thumbs up' sign.",
|
||||
"The person's cheeks flushed with embarrassment as he told a funny story.",
|
||||
"Llamas and Emus are playing chess",
|
||||
"A woman sipping a steaming cup of tea.",
|
||||
"A tree root bursting through the seat of an ancient, weathered bench, intertwining with the wood.",
|
||||
"Smoke rises from the chimney of a cozy log cabin nestled in the woods, with soft light glowing from the windows, suggesting a warm and inviting atmosphere.",
|
||||
"A close-up of sparkling water being poured into a glass, capturing the detailed flow and bubbles.",
|
||||
"a woman wearing blue jeans and a white t shirt taking a pleasant stroll in Antarctica during a beautiful sunset",
|
||||
"The Glenfinnan Viaduct is a historic railway bridge in Scotland, UK, that crosses over the west highland line between the towns of Mallaig and Fort William. It is a stunning sight as a steam train leaves the bridge, traveling over the arch-covered viaduct. The landscape is dotted with lush greenery and rocky mountains, creating a picturesque backdrop for the train journey. The sky is blue and the sun is shining, making for a beautiful day to explore this majestic spot.",
|
||||
"A piece of elastic fabric being pulled and stretched, then returning to its original size when the tension is released.",
|
||||
"a woman wearing a green dress and a sun hat taking a pleasant stroll in Antarctica during a beautiful sunset",
|
||||
"A video of a water jet cutting through metal, showing the powerful and precise movement of water.",
|
||||
"Car mirrors and sunsets",
|
||||
"Giant Pandas are eating hot noodles in a Chinese restaurant",
|
||||
"A rally car taking a fast turn on a track",
|
||||
"a toy robot wearing purple overalls and cowboy boots taking a pleasant stroll in Mumbai India during a colorful festival",
|
||||
"A crystal-clear icicle slowly dripping as it melts in the warmth of the midday sun, each drop sparkling as it falls.",
|
||||
"A tilt-down from a chandelier in a grand hall, revealing the ornate decor and people mingling below.",
|
||||
"A man is playing the drums under the water",
|
||||
"A person playing an electric guitar made of lightning, with thunderous sound waves.",
|
||||
"A person floating in a bubble, drifting over a bustling cityscape.",
|
||||
"A tilt-down from a starry night sky, revealing a quiet forest clearing bathed in moonlight.",
|
||||
"A pan right through a dense jungle, moving past lush vegetation and exotic wildlife.",
|
||||
"Close-up of a man eating an apple.",
|
||||
"A low-angle shot of a dancer leaping gracefully into the air, making their movement appear even more dynamic and powerful.",
|
||||
"A woman is search her bag trying to find something.",
|
||||
"A bulldozer clears debris from a demolished building, making way for new construction.",
|
||||
"A man sighed in relief as the doctor delivered the good news.",
|
||||
"A tsunami coming through an alley in Bulgaria, dynamic movement.",
|
||||
"Blooming Flowers",
|
||||
"A push-in through a dense crowd at a festival, moving towards a performer on stage who is captivating the audience.",
|
||||
"A truck right through a tranquil garden, moving past blooming flowers, trees, and a small fountain.",
|
||||
"The person's eyes sparkled with excitement as he greeted a friend.",
|
||||
"A person playing chess with a robot on a floating platform above the ocean.",
|
||||
"A gentle breeze rustles the leaves as someone walks down a serene forest path, sunlight filtering through the trees and shifting patterns on the ground as branches sway.",
|
||||
"A rollercoaster ride from a city to a desert and then to an ice world",
|
||||
"A pan left across an ancient library, moving from shelf to shelf, showcasing rows of leather-bound books.",
|
||||
"A mother otter floating on her back in a river, cradling her pup on her stomach to keep it safe and warm in the gentle current.",
|
||||
"an adorable kangaroo wearing purple overalls and cowboy boots taking a pleasant stroll in Johannesburg South Africa during a colorful festival",
|
||||
"a woman wearing a green dress and a sun hat taking a pleasant stroll in Mumbai India during a colorful festival",
|
||||
"A delicate layer of morning frost melting off a flower petal, the tiny droplets glistening like diamonds in the light.",
|
||||
"A panda is cooking for her child, her child is next to her.",
|
||||
"Macro shot of a man wearing an antique diving helmet with dark glass and a jetpack walking on the veins of a leaf. Realistic style",
|
||||
"an old man wearing purple overalls and cowboy boots taking a pleasant stroll in Johannesburg South Africa during a beautiful sunset",
|
||||
"A girl is unfolding a birthday gift.",
|
||||
"A pencil drawing an architectural plan.",
|
||||
"A handheld camera following a dog running through a park, bouncing and tilting as it captures the dog's joyful exploration.",
|
||||
"A pan left across a serene beach at sunrise, moving from the darkened shore to the brightening horizon.",
|
||||
"A group of people are clapping to celebrate",
|
||||
"Vendors set up stalls at a bustling farmer's market, displaying fresh fruits and vegetables, while people stroll through, selecting produce and enjoying the lively atmosphere.",
|
||||
"A police helicopter hovers above a high-speed chase, guiding officers on the ground to apprehend a suspect.",
|
||||
"A paper origami dragon riding a boat in waves. Realistic style.",
|
||||
"A close-up of a droplet of dew forming on a leaf, capturing the detailed surface tension.",
|
||||
"a toy robot wearing blue jeans and a white t shirt taking a pleasant stroll in Mumbai India during a beautiful sunset",
|
||||
"A dry rainbow rose is coming back to life.",
|
||||
"A glass falling off a table and shattering on the floor.",
|
||||
"A marathon runner crossing the finish line after a grueling race.",
|
||||
"A zoom-in on a drop of morning dew on a leaf, showing the reflection of the surrounding world within it.",
|
||||
"A child blowing on hot cocoa to cool it down.",
|
||||
"A squad of futsal players showcasing their skills on an indoor court.",
|
||||
"A princess is brushing her long golden hair in the garden.",
|
||||
"A close-up of a pair of eyes, revealing the subtle emotions and reflections within them.",
|
||||
"A tracking shot of a group of cyclists racing through a forest trail, with trees and foliage rushing by.",
|
||||
"A woman yawning widely at the end of a long day.",
|
||||
"an old man wearing a green dress and a sun hat taking a pleasant stroll in Johannesburg South Africa during a colorful festival",
|
||||
"Hidden within a garden, an ancient fountain trickles with water, surrounded by vibrant flowers and lush greenery that seem to whisper secrets of the past.",
|
||||
"A Chinese man sits at a table and eats noodles with chopsticks",
|
||||
"A pink pig running fast toward the camera in an alley in Tokyo.",
|
||||
"Strange creatures move through a mysterious, foggy marsh, their silhouettes barely visible through the dense mist as they navigate the eerie, otherworldly landscape.",
|
||||
"Tour of an art gallery with many beautiful works of art in different styles.",
|
||||
"FPV flying through a colorful coral lined streets of an underwater suburban neighborhood.",
|
||||
"Aerial view of Santorini during the blue hour, showcasing the stunning architecture of white Cycladic buildings with blue domes. The caldera views are breathtaking, and the lighting creates a beautiful, serene atmosphere.",
|
||||
"Camera zoom out. A couple walking along the beach as the sun sets over the ocean.",
|
||||
"an extreme close up shot of a woman's eye, with her iris appearing as earth",
|
||||
"a woman wearing purple overalls and cowboy boots taking a pleasant stroll in Mumbai India during a colorful festival",
|
||||
"an old man wearing a green dress and a sun hat taking a pleasant stroll in Mumbai India during a winter storm",
|
||||
"an adorable kangaroo wearing blue jeans and a white t shirt taking a pleasant stroll in Antarctica during a winter storm",
|
||||
"A martial artist breaking a board with a powerful punch.",
|
||||
"People gather on a peaceful beach at sunset, a bonfire crackling as they sit around, enjoying the warmth and the sight of the sun dipping below the horizon.",
|
||||
"A close-up of a waterfall, showing the detailed movement of water as it crashes down.",
|
||||
"A child is blowing bubbles",
|
||||
"a woman wearing a green dress and a sun hat taking a pleasant stroll in Johannesburg South Africa during a winter storm",
|
||||
"A wide-angle perspective of a serene lake surrounded by mountains, reflecting the sky and creating a sense of infinite space.",
|
||||
"The person's eyebrows arched in skepticism as she listened to a dubious claim.",
|
||||
"an old man wearing blue jeans and a white t shirt taking a pleasant stroll in Mumbai India during a beautiful sunset",
|
||||
"a woman wearing a green dress and a sun hat taking a pleasant stroll in Johannesburg South Africa during a colorful festival",
|
||||
"a woman wearing purple overalls and cowboy boots taking a pleasant stroll in Antarctica during a colorful festival",
|
||||
"a toy robot wearing blue jeans and a white t shirt taking a pleasant stroll in Antarctica during a colorful festival",
|
||||
"A chef flips a pancake and puts cream on it.",
|
||||
"An astronaut runs on the surface of the moon, the low angle shot shows the vast background of the moon, the movement is smooth and appears lightweight",
|
||||
"A man's face lit up with happiness as he received a heartfelt compliment.",
|
||||
"A futuristic spaceport hums with activity as ships of various shapes and sizes take off and land on multiple platforms, their engines glowing with vibrant colors.",
|
||||
"A person knitting a scarf using beams of light instead of yarn.",
|
||||
"A pedestal up from the edge of a canyon, gradually revealing the expansive landscape and river below.",
|
||||
"a woman wearing purple overalls and cowboy boots taking a pleasant stroll in Johannesburg South Africa during a colorful festival",
|
||||
"an old man wearing blue jeans and a white t shirt taking a pleasant stroll in Johannesburg South Africa during a colorful festival",
|
||||
"A person walking up a staircase made of clouds leading to a floating castle.",
|
||||
"Monks meditate in a serene mountaintop temple, sitting in quiet reflection as the wind gently moves through the surrounding trees, creating a sense of peace and tranquility.",
|
||||
"An aerial shot of a bustling city intersection at rush hour, capturing the organized chaos of cars and pedestrians.",
|
||||
"A pair of hands skillfully knitting a colorful scarf, the yarn winding through their fingers with each stitch.",
|
||||
"Close-up, a Chinese child is eating dumplings",
|
||||
"A kite losing wind and falling to the ground.",
|
||||
"Bioluminescent waves gently wash ashore on a deserted beach, illuminating the sand with each cresting wave as a figure walks along the water's edge, leaving glowing footprints.",
|
||||
"A red panda taking a bite of a pizza",
|
||||
"A close-up shot of a young woman driving a car, looking thoughtful, blurred green forest visible through the rainy car window.",
|
||||
"A high-speed video of a splash created by a stone thrown into a pond.",
|
||||
"A metal rod being bent slightly by a force and then springing back to its original straight shape when the force is removed.",
|
||||
"A hedgehog in a knight's armor, riding a toy horse into a medieval castle.",
|
||||
"A bird made of fresh oranges rushes out of the orange",
|
||||
"A low altitude first person perspective camera tracking shot of a soccer player's feet dribbling the ball on the groud in a soccer field, Sports Videography, Motion Tracking camera shot",
|
||||
"A tranquil island retreat features swaying palm trees and hammocks strung between them, inviting guests to relax and enjoy the serene beauty of the surroundings.",
|
||||
"a spooky haunted mansion, with friendly jack o lanterns and ghost characters welcoming trick or treaters to the entrance, tilt shift photography",
|
||||
"A coconut tree made of dollar bills at sunset, with bills falling off like leaves.",
|
||||
"A motocross bike accelerating out of a tight turn on a dirt track.",
|
||||
"A tranquil Zen garden with a gently flowing stream and koi fish.",
|
||||
"A green monster made of leaves walks through the airport, carrying a suitcase.",
|
||||
"A time-lapse of a frost-covered leaf gradually thawing in the morning sunlight, with tiny water droplets forming and trickling down.",
|
||||
"A woman practicing her archery skills at a range.",
|
||||
"A slow-motion video of ink being injected into a tank of water, creating intricate and beautiful patterns.",
|
||||
"a woman wearing blue jeans and a white t shirt taking a pleasant stroll in Johannesburg South Africa during a winter storm",
|
||||
"The person's forehead creased with worry as he listened to bad news.",
|
||||
"An arc shot around a grand piano being played in an empty concert hall, the motion revealing the intricate details of the instrument.",
|
||||
"A person conducting a symphony of animals in a forest clearing.",
|
||||
"A truck right alongside a flowing river, capturing the movement of the water and the surrounding forest.",
|
||||
"A rocket blasting off from the launch pad, accelerating rapidly into the sky.",
|
||||
"Workers move through a picturesque vineyard during the harvest season, carefully picking grapes and placing them into baskets as the sun bathes the vines in a warm glow.",
|
||||
"A person is eating an ice cream.",
|
||||
"An over-the-shoulder perspective of a chef meticulously plating a dish in a bustling kitchen.",
|
||||
"A man looked away in shame when confronted with his wrongdoing.",
|
||||
"A person is savoring a slice of pizza at a pizzeria."
|
||||
]
|
||||
|
||||
class BackendStressTest:
|
||||
def __init__(self, output_path: str,
|
||||
server_url: str = "http://localhost:8000", max_concurrent: int = 50):
|
||||
self.output_path = output_path
|
||||
self.server_url = server_url
|
||||
self.max_concurrent = max_concurrent
|
||||
|
||||
# Results storage
|
||||
self.results = []
|
||||
self.lock = threading.Lock()
|
||||
|
||||
async def check_health(self) -> bool:
|
||||
"""Check if the Ray Serve backend is healthy"""
|
||||
try:
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.get(f"{self.server_url}/health", timeout=aiohttp.ClientTimeout(total=5)) as response:
|
||||
return response.status == 200
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
def build_request_params(self, prompt: str, **kwargs) -> Dict[str, Any]:
|
||||
"""Build request parameters for Ray Serve backend"""
|
||||
# Default parameters matching the Ray Serve backend
|
||||
default_params = {
|
||||
'prompt': prompt,
|
||||
'negative_prompt': None,
|
||||
'use_negative_prompt': False,
|
||||
'seed': 42,
|
||||
'guidance_scale': 7.5,
|
||||
'num_frames': 21,
|
||||
'height': 448,
|
||||
'width': 832,
|
||||
'num_inference_steps': 20,
|
||||
'randomize_seed': True,
|
||||
'return_frames': False # Don't return frames for stress testing to reduce overhead
|
||||
}
|
||||
|
||||
# Override with any provided kwargs
|
||||
for key, value in kwargs.items():
|
||||
if key in default_params:
|
||||
default_params[key] = value
|
||||
|
||||
# Randomize seed if requested
|
||||
if default_params.get('randomize_seed', True):
|
||||
default_params['seed'] = torch.randint(0, 1000000, (1,)).item()
|
||||
|
||||
# Handle negative prompt
|
||||
if not default_params.get('use_negative_prompt', False):
|
||||
default_params['negative_prompt'] = None
|
||||
|
||||
# NEW: Remove keys with None values to avoid sending nulls that may break validation
|
||||
clean_params = {k: v for k, v in default_params.items() if v is not None}
|
||||
return clean_params
|
||||
|
||||
async def test_single_request(self, session: aiohttp.ClientSession, prompt: str, request_id: int) -> Dict[str, Any]:
|
||||
"""Test a single request and measure latency"""
|
||||
start_time = time.time()
|
||||
|
||||
try:
|
||||
# Build request parameters
|
||||
request_params = self.build_request_params(prompt)
|
||||
|
||||
# Make request to Ray Serve backend
|
||||
async with session.post(
|
||||
f"{self.server_url}/generate_video",
|
||||
json=request_params,
|
||||
timeout=aiohttp.ClientTimeout(total=900) # 15 minute timeout for video generation
|
||||
) as response:
|
||||
|
||||
end_time = time.time()
|
||||
latency = end_time - start_time
|
||||
|
||||
if response.status == 200:
|
||||
response_data = await response.json()
|
||||
if response_data.get('success', False):
|
||||
result = {
|
||||
'request_id': request_id,
|
||||
'prompt': prompt,
|
||||
'latency': latency,
|
||||
'status': 'success',
|
||||
'response_time': latency, # Use our own timing
|
||||
'timestamp': start_time,
|
||||
'output_path': response_data.get('output_path', ''),
|
||||
'used_seed': response_data.get('seed', request_params['seed'])
|
||||
}
|
||||
else:
|
||||
result = {
|
||||
'request_id': request_id,
|
||||
'prompt': prompt,
|
||||
'latency': latency,
|
||||
'status': 'error',
|
||||
'error': response_data.get('error_message', 'Unknown backend error'),
|
||||
'timestamp': start_time
|
||||
}
|
||||
else:
|
||||
response_text = await response.text()
|
||||
result = {
|
||||
'request_id': request_id,
|
||||
'prompt': prompt,
|
||||
'latency': latency,
|
||||
'status': 'error',
|
||||
'error': f"HTTP {response.status}: {response_text}",
|
||||
'timestamp': start_time
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
end_time = time.time()
|
||||
latency = end_time - start_time
|
||||
result = {
|
||||
'request_id': request_id,
|
||||
'prompt': prompt,
|
||||
'latency': latency,
|
||||
'status': 'error',
|
||||
'error': str(e),
|
||||
'timestamp': start_time
|
||||
}
|
||||
|
||||
# Thread-safe result storage
|
||||
with self.lock:
|
||||
self.results.append(result)
|
||||
|
||||
return result
|
||||
|
||||
async def run_stress_test(self, num_iterations: int = 1, concurrent_requests: int = None):
|
||||
"""Run the stress test with multiple iterations and concurrent requests"""
|
||||
if concurrent_requests is None:
|
||||
concurrent_requests = self.max_concurrent
|
||||
|
||||
# Check backend health before starting
|
||||
print(f"Testing Ray Serve backend at {self.server_url}...")
|
||||
if not await self.check_health():
|
||||
print(f"❌ Backend is not healthy at {self.server_url}")
|
||||
print("Make sure the Ray Serve backend is running with:")
|
||||
print("python ray_serve_backend.py")
|
||||
return
|
||||
print("✅ Backend is healthy and ready for stress testing")
|
||||
|
||||
print(f"\nStarting stress test with {len(STRESS_TEST_PROMPTS)} prompts")
|
||||
print(f"Running {num_iterations} iteration(s) with {concurrent_requests} concurrent requests")
|
||||
print(f"Total requests: {len(STRESS_TEST_PROMPTS) * num_iterations}")
|
||||
print(f"Backend URL: {self.server_url}")
|
||||
print("-" * 80)
|
||||
|
||||
all_prompts = STRESS_TEST_PROMPTS * num_iterations
|
||||
request_id = 0
|
||||
|
||||
# Create semaphore to limit concurrent requests
|
||||
semaphore = asyncio.Semaphore(concurrent_requests)
|
||||
|
||||
async def limited_request(session: aiohttp.ClientSession, prompt: str, req_id: int):
|
||||
async with semaphore:
|
||||
return await self.test_single_request(session, prompt, req_id)
|
||||
|
||||
# Run concurrent requests using asyncio
|
||||
async with aiohttp.ClientSession() as session:
|
||||
# Create all tasks
|
||||
tasks = [
|
||||
limited_request(session, prompt, request_id + i)
|
||||
for i, prompt in enumerate(all_prompts)
|
||||
]
|
||||
|
||||
# Process completed requests as they finish
|
||||
completed = 0
|
||||
for coro in asyncio.as_completed(tasks):
|
||||
try:
|
||||
result = await coro
|
||||
completed += 1
|
||||
prompt = result['prompt']
|
||||
status_icon = "✅" if result['status'] == 'success' else "❌"
|
||||
output_info = f" -> {result.get('output_path', 'N/A')}" if result['status'] == 'success' else ""
|
||||
print(f"{status_icon} [{completed}/{len(all_prompts)}] {result['latency']:.2f}s - {prompt[:50]}...{output_info}")
|
||||
except Exception as e:
|
||||
completed += 1
|
||||
print(f"❌ [{completed}/{len(all_prompts)}] Exception: {e}")
|
||||
|
||||
self.analyze_results()
|
||||
|
||||
def analyze_results(self):
|
||||
"""Analyze and print test results"""
|
||||
print("\n" + "=" * 80)
|
||||
print("STRESS TEST RESULTS")
|
||||
print("=" * 80)
|
||||
|
||||
successful_requests = [r for r in self.results if r['status'] == 'success']
|
||||
failed_requests = [r for r in self.results if r['status'] == 'error']
|
||||
|
||||
print(f"Total Requests: {len(self.results)}")
|
||||
print(f"Successful: {len(successful_requests)}")
|
||||
print(f"Failed: {len(failed_requests)}")
|
||||
print(f"Success Rate: {len(successful_requests)/len(self.results)*100:.1f}%")
|
||||
|
||||
if successful_requests:
|
||||
latencies = [r['latency'] for r in successful_requests]
|
||||
print(f"\nLatency Statistics (seconds):")
|
||||
print(f" Min: {min(latencies):.2f}")
|
||||
print(f" Max: {max(latencies):.2f}")
|
||||
print(f" Mean: {statistics.mean(latencies):.2f}")
|
||||
print(f" Median: {statistics.median(latencies):.2f}")
|
||||
print(f" Std Dev: {statistics.stdev(latencies):.2f}")
|
||||
|
||||
# Percentiles
|
||||
sorted_latencies = sorted(latencies)
|
||||
p50 = sorted_latencies[int(len(sorted_latencies) * 0.5)]
|
||||
p90 = sorted_latencies[int(len(sorted_latencies) * 0.9)]
|
||||
p95 = sorted_latencies[int(len(sorted_latencies) * 0.95)]
|
||||
p99 = sorted_latencies[int(len(sorted_latencies) * 0.99)]
|
||||
|
||||
print(f" P50: {p50:.2f}")
|
||||
print(f" P90: {p90:.2f}")
|
||||
print(f" P95: {p95:.2f}")
|
||||
print(f" P99: {p99:.2f}")
|
||||
|
||||
if failed_requests:
|
||||
print(f"\nFailed Requests ({len(failed_requests)}):")
|
||||
for req in failed_requests[:5]: # Show first 5 failures
|
||||
print(f" - {req['error']}")
|
||||
if len(failed_requests) > 5:
|
||||
print(f" ... and {len(failed_requests) - 5} more")
|
||||
|
||||
# Save detailed results
|
||||
results_file = os.path.join(self.output_path, "stress_test_results.json")
|
||||
os.makedirs(self.output_path, exist_ok=True)
|
||||
|
||||
with open(results_file, 'w') as f:
|
||||
json.dump({
|
||||
'summary': {
|
||||
'total_requests': len(self.results),
|
||||
'successful_requests': len(successful_requests),
|
||||
'failed_requests': len(failed_requests),
|
||||
'success_rate': len(successful_requests)/len(self.results)*100 if self.results else 0
|
||||
},
|
||||
'latency_stats': {
|
||||
'min': min(latencies) if successful_requests else 0,
|
||||
'max': max(latencies) if successful_requests else 0,
|
||||
'mean': statistics.mean(latencies) if successful_requests else 0,
|
||||
'median': statistics.median(latencies) if successful_requests else 0,
|
||||
'std_dev': statistics.stdev(latencies) if len(successful_requests) > 1 else 0
|
||||
},
|
||||
'detailed_results': self.results
|
||||
}, f, indent=2)
|
||||
|
||||
print(f"\nDetailed results saved to: {results_file}")
|
||||
|
||||
async def main():
|
||||
parser = argparse.ArgumentParser(description="FastVideo Ray Serve Backend Stress Test")
|
||||
parser.add_argument("--output_path",
|
||||
type=str,
|
||||
default="outputs",
|
||||
help="Path to save test results")
|
||||
parser.add_argument("--server_url",
|
||||
type=str,
|
||||
default="http://localhost:8000",
|
||||
help="Ray Serve backend URL")
|
||||
parser.add_argument("--max_concurrent",
|
||||
type=int,
|
||||
default=50,
|
||||
help="Maximum concurrent requests")
|
||||
parser.add_argument("--iterations",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Number of iterations through all prompts")
|
||||
parser.add_argument("--concurrent_requests",
|
||||
type=int,
|
||||
default=None,
|
||||
help="Number of concurrent requests (overrides max_concurrent)")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
# Create stress test instance
|
||||
stress_test = BackendStressTest(
|
||||
output_path=args.output_path,
|
||||
server_url=args.server_url,
|
||||
max_concurrent=args.max_concurrent
|
||||
)
|
||||
|
||||
# Run the stress test
|
||||
await stress_test.run_stress_test(
|
||||
num_iterations=args.iterations,
|
||||
concurrent_requests=args.concurrent_requests
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
asyncio.run(main())
|
||||
@@ -0,0 +1,537 @@
|
||||
import argparse
|
||||
import os
|
||||
import requests
|
||||
import json
|
||||
import base64
|
||||
import io
|
||||
from typing import Optional
|
||||
|
||||
import gradio as gr
|
||||
import torch
|
||||
import imageio
|
||||
from PIL import Image
|
||||
import numpy as np
|
||||
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
|
||||
class RayServeClient:
|
||||
def __init__(self, backend_url: str):
|
||||
self.backend_url = backend_url.rstrip('/')
|
||||
self.generate_endpoint = f"{self.backend_url}/generate_video"
|
||||
self.health_endpoint = f"{self.backend_url}/health"
|
||||
# Default request timeout in seconds. Increase if generation may run longer.
|
||||
self.request_timeout_s = int(os.getenv("FASTVIDEO_GENERATION_TIMEOUT", "900")) # 15 minutes default
|
||||
|
||||
def check_health(self) -> bool:
|
||||
"""Check if the backend is healthy"""
|
||||
try:
|
||||
response = requests.get(self.health_endpoint, timeout=5)
|
||||
return response.status_code == 200
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
def generate_video(self, request_data: dict) -> dict:
|
||||
"""Send video generation request to the backend"""
|
||||
try:
|
||||
response = requests.post(
|
||||
self.generate_endpoint,
|
||||
json=request_data,
|
||||
timeout=self.request_timeout_s,
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
except requests.exceptions.RequestException as e:
|
||||
return {
|
||||
"success": False,
|
||||
"error_message": f"Backend request failed: {str(e)}",
|
||||
"output_path": "",
|
||||
"seed": request_data.get("seed", 42)
|
||||
}
|
||||
|
||||
|
||||
def decode_and_save_video_from_frames(frames_b64: list, output_dir: str, prompt: str, fps: int = 24) -> str:
|
||||
"""Decode base64 frames and save them as a video file"""
|
||||
if not frames_b64:
|
||||
return "No frames to save"
|
||||
|
||||
# Create safe filename from prompt
|
||||
safe_prompt = prompt[:50].replace(' ', '_').replace('/', '_').replace('\\', '_')
|
||||
video_filename = f"{safe_prompt}_frames.mp4"
|
||||
video_path = os.path.join(output_dir, video_filename)
|
||||
|
||||
try:
|
||||
# Decode frames from base64
|
||||
decoded_frames = []
|
||||
|
||||
for i, frame_b64 in enumerate(frames_b64):
|
||||
try:
|
||||
# Remove the data URL prefix if present
|
||||
if frame_b64.startswith('data:image/'):
|
||||
frame_b64 = frame_b64.split(',')[1]
|
||||
|
||||
# Decode base64 to bytes
|
||||
frame_bytes = base64.b64decode(frame_b64)
|
||||
|
||||
# Create PIL Image from bytes
|
||||
image = Image.open(io.BytesIO(frame_bytes))
|
||||
|
||||
# Convert PIL Image to numpy array (same format as video_generator.py)
|
||||
frame_array = np.array(image)
|
||||
decoded_frames.append(frame_array)
|
||||
|
||||
except Exception as e:
|
||||
print(f"Warning: Failed to decode frame {i}: {e}")
|
||||
continue
|
||||
|
||||
if not decoded_frames:
|
||||
return "Failed to decode any frames", ""
|
||||
|
||||
# Save as video using imageio (same as video_generator.py)
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
imageio.mimsave(video_path, decoded_frames, fps=fps, format="mp4")
|
||||
|
||||
return f"Saved {len(decoded_frames)} frames as video: {video_path}", video_path
|
||||
|
||||
except Exception as e:
|
||||
return f"Failed to save video: {str(e)}", ""
|
||||
|
||||
|
||||
def create_gradio_interface(backend_url: str, default_params: SamplingParam):
|
||||
"""Create the Gradio interface"""
|
||||
|
||||
# Initialize the Ray Serve client
|
||||
client = RayServeClient(backend_url)
|
||||
|
||||
def generate_video(
|
||||
prompt,
|
||||
negative_prompt,
|
||||
use_negative_prompt,
|
||||
seed,
|
||||
guidance_scale,
|
||||
num_frames,
|
||||
height,
|
||||
width,
|
||||
num_inference_steps,
|
||||
randomize_seed=False,
|
||||
input_image=None,
|
||||
):
|
||||
# Check backend health first
|
||||
if not client.check_health():
|
||||
return None, f"Backend is not available. Please check if Ray Serve is running at {backend_url}", ""
|
||||
|
||||
# Handle input image for I2V
|
||||
image_path = None
|
||||
if input_image is not None:
|
||||
try:
|
||||
# Save the uploaded image to a temporary file
|
||||
import tempfile
|
||||
temp_dir = "temp_images"
|
||||
os.makedirs(temp_dir, exist_ok=True)
|
||||
|
||||
# Generate a unique filename with appropriate extension
|
||||
import uuid
|
||||
# Determine the best format to preserve quality
|
||||
if hasattr(input_image, 'format') and input_image.format:
|
||||
# Use original format if available
|
||||
ext = input_image.format.lower()
|
||||
if ext == 'jpeg':
|
||||
ext = 'jpg'
|
||||
else:
|
||||
# Default to PNG for lossless quality
|
||||
ext = 'png'
|
||||
|
||||
image_filename = f"input_image_{uuid.uuid4().hex[:8]}.{ext}"
|
||||
image_path = os.path.abspath(os.path.join(temp_dir, image_filename))
|
||||
|
||||
# Save the image preserving original quality
|
||||
if ext == 'png':
|
||||
# Use PNG for lossless compression
|
||||
input_image.save(image_path, "PNG", optimize=False)
|
||||
elif ext == 'jpg':
|
||||
# Use high quality JPEG with minimal compression
|
||||
input_image.convert("RGB").save(image_path, "JPEG", quality=95, optimize=False)
|
||||
else:
|
||||
# For other formats, save as PNG to preserve quality
|
||||
input_image.save(image_path, "PNG", optimize=False)
|
||||
|
||||
print(f"Saved input image to: {image_path}")
|
||||
except Exception as e:
|
||||
print(f"Warning: Failed to save input image: {e}")
|
||||
image_path = None
|
||||
|
||||
# Prepare request data - always request frames for video creation
|
||||
request_data = {
|
||||
"prompt": prompt,
|
||||
"negative_prompt": negative_prompt,
|
||||
"use_negative_prompt": use_negative_prompt,
|
||||
"seed": seed,
|
||||
"guidance_scale": guidance_scale,
|
||||
"num_frames": num_frames,
|
||||
"height": height,
|
||||
"width": width,
|
||||
"num_inference_steps": num_inference_steps,
|
||||
"randomize_seed": randomize_seed,
|
||||
"return_frames": False, # Always request frames
|
||||
"image_path": image_path,
|
||||
"model_type": "i2v" if image_path else "t2v" # Use I2V model if image is provided, T2V otherwise
|
||||
}
|
||||
|
||||
# Send request to backend
|
||||
response = client.generate_video(request_data)
|
||||
|
||||
# Clean up temporary image file after processing
|
||||
if image_path and os.path.exists(image_path):
|
||||
try:
|
||||
os.remove(image_path)
|
||||
print(f"Cleaned up temporary image: {image_path}")
|
||||
except Exception as e:
|
||||
print(f"Warning: Failed to clean up temporary image {image_path}: {e}")
|
||||
|
||||
if response.get("success", False):
|
||||
output_path = response.get("output_path", "")
|
||||
used_seed = response.get("seed", seed)
|
||||
frames_b64 = response.get("frames", [])
|
||||
|
||||
print(f"Used seed: {used_seed}")
|
||||
print(f"Output path: {output_path}")
|
||||
|
||||
# Handle frame extraction and video creation
|
||||
frames_status = ""
|
||||
# if frames_b64:
|
||||
# try:
|
||||
# Get the output directory from the video path
|
||||
# output_dir = os.path.dirname(output_path) if output_path else "outputs"
|
||||
# frames_status, video_path = decode_and_save_video_from_frames(frames_b64, output_dir, prompt)
|
||||
# print(f"Frames: {frames_status}")
|
||||
# except Exception as e:
|
||||
# frames_status = f"Failed to save frames video: {str(e)}"
|
||||
# print(f"Frame extraction error: {e}")
|
||||
# else:
|
||||
# frames_status = "No frames returned from backend"
|
||||
|
||||
# Check if the video file exists
|
||||
if os.path.exists(output_path):
|
||||
return output_path, used_seed, frames_status
|
||||
else:
|
||||
return None, f"Video generated but file not found at {output_path} {frames_status}", frames_status
|
||||
else:
|
||||
error_msg = response.get("error_message", "Unknown error occurred")
|
||||
return None, f"Generation failed: {error_msg}", ""
|
||||
|
||||
# Example prompts
|
||||
examples = [
|
||||
"A hand enters the frame, pulling a sheet of plastic wrap over three balls of dough placed on a wooden surface. The plastic wrap is stretched to cover the dough more securely. The hand adjusts the wrap, ensuring that it is tight and smooth over the dough. The scene focuses on the hand's movements as it secures the edges of the plastic wrap. No new objects appear, and the camera remains stationary, focusing on the action of covering the dough.",
|
||||
"A vintage train snakes through the mountains, its plume of white steam rising dramatically against the jagged peaks. The cars glint in the late afternoon sun, their deep crimson and gold accents lending a touch of elegance. The tracks carve a precarious path along the cliffside, revealing glimpses of a roaring river far below. Inside, passengers peer out the large windows, their faces lit with awe as the landscape unfolds.",
|
||||
"A crowded rooftop bar buzzes with energy, the city skyline twinkling like a field of stars in the background. Strings of fairy lights hang above, casting a warm, golden glow over the scene. Groups of people gather around high tables, their laughter blending with the soft rhythm of live jazz. The aroma of freshly mixed cocktails and charred appetizers wafts through the air, mingling with the cool night breeze.",
|
||||
]
|
||||
|
||||
# Example I2V prompts (for when users upload images)
|
||||
i2v_examples = [
|
||||
"The image comes to life with subtle movement, the scene gently animating while maintaining the original composition and mood.",
|
||||
"The static image transforms into a dynamic scene with natural motion, preserving the original lighting and atmosphere.",
|
||||
"The photograph animates with realistic movement, bringing the frozen moment to life while keeping the original artistic style.",
|
||||
]
|
||||
|
||||
# Create Gradio interface
|
||||
with gr.Blocks() as demo:
|
||||
gr.Markdown("# FastVideo Inference Demo (Ray Serve Backend)")
|
||||
gr.Markdown(f"**Backend URL:** {backend_url}")
|
||||
|
||||
# Backend status indicator
|
||||
status_text = gr.Text(
|
||||
label="Backend Status",
|
||||
value="Checking backend status...",
|
||||
interactive=False
|
||||
)
|
||||
|
||||
def update_status():
|
||||
if client.check_health():
|
||||
return "✅ Backend is healthy and ready"
|
||||
else:
|
||||
return "❌ Backend is not available"
|
||||
|
||||
with gr.Tabs():
|
||||
# Text-to-Video Tab
|
||||
with gr.Tab("Text-to-Video"):
|
||||
with gr.Group():
|
||||
with gr.Row():
|
||||
prompt = gr.Text(
|
||||
label="Prompt",
|
||||
show_label=False,
|
||||
max_lines=1,
|
||||
placeholder="Enter your prompt",
|
||||
container=False,
|
||||
)
|
||||
run_button = gr.Button("Run", scale=0)
|
||||
|
||||
result = gr.Video(label="Result", show_label=False)
|
||||
error_output = gr.Text(label="Error", visible=False)
|
||||
frames_output = gr.Text(label="Frame Video Status", visible=False)
|
||||
download_file = gr.File(visible=False)
|
||||
|
||||
gr.Examples(examples=examples, inputs=prompt)
|
||||
|
||||
# Image-to-Video Tab
|
||||
with gr.Tab("Image-to-Video"):
|
||||
with gr.Group():
|
||||
with gr.Row():
|
||||
i2v_prompt = gr.Text(
|
||||
label="Prompt",
|
||||
show_label=False,
|
||||
max_lines=1,
|
||||
placeholder="Describe how the image should animate",
|
||||
container=False,
|
||||
)
|
||||
i2v_run_button = gr.Button("Run", scale=0)
|
||||
|
||||
input_image = gr.Image(
|
||||
label="Input Image",
|
||||
type="pil",
|
||||
show_label=True,
|
||||
container=True,
|
||||
)
|
||||
|
||||
i2v_result = gr.Video(label="Result", show_label=False)
|
||||
i2v_error_output = gr.Text(label="Error", visible=False)
|
||||
i2v_frames_output = gr.Text(label="Frame Video Status", visible=False)
|
||||
i2v_download_file = gr.File(visible=False)
|
||||
|
||||
gr.Examples(examples=i2v_examples, inputs=i2v_prompt)
|
||||
|
||||
# Shared advanced options
|
||||
with gr.Accordion("Advanced options", open=False):
|
||||
with gr.Group():
|
||||
with gr.Row():
|
||||
height = gr.Slider(
|
||||
label="Height",
|
||||
minimum=256,
|
||||
maximum=1024,
|
||||
step=32,
|
||||
value=448,
|
||||
)
|
||||
width = gr.Slider(
|
||||
label="Width",
|
||||
minimum=256,
|
||||
maximum=1024,
|
||||
step=32,
|
||||
value=832
|
||||
)
|
||||
|
||||
with gr.Row():
|
||||
num_frames = gr.Slider(
|
||||
label="Number of Frames",
|
||||
minimum=16,
|
||||
maximum=160,
|
||||
step=16,
|
||||
value=61,
|
||||
)
|
||||
guidance_scale = gr.Slider(
|
||||
label="Guidance Scale",
|
||||
minimum=1,
|
||||
maximum=12,
|
||||
value=3.0,
|
||||
)
|
||||
num_inference_steps = gr.Slider(
|
||||
label="Inference Steps",
|
||||
minimum=3,
|
||||
maximum=100,
|
||||
value=3,
|
||||
)
|
||||
|
||||
with gr.Row():
|
||||
use_negative_prompt = gr.Checkbox(
|
||||
label="Use negative prompt", value=False)
|
||||
negative_prompt = gr.Text(
|
||||
label="Negative prompt",
|
||||
max_lines=1,
|
||||
placeholder="Enter a negative prompt",
|
||||
visible=False,
|
||||
)
|
||||
|
||||
seed = gr.Slider(
|
||||
label="Seed",
|
||||
minimum=0,
|
||||
maximum=1000000,
|
||||
step=1,
|
||||
value=1024
|
||||
)
|
||||
randomize_seed = gr.Checkbox(label="Randomize seed", value=False)
|
||||
seed_output = gr.Number(label="Used Seed")
|
||||
|
||||
gr.Examples(examples=examples, inputs=prompt)
|
||||
|
||||
# Event handlers
|
||||
use_negative_prompt.change(
|
||||
fn=lambda x: gr.update(visible=x),
|
||||
inputs=use_negative_prompt,
|
||||
outputs=negative_prompt,
|
||||
)
|
||||
|
||||
def handle_t2v_generation(*args):
|
||||
# For T2V, we pass None as input_image
|
||||
args = list(args)
|
||||
args.append(None) # Add None for input_image
|
||||
result_path, seed_or_error, frames_status = generate_video(*args)
|
||||
|
||||
if result_path and os.path.exists(result_path):
|
||||
# Show frame status if available
|
||||
if frames_status:
|
||||
return (
|
||||
result_path,
|
||||
seed_or_error,
|
||||
gr.update(visible=False), # error_output
|
||||
gr.update(visible=True, value=frames_status), # frames_output
|
||||
gr.update(visible=True, value=result_path) # download_file
|
||||
)
|
||||
else:
|
||||
return (
|
||||
result_path,
|
||||
seed_or_error,
|
||||
gr.update(visible=False), # error_output
|
||||
gr.update(visible=False), # frames_output
|
||||
gr.update(visible=True, value=result_path) # download_file
|
||||
)
|
||||
else:
|
||||
return (
|
||||
None,
|
||||
seed_or_error,
|
||||
gr.update(visible=True, value=seed_or_error), # error_output
|
||||
gr.update(visible=False), # frames_output
|
||||
gr.update(visible=False) # download_file
|
||||
)
|
||||
|
||||
def handle_i2v_generation(*args):
|
||||
# For I2V, we need to reorder args to match generate_video signature
|
||||
# args should be: [i2v_prompt, negative_prompt, use_negative_prompt, seed, guidance_scale, num_frames, height, width, num_inference_steps, randomize_seed, input_image]
|
||||
result_path, seed_or_error, frames_status = generate_video(*args)
|
||||
|
||||
if result_path and os.path.exists(result_path):
|
||||
# Show frame status if available
|
||||
if frames_status:
|
||||
return (
|
||||
result_path,
|
||||
seed_or_error,
|
||||
gr.update(visible=False), # i2v_error_output
|
||||
gr.update(visible=True, value=frames_status), # i2v_frames_output
|
||||
gr.update(visible=True, value=result_path) # i2v_download_file
|
||||
)
|
||||
else:
|
||||
return (
|
||||
result_path,
|
||||
seed_or_error,
|
||||
gr.update(visible=False), # i2v_error_output
|
||||
gr.update(visible=False), # i2v_frames_output
|
||||
gr.update(visible=True, value=result_path) # i2v_download_file
|
||||
)
|
||||
else:
|
||||
return (
|
||||
None,
|
||||
seed_or_error,
|
||||
gr.update(visible=True, value=seed_or_error), # i2v_error_output
|
||||
gr.update(visible=False), # i2v_frames_output
|
||||
gr.update(visible=False) # i2v_download_file
|
||||
)
|
||||
|
||||
# T2V event handler
|
||||
run_button.click(
|
||||
fn=handle_t2v_generation,
|
||||
inputs=[
|
||||
prompt,
|
||||
negative_prompt,
|
||||
use_negative_prompt,
|
||||
seed,
|
||||
guidance_scale,
|
||||
num_frames,
|
||||
height,
|
||||
width,
|
||||
num_inference_steps,
|
||||
randomize_seed,
|
||||
],
|
||||
outputs=[result, seed_output, error_output, frames_output, download_file],
|
||||
concurrency_limit=20,
|
||||
)
|
||||
|
||||
# I2V event handler
|
||||
i2v_run_button.click(
|
||||
fn=handle_i2v_generation,
|
||||
inputs=[
|
||||
i2v_prompt,
|
||||
negative_prompt,
|
||||
use_negative_prompt,
|
||||
seed,
|
||||
guidance_scale,
|
||||
num_frames,
|
||||
height,
|
||||
width,
|
||||
num_inference_steps,
|
||||
randomize_seed,
|
||||
input_image,
|
||||
],
|
||||
outputs=[i2v_result, seed_output, i2v_error_output, i2v_frames_output, i2v_download_file],
|
||||
concurrency_limit=20,
|
||||
)
|
||||
|
||||
# Update status periodically
|
||||
demo.load(update_status, outputs=status_text)
|
||||
|
||||
return demo
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="FastVideo Gradio Frontend")
|
||||
parser.add_argument("--backend_url",
|
||||
type=str,
|
||||
default="http://localhost:8000",
|
||||
help="URL of the Ray Serve backend")
|
||||
parser.add_argument("--t2v_model_path",
|
||||
type=str,
|
||||
default="FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
|
||||
help="Path to the T2V model (for default parameters)")
|
||||
parser.add_argument("--i2v_model_path",
|
||||
type=str,
|
||||
default="Wan-AI/Wan2.2-TI2V-5B-Diffusers",
|
||||
help="Path to the I2V model (for default parameters)")
|
||||
parser.add_argument("--host",
|
||||
type=str,
|
||||
default="0.0.0.0",
|
||||
help="Host to bind to")
|
||||
parser.add_argument("--port",
|
||||
type=int,
|
||||
default=7860,
|
||||
help="Port to bind to")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
# Load default parameters from the models
|
||||
# try:
|
||||
default_params = SamplingParam.from_pretrained(args.t2v_model_path)
|
||||
# except Exception as e:
|
||||
# print(f"Warning: Could not load default parameters from {args.t2v_model_path}: {e}")
|
||||
# print("Using fallback default parameters...")
|
||||
# # Create fallback default parameters
|
||||
# default_params = SamplingParam()
|
||||
# default_params.height = 448
|
||||
# default_params.width = 832
|
||||
# default_params.num_frames = 21
|
||||
# default_params.guidance_scale = 7.5
|
||||
# default_params.num_inference_steps = 20
|
||||
# default_params.seed = 1024
|
||||
|
||||
# Create and launch the interface
|
||||
demo = create_gradio_interface(args.backend_url, default_params)
|
||||
|
||||
print(f"Starting Gradio frontend at http://{args.host}:{args.port}")
|
||||
print(f"Backend URL: {args.backend_url}")
|
||||
print(f"T2V Model: {args.t2v_model_path}")
|
||||
print(f"I2V Model: {args.i2v_model_path}")
|
||||
|
||||
demo.queue(max_size=20).launch(
|
||||
server_name=args.host,
|
||||
server_port=args.port,
|
||||
allowed_paths=[os.path.abspath("outputs"), os.path.abspath("temp_images")]
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,182 @@
|
||||
import argparse
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
import time
|
||||
import threading
|
||||
import signal
|
||||
import requests
|
||||
from pathlib import Path
|
||||
|
||||
# Add the project root to the Python path
|
||||
project_root = Path(__file__).parent.parent.parent.parent
|
||||
sys.path.insert(0, str(project_root))
|
||||
|
||||
|
||||
def check_frontend_health(frontend_url: str, max_retries: int = 30) -> bool:
|
||||
"""Check if the frontend is healthy"""
|
||||
for i in range(max_retries):
|
||||
try:
|
||||
response = requests.get(frontend_url, timeout=5)
|
||||
if response.status_code == 200:
|
||||
print(f"✅ Frontend is healthy at {frontend_url}")
|
||||
return True
|
||||
except requests.exceptions.RequestException:
|
||||
pass
|
||||
|
||||
if i < max_retries - 1:
|
||||
print(f"⏳ Waiting for frontend to start... ({i+1}/{max_retries})")
|
||||
time.sleep(2)
|
||||
|
||||
print(f"❌ Frontend failed to start within {max_retries * 2} seconds")
|
||||
return False
|
||||
|
||||
|
||||
def start_frontend_instance(args, instance_id: int, backend_url: str):
|
||||
"""Start a single frontend instance"""
|
||||
frontend_script = Path(__file__).parent / "gradio_frontend.py"
|
||||
frontend_port = args.frontend_base_port + instance_id
|
||||
|
||||
cmd = [
|
||||
sys.executable, str(frontend_script),
|
||||
"--backend_url", backend_url,
|
||||
"--t2v_model_path", args.t2v_model_path,
|
||||
"--i2v_model_path", args.i2v_model_path,
|
||||
"--host", args.frontend_host,
|
||||
"--port", str(frontend_port)
|
||||
]
|
||||
|
||||
print(f"🎨 Starting Frontend {instance_id + 1} on port {frontend_port}...")
|
||||
print(f"Command: {' '.join(cmd)}")
|
||||
|
||||
# Start the frontend process
|
||||
frontend_process = subprocess.Popen(
|
||||
cmd,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.STDOUT,
|
||||
universal_newlines=True,
|
||||
bufsize=1
|
||||
)
|
||||
|
||||
# Monitor frontend output
|
||||
def monitor_frontend():
|
||||
for line in frontend_process.stdout:
|
||||
print(f"[FRONTEND-{instance_id + 1}] {line.rstrip()}")
|
||||
|
||||
monitor_thread = threading.Thread(target=monitor_frontend, daemon=True)
|
||||
monitor_thread.start()
|
||||
|
||||
return frontend_process, frontend_port
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="FastVideo Multi-Frontend Launcher")
|
||||
|
||||
# Model and output settings
|
||||
parser.add_argument("--t2v_model_path",
|
||||
type=str,
|
||||
default="FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
|
||||
help="Path to the T2V model")
|
||||
parser.add_argument("--i2v_model_path",
|
||||
type=str,
|
||||
default="Wan-AI/Wan2.1-I2V-14B-480P-Diffusers",
|
||||
help="Path to the I2V model")
|
||||
|
||||
# Frontend settings
|
||||
parser.add_argument("--frontend_host",
|
||||
type=str,
|
||||
default="0.0.0.0",
|
||||
help="Frontend host to bind to")
|
||||
parser.add_argument("--frontend_base_port",
|
||||
type=int,
|
||||
default=7860,
|
||||
help="Base port for frontend instances")
|
||||
parser.add_argument("--num_frontends",
|
||||
type=int,
|
||||
default=2,
|
||||
help="Number of frontend instances to start")
|
||||
|
||||
# Backend settings
|
||||
parser.add_argument("--backend_url",
|
||||
type=str,
|
||||
default="http://localhost:8000",
|
||||
help="Backend URL for frontends to connect to")
|
||||
|
||||
# Other settings
|
||||
parser.add_argument("--skip_health_check",
|
||||
action="store_true",
|
||||
help="Skip frontend health check")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
print("🎬 FastVideo Multi-Frontend Launcher")
|
||||
print("=" * 50)
|
||||
print(f"T2V Model: {args.t2v_model_path}")
|
||||
print(f"I2V Model: {args.i2v_model_path}")
|
||||
print(f"Backend URL: {args.backend_url}")
|
||||
print(f"Number of Frontends: {args.num_frontends}")
|
||||
print(f"Frontend Base Port: {args.frontend_base_port}")
|
||||
print("=" * 50)
|
||||
|
||||
# Start multiple frontend instances
|
||||
frontend_processes = []
|
||||
frontend_urls = []
|
||||
|
||||
for i in range(args.num_frontends):
|
||||
process, port = start_frontend_instance(args, i, args.backend_url)
|
||||
frontend_processes.append(process)
|
||||
frontend_urls.append(f"http://{args.frontend_host}:{port}")
|
||||
|
||||
# Wait for frontends to be ready
|
||||
if not args.skip_health_check:
|
||||
print("\n⏳ Waiting for frontends to start...")
|
||||
for i, url in enumerate(frontend_urls):
|
||||
if not check_frontend_health(url):
|
||||
print(f"❌ Frontend {i + 1} failed to start. Terminating...")
|
||||
for process in frontend_processes:
|
||||
process.terminate()
|
||||
sys.exit(1)
|
||||
|
||||
print("\n🎉 All frontend instances are starting up!")
|
||||
for i, url in enumerate(frontend_urls):
|
||||
print(f"📺 Frontend {i + 1}: {url}")
|
||||
print("\nPress Ctrl+C to stop all frontend instances...")
|
||||
|
||||
# Signal handler for graceful shutdown
|
||||
def signal_handler(signum, frame):
|
||||
print("\n🛑 Shutting down frontend instances...")
|
||||
for process in frontend_processes:
|
||||
process.terminate()
|
||||
|
||||
# Wait for processes to terminate
|
||||
try:
|
||||
for process in frontend_processes:
|
||||
process.wait(timeout=5)
|
||||
except subprocess.TimeoutExpired:
|
||||
print("⚠️ Force killing processes...")
|
||||
for process in frontend_processes:
|
||||
process.kill()
|
||||
|
||||
print("✅ Frontend instances stopped")
|
||||
sys.exit(0)
|
||||
|
||||
signal.signal(signal.SIGINT, signal_handler)
|
||||
signal.signal(signal.SIGTERM, signal_handler)
|
||||
|
||||
# Monitor processes
|
||||
try:
|
||||
while True:
|
||||
# Check if processes are still running
|
||||
for i, process in enumerate(frontend_processes):
|
||||
if process.poll() is not None:
|
||||
print(f"❌ Frontend {i + 1} process died unexpectedly")
|
||||
break
|
||||
|
||||
time.sleep(1)
|
||||
|
||||
except KeyboardInterrupt:
|
||||
signal_handler(signal.SIGINT, None)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,181 @@
|
||||
# Nginx configuration for FastVideo load balancing
|
||||
# This configuration implements the architecture:
|
||||
# ngrok -> nginx reverse proxy -> frontend1/frontend2 -> backend1×8/backend2×8
|
||||
|
||||
events {
|
||||
worker_connections 1024;
|
||||
}
|
||||
|
||||
http {
|
||||
# Basic settings
|
||||
sendfile on;
|
||||
tcp_nopush on;
|
||||
tcp_nodelay on;
|
||||
keepalive_timeout 65;
|
||||
types_hash_max_size 2048;
|
||||
client_max_body_size 100M; # Allow large video uploads
|
||||
|
||||
# Logging
|
||||
access_log /mnt/fast-disks/nfs/hao_lab/FastVideo/outputs/nginx_access.log;
|
||||
error_log /mnt/fast-disks/nfs/hao_lab/FastVideo/outputs/nginx_error.log;
|
||||
|
||||
# Gzip compression
|
||||
gzip on;
|
||||
gzip_vary on;
|
||||
gzip_min_length 1024;
|
||||
gzip_proxied any;
|
||||
gzip_comp_level 6;
|
||||
gzip_types
|
||||
text/plain
|
||||
text/css
|
||||
text/xml
|
||||
text/javascript
|
||||
application/json
|
||||
application/javascript
|
||||
application/xml+rss
|
||||
application/atom+xml
|
||||
image/svg+xml;
|
||||
|
||||
# Upstream for frontend load balancing
|
||||
upstream frontend_servers {
|
||||
# Round-robin load balancing between frontends
|
||||
server 127.0.0.1:7860 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:7861 weight=1 max_fails=3 fail_timeout=30s;
|
||||
upstream frontend_servers {
|
||||
# Round-robin load balancing between frontends
|
||||
server 127.0.0.1:7860 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:7861 weight=1 max_fails=3 fail_timeout=30s;
|
||||
|
||||
# Health check
|
||||
keepalive 32;
|
||||
}
|
||||
|
||||
# Upstream for backend1 load balancing
|
||||
upstream backend1_servers {
|
||||
server 127.0.0.1:8000 weight=1 max_fails=3 fail_timeout=30s; (8 replicas)
|
||||
upstream backend1_servers {
|
||||
# Round-robin load balancing for backend1 replicas
|
||||
server 127.0.0.1:8000 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8001 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8002 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8003 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8004 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8005 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8006 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8007 weight=1 max_fails=3 fail_timeout=30s;
|
||||
|
||||
keepalive 32;
|
||||
}
|
||||
|
||||
# Upstream for backend2 load balancing
|
||||
upstream backend2_servers {
|
||||
server 127.0.0.1:8000 weight=1 max_fails=3 fail_timeout=30s; (8 replicas)
|
||||
upstream backend2_servers {
|
||||
# Round-robin load balancing for backend2 replicas
|
||||
server 127.0.0.1:8010 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8011 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8012 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8013 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8014 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8015 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8016 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8017 weight=1 max_fails=3 fail_timeout=30s;
|
||||
|
||||
keepalive 32;
|
||||
}
|
||||
|
||||
# Main server block
|
||||
server {
|
||||
listen 80;
|
||||
server_name localhost;
|
||||
|
||||
# Security headers
|
||||
add_header X-Frame-Options "SAMEORIGIN" always;
|
||||
add_header X-Content-Type-Options "nosniff" always;
|
||||
add_header X-XSS-Protection "1; mode=block" always;
|
||||
add_header Referrer-Policy "no-referrer-when-downgrade" always;
|
||||
|
||||
# Frontend routes (Gradio interfaces)
|
||||
location / {
|
||||
proxy_pass http://frontend_servers;
|
||||
proxy_set_header Host $host;
|
||||
proxy_set_header X-Real-IP $remote_addr;
|
||||
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
|
||||
proxy_set_header X-Forwarded-Proto $scheme;
|
||||
|
||||
# WebSocket support for Gradio
|
||||
proxy_http_version 1.1;
|
||||
proxy_set_header Upgrade $http_upgrade;
|
||||
proxy_set_header Connection "upgrade";
|
||||
|
||||
# Timeouts
|
||||
proxy_connect_timeout 60s;
|
||||
proxy_send_timeout 60s;
|
||||
proxy_read_timeout 60s;
|
||||
|
||||
# Buffer settings
|
||||
proxy_buffering on;
|
||||
proxy_buffer_size 128k;
|
||||
proxy_buffers 4 256k;
|
||||
proxy_busy_buffers_size 256k;
|
||||
}
|
||||
|
||||
# Backend API routes for frontend1
|
||||
location /api/frontend1/ {
|
||||
# Strip the /api/frontend1/ prefix
|
||||
rewrite ^/api/frontend1/(.*) /$1 break;
|
||||
proxy_pass http://backend1_servers;
|
||||
proxy_set_header Host $host;
|
||||
proxy_set_header X-Real-IP $remote_addr;
|
||||
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
|
||||
proxy_set_header X-Forwarded-Proto $scheme;
|
||||
|
||||
# Timeouts for video generation
|
||||
proxy_connect_timeout 300s;
|
||||
proxy_send_timeout 300s;
|
||||
proxy_read_timeout 300s;
|
||||
|
||||
# Buffer settings for large responses
|
||||
proxy_buffering on;
|
||||
proxy_buffer_size 128k;
|
||||
proxy_buffers 4 256k;
|
||||
proxy_busy_buffers_size 256k;
|
||||
}
|
||||
|
||||
# Backend API routes for frontend2
|
||||
location /api/frontend2/ {
|
||||
# Strip the /api/frontend2/ prefix
|
||||
rewrite ^/api/frontend2/(.*) /$1 break;
|
||||
proxy_pass http://backend2_servers;
|
||||
proxy_set_header Host $host;
|
||||
proxy_set_header X-Real-IP $remote_addr;
|
||||
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
|
||||
proxy_set_header X-Forwarded-Proto $scheme;
|
||||
|
||||
# Timeouts for video generation
|
||||
proxy_connect_timeout 300s;
|
||||
proxy_send_timeout 300s;
|
||||
proxy_read_timeout 300s;
|
||||
|
||||
# Buffer settings for large responses
|
||||
proxy_buffering on;
|
||||
proxy_buffer_size 128k;
|
||||
proxy_buffers 4 256k;
|
||||
proxy_busy_buffers_size 256k;
|
||||
}
|
||||
|
||||
# Health check endpoint
|
||||
location /health {
|
||||
access_log off;
|
||||
return 200 "healthy\n";
|
||||
add_header Content-Type text/plain;
|
||||
}
|
||||
|
||||
# Static files (if needed)
|
||||
location /static/ {
|
||||
alias /var/www/static/;
|
||||
expires 1y;
|
||||
add_header Cache-Control "public, immutable";
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,173 @@
|
||||
# Nginx configuration for FastVideo load balancing
|
||||
# This configuration implements the architecture:
|
||||
# ngrok -> nginx reverse proxy -> frontend1/frontend2 -> backend1×8/backend2×8
|
||||
|
||||
events {
|
||||
worker_connections 1024;
|
||||
}
|
||||
|
||||
http {
|
||||
# Basic settings
|
||||
sendfile on;
|
||||
tcp_nopush on;
|
||||
tcp_nodelay on;
|
||||
keepalive_timeout 65;
|
||||
types_hash_max_size 2048;
|
||||
client_max_body_size 100M; # Allow large video uploads
|
||||
|
||||
# Logging
|
||||
access_log /var/log/nginx/access.log;
|
||||
error_log /var/log/nginx/error.log;
|
||||
|
||||
# Gzip compression
|
||||
gzip on;
|
||||
gzip_vary on;
|
||||
gzip_min_length 1024;
|
||||
gzip_proxied any;
|
||||
gzip_comp_level 6;
|
||||
gzip_types
|
||||
text/plain
|
||||
text/css
|
||||
text/xml
|
||||
text/javascript
|
||||
application/json
|
||||
application/javascript
|
||||
application/xml+rss
|
||||
application/atom+xml
|
||||
image/svg+xml;
|
||||
|
||||
# Upstream for frontend load balancing
|
||||
upstream frontend_servers {
|
||||
# Round-robin load balancing between frontends
|
||||
server 127.0.0.1:7860 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:7861 weight=1 max_fails=3 fail_timeout=30s;
|
||||
|
||||
# Health check
|
||||
keepalive 32;
|
||||
}
|
||||
|
||||
# Upstream for backend1 load balancing (8 replicas)
|
||||
upstream backend1_servers {
|
||||
# Round-robin load balancing for backend1 replicas
|
||||
server 127.0.0.1:8000 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8001 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8002 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8003 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8004 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8005 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8006 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8007 weight=1 max_fails=3 fail_timeout=30s;
|
||||
|
||||
keepalive 32;
|
||||
}
|
||||
|
||||
# Upstream for backend2 load balancing (8 replicas)
|
||||
upstream backend2_servers {
|
||||
# Round-robin load balancing for backend2 replicas
|
||||
server 127.0.0.1:8010 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8011 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8012 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8013 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8014 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8015 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8016 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8017 weight=1 max_fails=3 fail_timeout=30s;
|
||||
|
||||
keepalive 32;
|
||||
}
|
||||
|
||||
# Main server block
|
||||
server {
|
||||
listen 80;
|
||||
server_name localhost;
|
||||
|
||||
# Security headers
|
||||
add_header X-Frame-Options "SAMEORIGIN" always;
|
||||
add_header X-Content-Type-Options "nosniff" always;
|
||||
add_header X-XSS-Protection "1; mode=block" always;
|
||||
add_header Referrer-Policy "no-referrer-when-downgrade" always;
|
||||
|
||||
# Frontend routes (Gradio interfaces)
|
||||
location / {
|
||||
proxy_pass http://frontend_servers;
|
||||
proxy_set_header Host $host;
|
||||
proxy_set_header X-Real-IP $remote_addr;
|
||||
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
|
||||
proxy_set_header X-Forwarded-Proto $scheme;
|
||||
|
||||
# WebSocket support for Gradio
|
||||
proxy_http_version 1.1;
|
||||
proxy_set_header Upgrade $http_upgrade;
|
||||
proxy_set_header Connection "upgrade";
|
||||
|
||||
# Timeouts
|
||||
proxy_connect_timeout 60s;
|
||||
proxy_send_timeout 60s;
|
||||
proxy_read_timeout 60s;
|
||||
|
||||
# Buffer settings
|
||||
proxy_buffering on;
|
||||
proxy_buffer_size 128k;
|
||||
proxy_buffers 4 256k;
|
||||
proxy_busy_buffers_size 256k;
|
||||
}
|
||||
|
||||
# Backend API routes for frontend1
|
||||
location /api/frontend1/ {
|
||||
# Strip the /api/frontend1/ prefix
|
||||
rewrite ^/api/frontend1/(.*) /$1 break;
|
||||
proxy_pass http://backend1_servers;
|
||||
proxy_set_header Host $host;
|
||||
proxy_set_header X-Real-IP $remote_addr;
|
||||
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
|
||||
proxy_set_header X-Forwarded-Proto $scheme;
|
||||
|
||||
# Timeouts for video generation
|
||||
proxy_connect_timeout 300s;
|
||||
proxy_send_timeout 300s;
|
||||
proxy_read_timeout 300s;
|
||||
|
||||
# Buffer settings for large responses
|
||||
proxy_buffering on;
|
||||
proxy_buffer_size 128k;
|
||||
proxy_buffers 4 256k;
|
||||
proxy_busy_buffers_size 256k;
|
||||
}
|
||||
|
||||
# Backend API routes for frontend2
|
||||
location /api/frontend2/ {
|
||||
# Strip the /api/frontend2/ prefix
|
||||
rewrite ^/api/frontend2/(.*) /$1 break;
|
||||
proxy_pass http://backend2_servers;
|
||||
proxy_set_header Host $host;
|
||||
proxy_set_header X-Real-IP $remote_addr;
|
||||
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
|
||||
proxy_set_header X-Forwarded-Proto $scheme;
|
||||
|
||||
# Timeouts for video generation
|
||||
proxy_connect_timeout 300s;
|
||||
proxy_send_timeout 300s;
|
||||
proxy_read_timeout 300s;
|
||||
|
||||
# Buffer settings for large responses
|
||||
proxy_buffering on;
|
||||
proxy_buffer_size 128k;
|
||||
proxy_buffers 4 256k;
|
||||
proxy_busy_buffers_size 256k;
|
||||
}
|
||||
|
||||
# Health check endpoint
|
||||
location /health {
|
||||
access_log off;
|
||||
return 200 "healthy\n";
|
||||
add_header Content-Type text/plain;
|
||||
}
|
||||
|
||||
# Static files (if needed)
|
||||
location /static/ {
|
||||
alias /var/www/static/;
|
||||
expires 1y;
|
||||
add_header Cache-Control "public, immutable";
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,337 @@
|
||||
import time
|
||||
import os
|
||||
import torch
|
||||
import base64
|
||||
import io
|
||||
from copy import deepcopy
|
||||
from typing import Dict, Any, Optional, List
|
||||
|
||||
import ray
|
||||
from ray import serve
|
||||
from fastapi import FastAPI, Request
|
||||
from pydantic import BaseModel
|
||||
from PIL import Image
|
||||
import numpy as np
|
||||
from slowapi import Limiter, _rate_limit_exceeded_handler
|
||||
from slowapi.util import get_remote_address
|
||||
from slowapi.errors import RateLimitExceeded
|
||||
|
||||
|
||||
class VideoGenerationRequest(BaseModel):
|
||||
prompt: str
|
||||
negative_prompt: Optional[str] = None
|
||||
use_negative_prompt: bool = False
|
||||
seed: int = 42
|
||||
guidance_scale: float = 7.5
|
||||
num_frames: int = 21
|
||||
height: int = 448
|
||||
width: int = 832
|
||||
num_inference_steps: int = 20
|
||||
randomize_seed: bool = False
|
||||
return_frames: bool = False # Whether to return base64 encoded frames
|
||||
image_path: Optional[str] = None # Path to input image for I2V
|
||||
model_type: str = "t2v" # "t2v" or "i2v" to specify which model to use
|
||||
|
||||
|
||||
class VideoGenerationResponse(BaseModel):
|
||||
output_path: str
|
||||
seed: int
|
||||
success: bool
|
||||
error_message: Optional[str] = None
|
||||
frames: Optional[List[str]] = None # Base64 encoded frames
|
||||
|
||||
|
||||
def encode_frames_to_base64(frames: List[np.ndarray]) -> List[str]:
|
||||
"""Convert numpy frames (0-255) to base64-encoded PNG images"""
|
||||
if not frames:
|
||||
return []
|
||||
|
||||
encoded_frames = []
|
||||
|
||||
for i, frame in enumerate(frames):
|
||||
try:
|
||||
# Ensure frame is numpy array
|
||||
if not isinstance(frame, np.ndarray):
|
||||
print(f"Warning: Frame {i} is not a numpy array, skipping")
|
||||
continue
|
||||
|
||||
# Ensure frame is uint8
|
||||
if frame.dtype != np.uint8:
|
||||
# Clip values to 0-255 range and convert to uint8
|
||||
frame = np.clip(frame, 0, 255).astype(np.uint8)
|
||||
|
||||
# Convert numpy array to PIL Image
|
||||
if len(frame.shape) == 3 and frame.shape[2] == 3:
|
||||
# RGB image
|
||||
pil_image = Image.fromarray(frame, mode='RGB')
|
||||
elif len(frame.shape) == 3 and frame.shape[2] == 4:
|
||||
# RGBA image
|
||||
pil_image = Image.fromarray(frame, mode='RGBA')
|
||||
elif len(frame.shape) == 2:
|
||||
# Grayscale image
|
||||
pil_image = Image.fromarray(frame, mode='L')
|
||||
else:
|
||||
print(f"Warning: Frame {i} has unsupported shape {frame.shape}, skipping")
|
||||
continue
|
||||
|
||||
# Save to bytes buffer as PNG
|
||||
buffer = io.BytesIO()
|
||||
pil_image.save(buffer, format='PNG')
|
||||
buffer.seek(0)
|
||||
|
||||
# Encode to base64
|
||||
img_base64 = base64.b64encode(buffer.getvalue()).decode('utf-8')
|
||||
encoded_frames.append(f"data:image/png;base64,{img_base64}")
|
||||
|
||||
except Exception as e:
|
||||
print(f"Warning: Failed to encode frame {i}: {e}")
|
||||
continue
|
||||
|
||||
return encoded_frames
|
||||
|
||||
|
||||
# Create FastAPI app with rate limiting
|
||||
app = FastAPI()
|
||||
|
||||
# Initialize rate limiter
|
||||
limiter = Limiter(key_func=get_remote_address)
|
||||
app.state.limiter = limiter
|
||||
app.add_exception_handler(RateLimitExceeded, _rate_limit_exceeded_handler)
|
||||
|
||||
|
||||
@serve.deployment(
|
||||
num_replicas=8,
|
||||
# ray_actor_options={"num_cpus": 10, "num_gpus": 1, "runtime_env": {"conda": "fv", "working_dir": "/mnt/fast-disks/nfs/hao_lab/FastVideo"}},
|
||||
ray_actor_options={"num_cpus": 10, "num_gpus": 1, "runtime_env": {"conda": "fv"}},
|
||||
)
|
||||
@serve.ingress(app)
|
||||
class FastVideoAPI:
|
||||
def __init__(self, t2v_model_path: str, i2v_model_path: str, output_path: str):
|
||||
self.t2v_model_path = t2v_model_path
|
||||
self.i2v_model_path = i2v_model_path
|
||||
self.output_path = output_path
|
||||
|
||||
# Initialize the video generators
|
||||
self.t2v_generator = None # Initialize to None
|
||||
self.i2v_generator = None # Initialize to None
|
||||
self.t2v_default_params = None # Initialize to None
|
||||
self.i2v_default_params = None # Initialize to None
|
||||
|
||||
# Ensure output directory exists
|
||||
os.makedirs(output_path, exist_ok=True)
|
||||
time.sleep(10)
|
||||
self._initialize_models() # Ensure models are initialized
|
||||
|
||||
def _initialize_models(self):
|
||||
# Set VSA environment variable
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
|
||||
# os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "FLASH_ATTN"
|
||||
|
||||
# Import only when needed - use direct imports to avoid module-level execution
|
||||
from fastvideo.entrypoints.video_generator import VideoGenerator
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
# Initialize T2V model
|
||||
if self.t2v_generator is None:
|
||||
print(f"Initializing T2V model: {self.t2v_model_path}")
|
||||
self.t2v_generator = VideoGenerator.from_pretrained(
|
||||
model_path=self.t2v_model_path,
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
# Adjust these offload parameters if you have < 32GB of VRAM
|
||||
text_encoder_cpu_offload=False,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
VSA_sparsity=0.8,
|
||||
# master_port=port,
|
||||
)
|
||||
self.t2v_default_params = SamplingParam.from_pretrained(self.t2v_model_path)
|
||||
print("✅ T2V model initialized successfully")
|
||||
|
||||
# Initialize I2V model
|
||||
# if self.i2v_generator is None:
|
||||
if False:
|
||||
print(f"Initializing I2V model: {self.i2v_model_path}")
|
||||
self.i2v_generator = VideoGenerator.from_pretrained(
|
||||
model_path=self.i2v_model_path,
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
# Adjust these offload parameters if you have < 32GB of VRAM
|
||||
text_encoder_cpu_offload=False,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
VSA_sparsity=0.8,
|
||||
# master_port=port,
|
||||
)
|
||||
self.i2v_default_params = SamplingParam.from_pretrained(self.i2v_model_path)
|
||||
print("✅ I2V model initialized successfully")
|
||||
|
||||
@app.post("/generate_video", response_model=VideoGenerationResponse)
|
||||
@limiter.limit("50/minute") # Allow 2 requests per minute per IP
|
||||
async def generate_video(self, request: Request, video_request: VideoGenerationRequest) -> VideoGenerationResponse:
|
||||
try:
|
||||
# Select the appropriate model and parameters based on model_type
|
||||
if video_request.model_type.lower() == "i2v":
|
||||
generator = self.i2v_generator
|
||||
params = deepcopy(self.i2v_default_params)
|
||||
print(f"Using I2V model for generation")
|
||||
else:
|
||||
generator = self.t2v_generator
|
||||
params = deepcopy(self.t2v_default_params)
|
||||
print(f"Using T2V model for generation")
|
||||
|
||||
# Update parameters with request values
|
||||
params.prompt = video_request.prompt
|
||||
# Only override negative prompt if user explicitly opts in
|
||||
if video_request.use_negative_prompt:
|
||||
params.negative_prompt = video_request.negative_prompt
|
||||
|
||||
params.seed = video_request.seed
|
||||
params.guidance_scale = video_request.guidance_scale
|
||||
params.num_frames = video_request.num_frames
|
||||
params.height = video_request.height
|
||||
params.width = video_request.width
|
||||
params.num_inference_steps = video_request.num_inference_steps
|
||||
|
||||
# Handle seed randomization
|
||||
if video_request.randomize_seed:
|
||||
params.seed = torch.randint(0, 1000000, (1,)).item()
|
||||
|
||||
# Ensure negative_prompt is a non-None string; FastVideo validation disallows None
|
||||
if params.negative_prompt is None:
|
||||
params.negative_prompt = "" # empty string satisfies validator
|
||||
|
||||
# Set up output path and video saving
|
||||
params.save_video = True
|
||||
params.output_path = self.output_path
|
||||
# params.return_frames = False # avoid keeping frames in memory
|
||||
|
||||
# Create a clean filename from the prompt
|
||||
safe_prompt = video_request.prompt[:100].replace(' ', '_').replace('/', '_').replace('\\', '_')
|
||||
|
||||
# Store desired video name inside the SamplingParam to avoid unknown kwarg errors
|
||||
setattr(params, "output_video_name", safe_prompt)
|
||||
|
||||
# Handle image_path for I2V
|
||||
if video_request.image_path:
|
||||
params.image_path = video_request.image_path
|
||||
|
||||
# Generate the video with proper output path and filename
|
||||
result = generator.generate_video(
|
||||
prompt=video_request.prompt,
|
||||
sampling_param=params,
|
||||
save_video=True, # Match the params.save_video setting
|
||||
)
|
||||
|
||||
# The actual output path where the video was saved
|
||||
output_path = os.path.join(self.output_path, f"{safe_prompt}.mp4")
|
||||
|
||||
# Verify the file exists
|
||||
if not os.path.exists(output_path):
|
||||
raise FileNotFoundError(f"Video was not saved to expected location: {output_path}")
|
||||
|
||||
frames = result.get("frames", [])
|
||||
|
||||
# Encode frames to base64 for web transmission only if requested
|
||||
encoded_frames = None
|
||||
if video_request.return_frames and frames:
|
||||
try:
|
||||
encoded_frames = encode_frames_to_base64(frames)
|
||||
except Exception as e:
|
||||
print(f"Warning: Failed to encode frames: {e}")
|
||||
encoded_frames = None
|
||||
|
||||
response = VideoGenerationResponse(
|
||||
output_path=output_path,
|
||||
frames=encoded_frames,
|
||||
seed=params.seed,
|
||||
success=True
|
||||
)
|
||||
|
||||
# Memory cleanup to avoid OOM in repeated generations
|
||||
import gc
|
||||
gc.collect()
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
return response
|
||||
except Exception as e:
|
||||
return VideoGenerationResponse(
|
||||
output_path="",
|
||||
seed=video_request.seed,
|
||||
success=False,
|
||||
error_message=str(e)
|
||||
)
|
||||
|
||||
@app.get("/health")
|
||||
@limiter.limit("10/minute") # Allow 10 health checks per minute per IP
|
||||
async def health_check(self, request: Request):
|
||||
return {"status": "healthy"}
|
||||
|
||||
|
||||
def start_ray_serve(
|
||||
t2v_model_path: str = "Wan-AI/Wan2.2-TI2V-5B-Diffusers",
|
||||
i2v_model_path: str = "Wan-AI/Wan2.2-TI2V-5B-Diffusers",
|
||||
output_path: str = "outputs",
|
||||
host: str = "0.0.0.0",
|
||||
port: int = 8000
|
||||
):
|
||||
"""Start the Ray Serve backend"""
|
||||
# Initialize Ray
|
||||
if not ray.is_initialized():
|
||||
ray.init()
|
||||
|
||||
# Deploy the API
|
||||
api = FastVideoAPI.bind(t2v_model_path, i2v_model_path, output_path)
|
||||
serve.run(api, route_prefix="/", name="fast_video") # detach
|
||||
|
||||
print(f"Ray Serve backend started at http://{host}:{port}")
|
||||
print(f"T2V Model: {t2v_model_path}")
|
||||
print(f"I2V Model: {i2v_model_path}")
|
||||
print(f"Health check: http://{host}:{port}/health")
|
||||
print(f"Video generation endpoint: http://{host}:{port}/generate_video")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import argparse
|
||||
|
||||
parser = argparse.ArgumentParser(description="FastVideo Ray Serve Backend")
|
||||
parser.add_argument("--t2v_model_path",
|
||||
type=str,
|
||||
default="FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
|
||||
help="Path to the T2V model")
|
||||
parser.add_argument("--i2v_model_path",
|
||||
type=str,
|
||||
default="Wan-AI/Wan2.2-TI2V-5B-Diffusers",
|
||||
help="Path to the I2V model")
|
||||
parser.add_argument("--output_path",
|
||||
type=str,
|
||||
default="outputs",
|
||||
help="Path to save generated videos")
|
||||
parser.add_argument("--host",
|
||||
type=str,
|
||||
default="0.0.0.0",
|
||||
help="Host to bind to")
|
||||
parser.add_argument("--port",
|
||||
type=int,
|
||||
default=8000,
|
||||
help="Port to bind to")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
start_ray_serve(
|
||||
t2v_model_path=args.t2v_model_path,
|
||||
i2v_model_path=args.i2v_model_path,
|
||||
output_path=args.output_path,
|
||||
host=args.host,
|
||||
port=args.port,
|
||||
)
|
||||
|
||||
# ---- keep the process alive ---------------------------------
|
||||
import signal, sys, time
|
||||
signal.signal(signal.SIGINT, lambda *_: sys.exit(0)) # Ctrl-C
|
||||
signal.signal(signal.SIGTERM, lambda *_: sys.exit(0)) # docker stop etc.
|
||||
|
||||
print("✅ FastVideo backend is running. Press Ctrl-C to stop.")
|
||||
while True:
|
||||
time.sleep(3600)
|
||||
@@ -0,0 +1,334 @@
|
||||
import time
|
||||
import os
|
||||
import torch
|
||||
import base64
|
||||
import io
|
||||
from copy import deepcopy
|
||||
from typing import Dict, Any, Optional, List
|
||||
|
||||
import ray
|
||||
from ray import serve
|
||||
from fastapi import FastAPI, Request
|
||||
from pydantic import BaseModel
|
||||
from PIL import Image
|
||||
import numpy as np
|
||||
from slowapi import Limiter, _rate_limit_exceeded_handler
|
||||
from slowapi.util import get_remote_address
|
||||
from slowapi.errors import RateLimitExceeded
|
||||
|
||||
|
||||
class VideoGenerationRequest(BaseModel):
|
||||
prompt: str
|
||||
negative_prompt: Optional[str] = None
|
||||
use_negative_prompt: bool = False
|
||||
seed: int = 42
|
||||
guidance_scale: float = 7.5
|
||||
num_frames: int = 21
|
||||
height: int = 448
|
||||
width: int = 832
|
||||
num_inference_steps: int = 20
|
||||
randomize_seed: bool = False
|
||||
return_frames: bool = False # Whether to return base64 encoded frames
|
||||
image_path: Optional[str] = None # Path to input image for I2V
|
||||
model_type: str = "t2v" # "t2v" or "i2v" to specify which model to use
|
||||
|
||||
|
||||
class VideoGenerationResponse(BaseModel):
|
||||
output_path: str
|
||||
seed: int
|
||||
success: bool
|
||||
error_message: Optional[str] = None
|
||||
frames: Optional[List[str]] = None # Base64 encoded frames
|
||||
|
||||
|
||||
def encode_frames_to_base64(frames: List[np.ndarray]) -> List[str]:
|
||||
"""Convert numpy frames (0-255) to base64-encoded PNG images"""
|
||||
if not frames:
|
||||
return []
|
||||
|
||||
encoded_frames = []
|
||||
|
||||
for i, frame in enumerate(frames):
|
||||
try:
|
||||
# Ensure frame is numpy array
|
||||
if not isinstance(frame, np.ndarray):
|
||||
print(f"Warning: Frame {i} is not a numpy array, skipping")
|
||||
continue
|
||||
|
||||
# Ensure frame is uint8
|
||||
if frame.dtype != np.uint8:
|
||||
# Clip values to 0-255 range and convert to uint8
|
||||
frame = np.clip(frame, 0, 255).astype(np.uint8)
|
||||
|
||||
# Convert numpy array to PIL Image
|
||||
if len(frame.shape) == 3 and frame.shape[2] == 3:
|
||||
# RGB image
|
||||
pil_image = Image.fromarray(frame, mode='RGB')
|
||||
elif len(frame.shape) == 3 and frame.shape[2] == 4:
|
||||
# RGBA image
|
||||
pil_image = Image.fromarray(frame, mode='RGBA')
|
||||
elif len(frame.shape) == 2:
|
||||
# Grayscale image
|
||||
pil_image = Image.fromarray(frame, mode='L')
|
||||
else:
|
||||
print(f"Warning: Frame {i} has unsupported shape {frame.shape}, skipping")
|
||||
continue
|
||||
|
||||
# Save to bytes buffer as PNG
|
||||
buffer = io.BytesIO()
|
||||
pil_image.save(buffer, format='PNG')
|
||||
buffer.seek(0)
|
||||
|
||||
# Encode to base64
|
||||
img_base64 = base64.b64encode(buffer.getvalue()).decode('utf-8')
|
||||
encoded_frames.append(f"data:image/png;base64,{img_base64}")
|
||||
|
||||
except Exception as e:
|
||||
print(f"Warning: Failed to encode frame {i}: {e}")
|
||||
continue
|
||||
|
||||
return encoded_frames
|
||||
|
||||
|
||||
# Create FastAPI app with rate limiting
|
||||
app = FastAPI()
|
||||
|
||||
# Initialize rate limiter
|
||||
limiter = Limiter(key_func=get_remote_address)
|
||||
app.state.limiter = limiter
|
||||
app.add_exception_handler(RateLimitExceeded, _rate_limit_exceeded_handler)
|
||||
|
||||
|
||||
@serve.deployment(
|
||||
num_replicas=3, # Set to 3 for cluster 0
|
||||
ray_actor_options={
|
||||
"num_cpus": 10,
|
||||
"num_gpus": 1,
|
||||
"runtime_env": {"conda": "fv"},
|
||||
},
|
||||
)
|
||||
@serve.ingress(app)
|
||||
class FastVideoMultiGPUAPI:
|
||||
def __init__(self, t2v_model_path: str, i2v_model_path: str, output_path: str, gpu_id: int = 0):
|
||||
self.t2v_model_path = t2v_model_path
|
||||
self.i2v_model_path = i2v_model_path
|
||||
self.output_path = output_path
|
||||
self.gpu_id = gpu_id
|
||||
|
||||
# Initialize the video generators
|
||||
self.t2v_generator = None
|
||||
self.i2v_generator = None
|
||||
self.t2v_default_params = None
|
||||
self.i2v_default_params = None
|
||||
|
||||
# Ensure output directory exists
|
||||
os.makedirs(output_path, exist_ok=True)
|
||||
time.sleep(10)
|
||||
self._initialize_models()
|
||||
|
||||
def _initialize_models(self):
|
||||
# Set VSA environment variable
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "FLASH_ATTN"
|
||||
|
||||
# Import only when needed
|
||||
from fastvideo.entrypoints.video_generator import VideoGenerator
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
# Initialize T2V model
|
||||
if False: # Disabled for now
|
||||
print(f"Initializing T2V model on GPU {self.gpu_id}: {self.t2v_model_path}")
|
||||
self.t2v_generator = VideoGenerator.from_pretrained(
|
||||
model_path=self.t2v_model_path,
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
text_encoder_cpu_offload=False,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
VSA_sparsity=0.8,
|
||||
)
|
||||
self.t2v_default_params = SamplingParam.from_pretrained(self.t2v_model_path)
|
||||
print(f"✅ T2V model initialized successfully on GPU {self.gpu_id}")
|
||||
|
||||
# Initialize I2V model
|
||||
if self.i2v_generator is None:
|
||||
print(f"Initializing I2V model on GPU {self.gpu_id}: {self.i2v_model_path}")
|
||||
self.i2v_generator = VideoGenerator.from_pretrained(
|
||||
model_path=self.i2v_model_path,
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
text_encoder_cpu_offload=False,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
VSA_sparsity=0.8,
|
||||
)
|
||||
self.i2v_default_params = SamplingParam.from_pretrained(self.i2v_model_path)
|
||||
print(f"✅ I2V model initialized successfully on GPU {self.gpu_id}")
|
||||
|
||||
@app.post("/generate_video", response_model=VideoGenerationResponse)
|
||||
@limiter.limit("2/minute") # Allow 2 requests per minute per IP
|
||||
async def generate_video(self, request: Request, video_request: VideoGenerationRequest) -> VideoGenerationResponse:
|
||||
try:
|
||||
# Select the appropriate model and parameters based on model_type
|
||||
if video_request.model_type.lower() == "i2v":
|
||||
generator = self.i2v_generator
|
||||
params = deepcopy(self.i2v_default_params)
|
||||
print(f"Using I2V model for generation on GPU {self.gpu_id}")
|
||||
else:
|
||||
generator = self.t2v_generator
|
||||
params = deepcopy(self.t2v_default_params)
|
||||
print(f"Using T2V model for generation on GPU {self.gpu_id}")
|
||||
|
||||
# Update parameters with request values
|
||||
params.prompt = video_request.prompt
|
||||
|
||||
# Handle seed randomization
|
||||
if video_request.randomize_seed:
|
||||
params.seed = torch.randint(0, 1000000, (1,)).item()
|
||||
|
||||
# Ensure negative_prompt is a non-None string
|
||||
if params.negative_prompt is None:
|
||||
params.negative_prompt = ""
|
||||
|
||||
# Set up output path and video saving
|
||||
params.save_video = True
|
||||
params.output_path = self.output_path
|
||||
|
||||
# Create a clean filename from the prompt
|
||||
safe_prompt = video_request.prompt[:100].replace(' ', '_').replace('/', '_').replace('\\', '_')
|
||||
setattr(params, "output_video_name", safe_prompt)
|
||||
|
||||
# Handle image_path for I2V
|
||||
if video_request.image_path:
|
||||
params.image_path = video_request.image_path
|
||||
|
||||
# Generate the video
|
||||
result = generator.generate_video(
|
||||
prompt=video_request.prompt,
|
||||
sampling_param=params,
|
||||
save_video=True,
|
||||
)
|
||||
|
||||
frames = result.get("frames", [])
|
||||
|
||||
# Encode frames to base64 for web transmission only if requested
|
||||
encoded_frames = None
|
||||
if video_request.return_frames and frames:
|
||||
try:
|
||||
encoded_frames = encode_frames_to_base64(frames)
|
||||
except Exception as e:
|
||||
print(f"Warning: Failed to encode frames: {e}")
|
||||
encoded_frames = None
|
||||
|
||||
response = VideoGenerationResponse(
|
||||
output_path="",
|
||||
frames=encoded_frames,
|
||||
seed=params.seed,
|
||||
success=True
|
||||
)
|
||||
|
||||
# Memory cleanup to avoid OOM in repeated generations
|
||||
import gc
|
||||
gc.collect()
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
return response
|
||||
except Exception as e:
|
||||
return VideoGenerationResponse(
|
||||
output_path="",
|
||||
seed=video_request.seed,
|
||||
success=False,
|
||||
error_message=str(e)
|
||||
)
|
||||
|
||||
@app.get("/health")
|
||||
@limiter.limit("10/minute") # Allow 10 health checks per minute per IP
|
||||
async def health_check(self, request: Request):
|
||||
return {"status": "healthy", "gpu_id": self.gpu_id}
|
||||
|
||||
|
||||
def start_ray_serve_multi_gpu(
|
||||
t2v_model_path: str = "FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
|
||||
i2v_model_path: str = "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers",
|
||||
output_path: str = "outputs",
|
||||
host: str = "0.0.0.0",
|
||||
port: int = 8000,
|
||||
num_gpus: int = 8,
|
||||
cluster_id: int = 0
|
||||
):
|
||||
"""Start the Ray Serve backend with multiple GPU replicas"""
|
||||
# Initialize Ray
|
||||
if not ray.is_initialized():
|
||||
ray.init()
|
||||
|
||||
# Use unique application name based on cluster_id
|
||||
app_name = f"fast_video_cluster_{cluster_id}"
|
||||
|
||||
# Deploy the API
|
||||
api = FastVideoMultiGPUAPI.bind(t2v_model_path, i2v_model_path, output_path)
|
||||
serve.run(api, route_prefix=f"/cluster_{cluster_id}", name=app_name)
|
||||
|
||||
print(f"Ray Serve multi-GPU backend started at http://{host}:{port}")
|
||||
print(f"T2V Model: {t2v_model_path}")
|
||||
print(f"I2V Model: {i2v_model_path}")
|
||||
print(f"Number of GPU replicas: {num_gpus}")
|
||||
print(f"Cluster ID: {cluster_id}")
|
||||
print(f"Application name: {app_name}")
|
||||
print(f"Health check: http://{host}:{port}/cluster_{cluster_id}/health")
|
||||
print(f"Video generation endpoint: http://{host}:{port}/cluster_{cluster_id}/generate_video")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import argparse
|
||||
|
||||
parser = argparse.ArgumentParser(description="FastVideo Ray Serve Multi-GPU Backend")
|
||||
parser.add_argument("--t2v_model_path",
|
||||
type=str,
|
||||
default="FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
|
||||
help="Path to the T2V model")
|
||||
parser.add_argument("--i2v_model_path",
|
||||
type=str,
|
||||
default="Wan-AI/Wan2.1-I2V-14B-480P-Diffusers",
|
||||
help="Path to the I2V model")
|
||||
parser.add_argument("--output_path",
|
||||
type=str,
|
||||
default="outputs",
|
||||
help="Path to save generated videos")
|
||||
parser.add_argument("--host",
|
||||
type=str,
|
||||
default="0.0.0.0",
|
||||
help="Host to bind to")
|
||||
parser.add_argument("--port",
|
||||
type=int,
|
||||
default=8000,
|
||||
help="Port to bind to")
|
||||
parser.add_argument("--num_gpus",
|
||||
type=int,
|
||||
default=8,
|
||||
help="Number of GPU replicas")
|
||||
parser.add_argument("--cluster_id",
|
||||
type=int,
|
||||
default=0,
|
||||
help="Cluster ID for unique naming")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
start_ray_serve_multi_gpu(
|
||||
t2v_model_path=args.t2v_model_path,
|
||||
i2v_model_path=args.i2v_model_path,
|
||||
output_path=args.output_path,
|
||||
host=args.host,
|
||||
port=args.port,
|
||||
num_gpus=args.num_gpus,
|
||||
cluster_id=args.cluster_id,
|
||||
)
|
||||
|
||||
# Keep the process alive
|
||||
import signal, sys, time
|
||||
signal.signal(signal.SIGINT, lambda *_: sys.exit(0))
|
||||
signal.signal(signal.SIGTERM, lambda *_: sys.exit(0))
|
||||
|
||||
print("✅ FastVideo multi-GPU backend is running. Press Ctrl-C to stop.")
|
||||
while True:
|
||||
time.sleep(3600)
|
||||
@@ -0,0 +1,232 @@
|
||||
"""
|
||||
Startup script for FastVideo with Ray Serve backend and Gradio frontend.
|
||||
This script starts both the backend and frontend services.
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
import time
|
||||
import threading
|
||||
import signal
|
||||
import requests
|
||||
from pathlib import Path
|
||||
|
||||
# Add the project root to the Python path
|
||||
project_root = Path(__file__).parent.parent.parent.parent
|
||||
sys.path.insert(0, str(project_root))
|
||||
|
||||
|
||||
def check_backend_health(backend_url: str, max_retries: int = 100) -> bool:
|
||||
"""Check if the backend is healthy"""
|
||||
health_url = f"{backend_url}/health"
|
||||
|
||||
for i in range(max_retries):
|
||||
try:
|
||||
response = requests.get(health_url, timeout=5)
|
||||
if response.status_code == 200:
|
||||
print(f"✅ Backend is healthy at {backend_url}")
|
||||
return True
|
||||
except requests.exceptions.RequestException:
|
||||
pass
|
||||
|
||||
if i < max_retries - 1:
|
||||
print(f"⏳ Waiting for backend to start... ({i+1}/{max_retries})")
|
||||
time.sleep(2)
|
||||
|
||||
print(f"❌ Backend failed to start within {max_retries * 2} seconds")
|
||||
return False
|
||||
|
||||
|
||||
def start_backend(args):
|
||||
"""Start the Ray Serve backend"""
|
||||
backend_script = Path(__file__).parent / "ray_serve_backend.py"
|
||||
|
||||
cmd = [
|
||||
sys.executable, str(backend_script),
|
||||
"--t2v_model_path", args.t2v_model_path,
|
||||
"--i2v_model_path", args.i2v_model_path,
|
||||
"--output_path", args.output_path,
|
||||
"--host", args.backend_host,
|
||||
"--port", str(args.backend_port)
|
||||
]
|
||||
|
||||
print(f"🚀 Starting Ray Serve backend...")
|
||||
print(f"Command: {' '.join(cmd)}")
|
||||
|
||||
# Start the backend process
|
||||
backend_process = subprocess.Popen(
|
||||
cmd,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.STDOUT,
|
||||
universal_newlines=True,
|
||||
bufsize=1
|
||||
)
|
||||
|
||||
# Monitor backend output
|
||||
def monitor_backend():
|
||||
for line in backend_process.stdout:
|
||||
print(f"[BACKEND] {line.rstrip()}")
|
||||
|
||||
monitor_thread = threading.Thread(target=monitor_backend, daemon=True)
|
||||
monitor_thread.start()
|
||||
|
||||
return backend_process
|
||||
|
||||
|
||||
def start_frontend(args):
|
||||
"""Start the Gradio frontend"""
|
||||
frontend_script = Path(__file__).parent / "gradio_frontend.py"
|
||||
backend_url = f"http://{args.backend_host}:{args.backend_port}"
|
||||
|
||||
cmd = [
|
||||
sys.executable, str(frontend_script),
|
||||
"--backend_url", backend_url,
|
||||
"--t2v_model_path", args.t2v_model_path,
|
||||
"--i2v_model_path", args.i2v_model_path,
|
||||
"--host", args.frontend_host,
|
||||
"--port", str(args.frontend_port)
|
||||
]
|
||||
|
||||
print(f"🎨 Starting Gradio frontend...")
|
||||
print(f"Command: {' '.join(cmd)}")
|
||||
|
||||
# Start the frontend process
|
||||
frontend_process = subprocess.Popen(
|
||||
cmd,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.STDOUT,
|
||||
universal_newlines=True,
|
||||
bufsize=1
|
||||
)
|
||||
|
||||
# Monitor frontend output
|
||||
def monitor_frontend():
|
||||
for line in frontend_process.stdout:
|
||||
print(f"[FRONTEND] {line.rstrip()}")
|
||||
|
||||
monitor_thread = threading.Thread(target=monitor_frontend, daemon=True)
|
||||
monitor_thread.start()
|
||||
|
||||
return frontend_process
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="FastVideo Ray Serve App")
|
||||
|
||||
# Model and output settings
|
||||
parser.add_argument("--t2v_model_path",
|
||||
type=str,
|
||||
default="Wan-AI/Wan2.2-TI2V-5BDiffusers",
|
||||
help="Path to the T2V model")
|
||||
parser.add_argument("--i2v_model_path",
|
||||
type=str,
|
||||
default="Wan-AI/Wan2.2-TI2V-5B-Diffusers",
|
||||
help="Path to the I2V model")
|
||||
parser.add_argument("--output_path",
|
||||
type=str,
|
||||
default="outputs",
|
||||
help="Path to save generated videos")
|
||||
|
||||
# Backend settings
|
||||
parser.add_argument("--backend_host",
|
||||
type=str,
|
||||
default="0.0.0.0",
|
||||
help="Backend host to bind to")
|
||||
parser.add_argument("--backend_port",
|
||||
type=int,
|
||||
default=8000,
|
||||
help="Backend port to bind to")
|
||||
|
||||
# Frontend settings
|
||||
parser.add_argument("--frontend_host",
|
||||
type=str,
|
||||
default="0.0.0.0",
|
||||
help="Frontend host to bind to")
|
||||
parser.add_argument("--frontend_port",
|
||||
type=int,
|
||||
default=7860,
|
||||
help="Frontend port to bind to")
|
||||
|
||||
# Other settings
|
||||
parser.add_argument("--skip_backend_check",
|
||||
action="store_true",
|
||||
help="Skip backend health check")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
# Ensure output directory exists
|
||||
os.makedirs(args.output_path, exist_ok=True)
|
||||
|
||||
print("🎬 FastVideo Ray Serve App")
|
||||
print("=" * 50)
|
||||
print(f"T2V Model: {args.t2v_model_path}")
|
||||
print(f"I2V Model: {args.i2v_model_path}")
|
||||
print(f"Output: {args.output_path}")
|
||||
print(f"Backend: http://{args.backend_host}:{args.backend_port}")
|
||||
print(f"Frontend: http://{args.frontend_host}:{args.frontend_port}")
|
||||
print("=" * 50)
|
||||
|
||||
# Start backend
|
||||
backend_process = start_backend(args)
|
||||
|
||||
# Wait for backend to be ready
|
||||
backend_url = f"http://{args.backend_host}:{args.backend_port}"
|
||||
|
||||
if not args.skip_backend_check:
|
||||
if not check_backend_health(backend_url):
|
||||
print("❌ Backend failed to start. Terminating...")
|
||||
backend_process.terminate()
|
||||
sys.exit(1)
|
||||
|
||||
# Start frontend
|
||||
frontend_process = start_frontend(args)
|
||||
|
||||
print("\n🎉 Both services are starting up!")
|
||||
print(f"📺 Frontend will be available at: http://{args.frontend_host}:{args.frontend_port}")
|
||||
print(f"🔧 Backend API will be available at: {backend_url}")
|
||||
print("\nPress Ctrl+C to stop both services...")
|
||||
# return
|
||||
|
||||
# Signal handler for graceful shutdown
|
||||
def signal_handler(signum, frame):
|
||||
print("\n🛑 Shutting down services...")
|
||||
frontend_process.terminate()
|
||||
backend_process.terminate()
|
||||
|
||||
# Wait for processes to terminate
|
||||
try:
|
||||
frontend_process.wait(timeout=5)
|
||||
backend_process.wait(timeout=5)
|
||||
except subprocess.TimeoutExpired:
|
||||
print("⚠️ Force killing processes...")
|
||||
frontend_process.kill()
|
||||
backend_process.kill()
|
||||
|
||||
print("✅ Services stopped")
|
||||
sys.exit(0)
|
||||
|
||||
signal.signal(signal.SIGINT, signal_handler)
|
||||
signal.signal(signal.SIGTERM, signal_handler)
|
||||
|
||||
# Monitor processes
|
||||
try:
|
||||
while True:
|
||||
# Check if processes are still running
|
||||
if frontend_process.poll() is not None:
|
||||
print("❌ Frontend process died unexpectedly")
|
||||
break
|
||||
|
||||
if backend_process.poll() is not None:
|
||||
print("❌ Backend process died unexpectedly")
|
||||
break
|
||||
|
||||
time.sleep(1)
|
||||
|
||||
except KeyboardInterrupt:
|
||||
signal_handler(signal.SIGINT, None)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,465 @@
|
||||
"""
|
||||
Startup script for FastVideo scalable architecture:
|
||||
ngrok -> nginx reverse proxy -> frontend1/frontend2 -> backend1×8/backend2×8
|
||||
|
||||
This script starts:
|
||||
1. Multiple backend instances (8 GPU replicas each)
|
||||
2. Multiple frontend instances (2 instances)
|
||||
3. Nginx reverse proxy
|
||||
4. Optional ngrok tunnel
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
import time
|
||||
import threading
|
||||
import signal
|
||||
import requests
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
# Add the project root to the Python path
|
||||
project_root = Path(__file__).parent.parent.parent.parent
|
||||
sys.path.insert(0, str(project_root))
|
||||
|
||||
|
||||
def check_service_health(url: str, max_retries: int = 50) -> bool:
|
||||
"""Check if a service is healthy"""
|
||||
for i in range(max_retries):
|
||||
try:
|
||||
response = requests.get(url, timeout=5)
|
||||
if response.status_code == 200:
|
||||
print(f"✅ Service is healthy at {url}")
|
||||
return True
|
||||
except requests.exceptions.RequestException:
|
||||
pass
|
||||
|
||||
if i < max_retries - 1:
|
||||
print(f"⏳ Waiting for service to start... ({i+1}/{max_retries})")
|
||||
time.sleep(2)
|
||||
|
||||
print(f"❌ Service failed to start within {max_retries * 2} seconds")
|
||||
return False
|
||||
|
||||
|
||||
def start_backend_cluster(args, cluster_id: int):
|
||||
"""Start one backend cluster (Ray-Serve application)."""
|
||||
backend_script = Path(__file__).parent / "ray_serve_backend_scalable.py"
|
||||
|
||||
# All Ray Serve apps share the same HTTP server (default 8000).
|
||||
# We still forward the port flag for completeness, but keep it
|
||||
# identical for every cluster.
|
||||
base_port = args.backend_base_port
|
||||
|
||||
cmd = [
|
||||
sys.executable, str(backend_script),
|
||||
"--t2v_model_path", args.t2v_model_path,
|
||||
"--i2v_model_path", args.i2v_model_path,
|
||||
"--output_path", args.output_path,
|
||||
"--host", args.backend_host,
|
||||
"--port", str(base_port),
|
||||
"--num_gpus", str(args.num_gpus_per_cluster),
|
||||
"--cluster_id", str(cluster_id),
|
||||
]
|
||||
|
||||
print(f"🚀 Starting Backend Cluster {cluster_id + 1} (HTTP port {base_port})...")
|
||||
print(f"Command: {' '.join(cmd)}")
|
||||
|
||||
# Start the backend process
|
||||
backend_process = subprocess.Popen(
|
||||
cmd,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.STDOUT,
|
||||
universal_newlines=True,
|
||||
bufsize=1
|
||||
)
|
||||
|
||||
# Monitor backend output
|
||||
def monitor_backend():
|
||||
for line in backend_process.stdout:
|
||||
print(f"[BACKEND-{cluster_id + 1}] {line.rstrip()}")
|
||||
|
||||
monitor_thread = threading.Thread(target=monitor_backend, daemon=True)
|
||||
monitor_thread.start()
|
||||
|
||||
return backend_process, base_port
|
||||
|
||||
|
||||
def start_frontend_instance(args, instance_id: int, backend_url: str):
|
||||
"""Start a single frontend instance"""
|
||||
frontend_script = Path(__file__).parent / "gradio_frontend.py"
|
||||
frontend_port = args.frontend_base_port + instance_id
|
||||
|
||||
# Update backend URL to include cluster-specific path
|
||||
cluster_id = instance_id % args.num_backend_clusters
|
||||
backend_url_with_cluster = f"{backend_url}/cluster_{cluster_id}"
|
||||
|
||||
cmd = [
|
||||
sys.executable, str(frontend_script),
|
||||
"--backend_url", backend_url_with_cluster,
|
||||
"--t2v_model_path", args.t2v_model_path,
|
||||
"--i2v_model_path", args.i2v_model_path,
|
||||
"--host", args.frontend_host,
|
||||
"--port", str(frontend_port)
|
||||
]
|
||||
|
||||
print(f"🎨 Starting Frontend {instance_id + 1} on port {frontend_port}...")
|
||||
print(f"Command: {' '.join(cmd)}")
|
||||
|
||||
# Start the frontend process
|
||||
frontend_process = subprocess.Popen(
|
||||
cmd,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.STDOUT,
|
||||
universal_newlines=True,
|
||||
bufsize=1
|
||||
)
|
||||
|
||||
# Monitor frontend output
|
||||
def monitor_frontend():
|
||||
for line in frontend_process.stdout:
|
||||
print(f"[FRONTEND-{instance_id + 1}] {line.rstrip()}")
|
||||
|
||||
monitor_thread = threading.Thread(target=monitor_frontend, daemon=True)
|
||||
monitor_thread.start()
|
||||
|
||||
return frontend_process, frontend_port
|
||||
|
||||
|
||||
def start_nginx(args):
|
||||
"""Start nginx reverse proxy"""
|
||||
nginx_conf = Path(__file__).parent / "nginx.conf"
|
||||
|
||||
# Update nginx configuration with actual ports
|
||||
update_nginx_config(args)
|
||||
|
||||
cmd = [
|
||||
"nginx",
|
||||
"-c", str(nginx_conf),
|
||||
"-g", "daemon off;"
|
||||
]
|
||||
|
||||
print(f"🌐 Starting Nginx reverse proxy...")
|
||||
print(f"Command: {' '.join(cmd)}")
|
||||
|
||||
# Start nginx process
|
||||
nginx_process = subprocess.Popen(
|
||||
cmd,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.STDOUT,
|
||||
universal_newlines=True,
|
||||
bufsize=1
|
||||
)
|
||||
|
||||
# Monitor nginx output
|
||||
def monitor_nginx():
|
||||
for line in nginx_process.stdout:
|
||||
print(f"[NGINX] {line.rstrip()}")
|
||||
|
||||
monitor_thread = threading.Thread(target=monitor_nginx, daemon=True)
|
||||
monitor_thread.start()
|
||||
|
||||
return nginx_process
|
||||
|
||||
|
||||
def update_nginx_config(args):
|
||||
"""Rewrite nginx.conf with the correct ports – NO “/cluster_X” in upstreams."""
|
||||
nginx_conf = Path(__file__).parent / "nginx.conf"
|
||||
nginx_conf_backup = Path(__file__).parent / "nginx.conf.backup"
|
||||
|
||||
if not nginx_conf_backup.exists():
|
||||
nginx_conf_backup.write_text(nginx_conf.read_text())
|
||||
|
||||
config_content = nginx_conf_backup.read_text()
|
||||
|
||||
# ── 1. front-end pool ───────────────────────────────────────────────
|
||||
frontend_servers = "\n ".join(
|
||||
f"server 127.0.0.1:{args.frontend_base_port + i} "
|
||||
f"weight=1 max_fails=3 fail_timeout=30s;"
|
||||
for i in range(args.num_frontends)
|
||||
)
|
||||
config_content = config_content.replace(
|
||||
"# Upstream for frontend load balancing",
|
||||
f"# Upstream for frontend load balancing\n upstream frontend_servers {{\n"
|
||||
f" # Round-robin load balancing between frontends\n {frontend_servers}"
|
||||
)
|
||||
|
||||
# Shared Ray-Serve HTTP port
|
||||
backend_port = args.backend_base_port # default 8000
|
||||
backend_line = (f"server 127.0.0.1:{backend_port} "
|
||||
f"weight=1 max_fails=3 fail_timeout=30s;")
|
||||
|
||||
# ── 2. backend-1 pool ───────────────────────────────────────────────
|
||||
config_content = config_content.replace(
|
||||
"# Upstream for backend1 load balancing",
|
||||
f"# Upstream for backend1 load balancing\n upstream backend1_servers {{\n"
|
||||
f" {backend_line}"
|
||||
)
|
||||
|
||||
# ── 3. backend-2 pool ───────────────────────────────────────────────
|
||||
config_content = config_content.replace(
|
||||
"# Upstream for backend2 load balancing",
|
||||
f"# Upstream for backend2 load balancing\n upstream backend2_servers {{\n"
|
||||
f" {backend_line}"
|
||||
)
|
||||
|
||||
# ── 4. strip any stray “/cluster_X” fragments ───────────────────────
|
||||
config_content = config_content.replace("/cluster_0", "").replace("/cluster_1", "")
|
||||
|
||||
# ── 5. use user-writable log directory ---------------------------------
|
||||
log_dir = Path(args.output_path).resolve()
|
||||
config_content = config_content.replace(
|
||||
"access_log /var/log/nginx/access.log;",
|
||||
f"access_log {log_dir}/nginx_access.log;")
|
||||
config_content = config_content.replace(
|
||||
"error_log /var/log/nginx/error.log;",
|
||||
f"error_log {log_dir}/nginx_error.log;")
|
||||
|
||||
nginx_conf.write_text(config_content)
|
||||
print("✅ nginx.conf updated (no path suffixes & custom log paths)")
|
||||
|
||||
|
||||
def start_ngrok(args):
|
||||
"""Start ngrok tunnel"""
|
||||
if not args.use_ngrok:
|
||||
return None
|
||||
|
||||
cmd = [
|
||||
"ngrok",
|
||||
"http",
|
||||
str(args.nginx_port),
|
||||
"--log=stdout"
|
||||
]
|
||||
|
||||
print(f"🌍 Starting ngrok tunnel to port {args.nginx_port}...")
|
||||
print(f"Command: {' '.join(cmd)}")
|
||||
|
||||
# Start ngrok process
|
||||
ngrok_process = subprocess.Popen(
|
||||
cmd,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.STDOUT,
|
||||
universal_newlines=True,
|
||||
bufsize=1
|
||||
)
|
||||
|
||||
# Monitor ngrok output
|
||||
def monitor_ngrok():
|
||||
for line in ngrok_process.stdout:
|
||||
print(f"[NGROK] {line.rstrip()}")
|
||||
|
||||
monitor_thread = threading.Thread(target=monitor_ngrok, daemon=True)
|
||||
monitor_thread.start()
|
||||
|
||||
return ngrok_process
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="FastVideo Scalable Architecture Launcher")
|
||||
|
||||
# Model and output settings
|
||||
parser.add_argument("--t2v_model_path",
|
||||
type=str,
|
||||
default="FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
|
||||
help="Path to the T2V model")
|
||||
parser.add_argument("--i2v_model_path",
|
||||
type=str,
|
||||
default="Wan-AI/Wan2.1-I2V-14B-480P-Diffusers",
|
||||
help="Path to the I2V model")
|
||||
parser.add_argument("--output_path",
|
||||
type=str,
|
||||
default="outputs",
|
||||
help="Path to save generated videos")
|
||||
|
||||
# Backend settings
|
||||
parser.add_argument("--backend_host",
|
||||
type=str,
|
||||
default="0.0.0.0",
|
||||
help="Backend host to bind to")
|
||||
parser.add_argument("--backend_base_port",
|
||||
type=int,
|
||||
default=8000,
|
||||
help="Base port for backend clusters")
|
||||
parser.add_argument("--num_backend_clusters",
|
||||
type=int,
|
||||
default=2,
|
||||
help="Number of backend clusters")
|
||||
parser.add_argument("--num_gpus_per_cluster",
|
||||
type=int,
|
||||
default=3, # Changed from 8 to 3 (3+3=6 GPUs total, leaving 1 GPU buffer)
|
||||
help="Number of GPUs per backend cluster")
|
||||
|
||||
# Frontend settings
|
||||
parser.add_argument("--frontend_host",
|
||||
type=str,
|
||||
default="0.0.0.0",
|
||||
help="Frontend host to bind to")
|
||||
parser.add_argument("--frontend_base_port",
|
||||
type=int,
|
||||
default=7860,
|
||||
help="Base port for frontend instances")
|
||||
parser.add_argument("--num_frontends",
|
||||
type=int,
|
||||
default=2,
|
||||
help="Number of frontend instances")
|
||||
|
||||
# Nginx settings
|
||||
parser.add_argument("--nginx_port",
|
||||
type=int,
|
||||
default=80,
|
||||
help="Port for nginx reverse proxy")
|
||||
|
||||
# Ngrok settings
|
||||
parser.add_argument("--use_ngrok",
|
||||
action="store_true",
|
||||
help="Start ngrok tunnel")
|
||||
|
||||
# Other settings
|
||||
parser.add_argument("--skip_health_check",
|
||||
action="store_true",
|
||||
help="Skip health checks")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
# Ensure output directory exists
|
||||
os.makedirs(args.output_path, exist_ok=True)
|
||||
|
||||
print(" FastVideo Scalable Architecture")
|
||||
print("=" * 60)
|
||||
print(f"Architecture: ngrok -> nginx -> frontend1/frontend2 -> backend1×{args.num_gpus_per_cluster}/backend2×{args.num_gpus_per_cluster}")
|
||||
print(f"T2V Model: {args.t2v_model_path}")
|
||||
print(f"I2V Model: {args.i2v_model_path}")
|
||||
print(f"Output: {args.output_path}")
|
||||
print(f"Backend Clusters: {args.num_backend_clusters}")
|
||||
print(f"GPUs per Cluster: {args.num_gpus_per_cluster}")
|
||||
print(f"Total GPUs needed: {args.num_backend_clusters * args.num_gpus_per_cluster}")
|
||||
print(f"Frontend Instances: {args.num_frontends}")
|
||||
print(f"Nginx Port: {args.nginx_port}")
|
||||
print(f"Use Ngrok: {args.use_ngrok}")
|
||||
print("=" * 60)
|
||||
|
||||
# Start backend clusters
|
||||
backend_processes = []
|
||||
backend_urls = []
|
||||
|
||||
for i in range(args.num_backend_clusters):
|
||||
process, _ = start_backend_cluster(args, i)
|
||||
backend_processes.append(process)
|
||||
backend_urls.append(f"http://{args.backend_host}:{args.backend_base_port}")
|
||||
|
||||
# Wait for backends to be ready
|
||||
if not args.skip_health_check:
|
||||
print("\n⏳ Waiting for backend clusters to start...")
|
||||
for i, url in enumerate(backend_urls):
|
||||
if not check_service_health(f"{url}/cluster_{i}/health"):
|
||||
print(f"❌ Backend cluster {i + 1} failed to start. Terminating...")
|
||||
for process in backend_processes:
|
||||
process.terminate()
|
||||
sys.exit(1)
|
||||
|
||||
# Start frontend instances
|
||||
frontend_processes = []
|
||||
frontend_urls = []
|
||||
|
||||
for i in range(args.num_frontends):
|
||||
# Each frontend connects to a different backend cluster
|
||||
backend_url = backend_urls[i % len(backend_urls)]
|
||||
process, port = start_frontend_instance(args, i, backend_url)
|
||||
frontend_processes.append(process)
|
||||
frontend_urls.append(f"http://{args.frontend_host}:{port}")
|
||||
|
||||
# Wait for frontends to be ready
|
||||
if not args.skip_health_check:
|
||||
print("\n⏳ Waiting for frontend instances to start...")
|
||||
for i, url in enumerate(frontend_urls):
|
||||
if not check_service_health(url):
|
||||
print(f"❌ Frontend {i + 1} failed to start. Terminating...")
|
||||
for process in backend_processes + frontend_processes:
|
||||
process.terminate()
|
||||
sys.exit(1)
|
||||
|
||||
# Start nginx reverse proxy
|
||||
nginx_process = start_nginx(args)
|
||||
|
||||
# Wait for nginx to be ready
|
||||
if not args.skip_health_check:
|
||||
print("\n⏳ Waiting for nginx to start...")
|
||||
if not check_service_health(f"http://localhost:{args.nginx_port}/health"):
|
||||
print("❌ Nginx failed to start. Terminating...")
|
||||
for process in backend_processes + frontend_processes + [nginx_process]:
|
||||
process.terminate()
|
||||
sys.exit(1)
|
||||
|
||||
# Start ngrok tunnel (optional)
|
||||
ngrok_process = start_ngrok(args)
|
||||
|
||||
print("\n🎉 All services are starting up!")
|
||||
print(f"🌐 Nginx reverse proxy: http://localhost:{args.nginx_port}")
|
||||
for i, url in enumerate(frontend_urls):
|
||||
print(f"📺 Frontend {i + 1}: {url}")
|
||||
for i, url in enumerate(backend_urls):
|
||||
print(f" Backend Cluster {i + 1}: {url}")
|
||||
if args.use_ngrok:
|
||||
print("🌍 Ngrok tunnel is starting...")
|
||||
print("\nPress Ctrl+C to stop all services...")
|
||||
|
||||
# Signal handler for graceful shutdown
|
||||
def signal_handler(signum, frame):
|
||||
print("\n🛑 Shutting down all services...")
|
||||
all_processes = backend_processes + frontend_processes + [nginx_process]
|
||||
if ngrok_process:
|
||||
all_processes.append(ngrok_process)
|
||||
|
||||
for process in all_processes:
|
||||
if process:
|
||||
process.terminate()
|
||||
|
||||
# Wait for processes to terminate
|
||||
try:
|
||||
for process in all_processes:
|
||||
if process:
|
||||
process.wait(timeout=5)
|
||||
except subprocess.TimeoutExpired:
|
||||
print("⚠️ Force killing processes...")
|
||||
for process in all_processes:
|
||||
if process:
|
||||
process.kill()
|
||||
|
||||
print("✅ All services stopped")
|
||||
sys.exit(0)
|
||||
|
||||
signal.signal(signal.SIGINT, signal_handler)
|
||||
signal.signal(signal.SIGTERM, signal_handler)
|
||||
|
||||
# Monitor processes
|
||||
try:
|
||||
while True:
|
||||
# Check if processes are still running
|
||||
for i, process in enumerate(backend_processes):
|
||||
if process.poll() is not None:
|
||||
print(f"❌ Backend cluster {i + 1} process died unexpectedly")
|
||||
break
|
||||
|
||||
for i, process in enumerate(frontend_processes):
|
||||
if process.poll() is not None:
|
||||
print(f"❌ Frontend {i + 1} process died unexpectedly")
|
||||
break
|
||||
|
||||
if nginx_process and nginx_process.poll() is not None:
|
||||
print("❌ Nginx process died unexpectedly")
|
||||
break
|
||||
|
||||
if ngrok_process and ngrok_process.poll() is not None:
|
||||
print("❌ Ngrok process died unexpectedly")
|
||||
break
|
||||
|
||||
time.sleep(1)
|
||||
|
||||
except KeyboardInterrupt:
|
||||
signal_handler(signal.SIGINT, None)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,147 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Test script for T2V and I2V functionality in FastVideo Gradio app.
|
||||
This script tests both the backend and frontend modifications.
|
||||
"""
|
||||
|
||||
import requests
|
||||
import json
|
||||
import os
|
||||
from PIL import Image
|
||||
import numpy as np
|
||||
|
||||
def test_backend_t2v():
|
||||
"""Test the backend T2V functionality directly"""
|
||||
backend_url = "http://localhost:8000"
|
||||
|
||||
try:
|
||||
# Test T2V request data
|
||||
request_data = {
|
||||
"prompt": "A beautiful sunset over the ocean with gentle waves",
|
||||
"negative_prompt": "",
|
||||
"use_negative_prompt": False,
|
||||
"seed": 42,
|
||||
"guidance_scale": 7.5,
|
||||
"num_frames": 21,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_inference_steps": 20,
|
||||
"randomize_seed": False,
|
||||
"return_frames": True,
|
||||
"image_path": None,
|
||||
"model_type": "t2v"
|
||||
}
|
||||
|
||||
# Send request to backend
|
||||
response = requests.post(
|
||||
f"{backend_url}/generate_video",
|
||||
json=request_data,
|
||||
timeout=300 # 5 minutes timeout
|
||||
)
|
||||
|
||||
if response.status_code == 200:
|
||||
result = response.json()
|
||||
print("✅ Backend T2V test successful!")
|
||||
print(f"Success: {result.get('success')}")
|
||||
print(f"Seed used: {result.get('seed')}")
|
||||
if result.get('frames'):
|
||||
print(f"Frames returned: {len(result.get('frames'))}")
|
||||
else:
|
||||
print("No frames returned")
|
||||
else:
|
||||
print(f"❌ Backend T2V test failed with status {response.status_code}")
|
||||
print(f"Response: {response.text}")
|
||||
|
||||
except Exception as e:
|
||||
print(f"❌ Backend T2V test failed with exception: {e}")
|
||||
|
||||
def test_backend_i2v():
|
||||
"""Test the backend I2V functionality directly"""
|
||||
backend_url = "http://localhost:8000"
|
||||
|
||||
# Create a simple test image
|
||||
test_image = Image.new('RGB', (256, 256), color='red')
|
||||
temp_image_path = "test_image.png"
|
||||
test_image.save(temp_image_path)
|
||||
|
||||
try:
|
||||
# Test I2V request data
|
||||
request_data = {
|
||||
"prompt": "The red square gently animates with subtle movement",
|
||||
"negative_prompt": "",
|
||||
"use_negative_prompt": False,
|
||||
"seed": 42,
|
||||
"guidance_scale": 7.5,
|
||||
"num_frames": 21,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_inference_steps": 20,
|
||||
"randomize_seed": False,
|
||||
"return_frames": True,
|
||||
"image_path": temp_image_path,
|
||||
"model_type": "i2v"
|
||||
}
|
||||
|
||||
# Send request to backend
|
||||
response = requests.post(
|
||||
f"{backend_url}/generate_video",
|
||||
json=request_data,
|
||||
timeout=300 # 5 minutes timeout
|
||||
)
|
||||
|
||||
if response.status_code == 200:
|
||||
result = response.json()
|
||||
print("✅ Backend I2V test successful!")
|
||||
print(f"Success: {result.get('success')}")
|
||||
print(f"Seed used: {result.get('seed')}")
|
||||
if result.get('frames'):
|
||||
print(f"Frames returned: {len(result.get('frames'))}")
|
||||
else:
|
||||
print("No frames returned")
|
||||
else:
|
||||
print(f"❌ Backend I2V test failed with status {response.status_code}")
|
||||
print(f"Response: {response.text}")
|
||||
|
||||
except Exception as e:
|
||||
print(f"❌ Backend I2V test failed with exception: {e}")
|
||||
|
||||
finally:
|
||||
# Clean up test image
|
||||
if os.path.exists(temp_image_path):
|
||||
os.remove(temp_image_path)
|
||||
|
||||
def test_backend_health():
|
||||
"""Test if the backend is running"""
|
||||
backend_url = "http://localhost:8000"
|
||||
|
||||
try:
|
||||
response = requests.get(f"{backend_url}/health", timeout=5)
|
||||
if response.status_code == 200:
|
||||
print("✅ Backend is healthy")
|
||||
return True
|
||||
else:
|
||||
print(f"❌ Backend health check failed: {response.status_code}")
|
||||
return False
|
||||
except Exception as e:
|
||||
print(f"❌ Backend health check failed: {e}")
|
||||
return False
|
||||
|
||||
if __name__ == "__main__":
|
||||
print("🧪 Testing FastVideo T2V and I2V functionality...")
|
||||
print("=" * 50)
|
||||
|
||||
# Test backend health first
|
||||
if test_backend_health():
|
||||
# Test T2V functionality
|
||||
print("\n📝 Testing T2V functionality...")
|
||||
test_backend_t2v()
|
||||
|
||||
# Test I2V functionality
|
||||
print("\n🖼️ Testing I2V functionality...")
|
||||
test_backend_i2v()
|
||||
else:
|
||||
print("⚠️ Backend is not running. Please start the backend first.")
|
||||
print("You can start it with: python start_ray_serve_app.py")
|
||||
|
||||
print("=" * 50)
|
||||
print("Test completed!")
|
||||
@@ -85,6 +85,9 @@ class PipelineConfig:
|
||||
# DMD parameters
|
||||
dmd_denoising_steps: list[int] | None = field(default=None)
|
||||
|
||||
# Wan2.2 TI2V parameters
|
||||
ti2v_task: bool = False
|
||||
|
||||
# Compilation
|
||||
# enable_torch_compile: bool = False
|
||||
|
||||
|
||||
@@ -8,6 +8,7 @@ from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
|
||||
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
|
||||
from fastvideo.configs.pipelines.wan import (FastWanT2V480PConfig,
|
||||
Wan2_2_TI2V_5B_Config,
|
||||
WanI2V480PConfig, WanI2V720PConfig,
|
||||
WanT2V480PConfig, WanT2V720PConfig)
|
||||
from fastvideo.logger import init_logger
|
||||
@@ -26,9 +27,12 @@ PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
|
||||
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers": WanI2V720PConfig,
|
||||
"Wan-AI/Wan2.1-T2V-14B-Diffusers": WanT2V720PConfig,
|
||||
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers": FastWanT2V480PConfig,
|
||||
"FastVideo/FastWan2.1-T2V-14B-480P-Diffusers": FastWanT2V480PConfig,
|
||||
"FastVideo/stepvideo-t2v-diffusers": StepVideoT2VConfig,
|
||||
"FastVideo/Wan2.1-VSA-T2V-14B-720P-Diffusers": WanT2V720PConfig,
|
||||
"Wan-AI/Wan2.2-TI2V-5B-Diffusers": WanT2V720PConfig
|
||||
"Wan-AI/Wan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_Config,
|
||||
# "Wan-AI/Wan2.2-T2V-A14B-Diffusers": Wan2_2_T2V_A14B_Config,
|
||||
# "Wan-AI/Wan2.2-I2V-A14B-Diffusers": Wan2_2_I2V_A14B_Config,
|
||||
# Add other specific weight variants
|
||||
}
|
||||
|
||||
|
||||
@@ -111,3 +111,26 @@ class FastWanT2V480PConfig(WanT2V480PConfig):
|
||||
def __post_init__(self) -> None:
|
||||
self.vae_config.load_encoder = True
|
||||
self.vae_config.load_decoder = True
|
||||
|
||||
|
||||
@dataclass
|
||||
class Wan2_2_TI2V_5B_Config(WanT2V480PConfig):
|
||||
"""Base configuration for FastWan T2V 1.3B 480P pipeline architecture with DMD"""
|
||||
|
||||
# Denoising stage
|
||||
flow_shift: int = 5
|
||||
ti2v_task: bool = True
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.vae_config.load_encoder = True
|
||||
self.vae_config.load_decoder = True
|
||||
|
||||
|
||||
@dataclass
|
||||
class Wan2_2_T2V_A14B_Config(WanT2V480PConfig):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass
|
||||
class Wan2_2_I2V_A14B_Config(WanT2V480PConfig):
|
||||
pass
|
||||
|
||||
@@ -6,7 +6,7 @@ from typing import Any
|
||||
from fastvideo.configs.sample.hunyuan import (FastHunyuanSamplingParam,
|
||||
HunyuanSamplingParam)
|
||||
from fastvideo.configs.sample.stepvideo import StepVideoT2VSamplingParam
|
||||
from fastvideo.configs.sample.wan import (FastWanT2V480PConfig,
|
||||
from fastvideo.configs.sample.wan import (Wan2_1_Fun_1_3B_InP_SamplingParam,
|
||||
Wan2_2_TI2V_5B_SamplingParam,
|
||||
WanI2V_14B_480P_SamplingParam,
|
||||
WanI2V_14B_720P_SamplingParam,
|
||||
@@ -25,9 +25,18 @@ SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
|
||||
"Wan-AI/Wan2.1-T2V-14B-Diffusers": WanT2V_14B_SamplingParam,
|
||||
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers": WanI2V_14B_480P_SamplingParam,
|
||||
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers": WanI2V_14B_720P_SamplingParam,
|
||||
"weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers":
|
||||
Wan2_1_Fun_1_3B_InP_SamplingParam,
|
||||
"FastVideo/stepvideo-t2v-diffusers": StepVideoT2VSamplingParam,
|
||||
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers": FastWanT2V480PConfig,
|
||||
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers":
|
||||
Wan2_1_Fun_1_3B_InP_SamplingParam,
|
||||
"FastVideo/FastWan2.1-T2V-14B-Diffusers":
|
||||
Wan2_1_Fun_1_3B_InP_SamplingParam,
|
||||
"Wan-AI/Wan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_SamplingParam,
|
||||
# "Wan-AI/Wan2.2-T2V-A14B-Diffusers":
|
||||
# Wan2_2_T2V_A14B_SamplingParam,
|
||||
# "Wan-AI/Wan2.2-I2V-A14B-Diffusers":
|
||||
# Wan2_2_I2V_A14B_SamplingParam,
|
||||
# Add other specific weight variants
|
||||
}
|
||||
|
||||
|
||||
@@ -107,6 +107,21 @@ class FastWanT2V480PConfig(WanT2V_1_3B_SamplingParam):
|
||||
fps: int = 16
|
||||
|
||||
|
||||
# =============================================
|
||||
# ============= Wan2.1 Fun Models =============
|
||||
# =============================================
|
||||
@dataclass
|
||||
class Wan2_1_Fun_1_3B_InP_SamplingParam(SamplingParam):
|
||||
"""Sampling parameters for Wan2.1 Fun 1.3B InP model."""
|
||||
height: int = 480
|
||||
width: int = 832
|
||||
num_frames: int = 81
|
||||
fps: int = 16
|
||||
negative_prompt: str | None = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
|
||||
guidance_scale: float = 6.0
|
||||
num_inference_steps: int = 50
|
||||
|
||||
|
||||
# =============================================
|
||||
# ============= Wan2.2 TI2V Models =============
|
||||
# =============================================
|
||||
@@ -134,4 +149,4 @@ class Wan2_2_T2V_A14B_SamplingParam(Wan2_2_Base_SamplingParam):
|
||||
|
||||
@dataclass
|
||||
class Wan2_2_I2V_A14B_SamplingParam(Wan2_2_Base_SamplingParam):
|
||||
pass
|
||||
pass
|
||||
@@ -9,6 +9,7 @@ diffusion models.
|
||||
import math
|
||||
import os
|
||||
import time
|
||||
from copy import deepcopy
|
||||
from typing import Any
|
||||
|
||||
import imageio
|
||||
@@ -202,6 +203,8 @@ class VideoGenerator:
|
||||
if sampling_param is None:
|
||||
sampling_param = SamplingParam.from_pretrained(
|
||||
fastvideo_args.model_path)
|
||||
else:
|
||||
sampling_param = deepcopy(sampling_param)
|
||||
|
||||
kwargs["prompt"] = prompt
|
||||
sampling_param.update(kwargs)
|
||||
@@ -275,6 +278,7 @@ class VideoGenerator:
|
||||
width: {target_width}
|
||||
video_length: {sampling_param.num_frames}
|
||||
prompt: {prompt}
|
||||
image_path: {sampling_param.image_path}
|
||||
neg_prompt: {sampling_param.negative_prompt}
|
||||
seed: {sampling_param.seed}
|
||||
infer_steps: {sampling_param.num_inference_steps}
|
||||
@@ -334,6 +338,7 @@ class VideoGenerator:
|
||||
else:
|
||||
return {
|
||||
"samples": samples,
|
||||
"frames": frames,
|
||||
"prompts": prompt,
|
||||
"size": (target_height, target_width, batch.num_frames),
|
||||
"generation_time": gen_time
|
||||
|
||||
@@ -86,12 +86,14 @@ class TimestepEmbedder(nn.Module):
|
||||
dtype=dtype)
|
||||
self.freq_dtype = freq_dtype
|
||||
|
||||
def forward(self, t: torch.Tensor) -> torch.Tensor:
|
||||
def forward(self, t: torch.Tensor, timestep_seq_len: int | None = None) -> torch.Tensor:
|
||||
t_freq = timestep_embedding(t,
|
||||
self.frequency_embedding_size,
|
||||
self.max_period,
|
||||
dtype=self.freq_dtype).to(
|
||||
self.mlp.fc_in.weight.dtype)
|
||||
if timestep_seq_len is not None:
|
||||
t_freq = t_freq.unflatten(0, (1, timestep_seq_len))
|
||||
# t_freq = t_freq.to(self.mlp.fc_in.weight.dtype)
|
||||
t_emb = self.mlp(t_freq)
|
||||
return t_emb
|
||||
|
||||
@@ -81,8 +81,9 @@ class WanTimeTextImageEmbedding(nn.Module):
|
||||
timestep: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
encoder_hidden_states_image: torch.Tensor | None = None,
|
||||
timestep_seq_len: int | None = None,
|
||||
):
|
||||
temb = self.time_embedder(timestep)
|
||||
temb = self.time_embedder(timestep, timestep_seq_len)
|
||||
timestep_proj = self.time_modulation(temb)
|
||||
|
||||
encoder_hidden_states = self.text_embedder(encoder_hidden_states)
|
||||
@@ -307,9 +308,24 @@ class WanTransformerBlock(nn.Module):
|
||||
bs, seq_length, _ = hidden_states.shape
|
||||
orig_dtype = hidden_states.dtype
|
||||
# assert orig_dtype != torch.float32
|
||||
e = self.scale_shift_table + temb.float()
|
||||
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
|
||||
6, dim=1)
|
||||
|
||||
if temb.dim() == 4:
|
||||
# temb: batch_size, seq_len, 6, inner_dim (wan2.2 ti2v)
|
||||
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = (
|
||||
self.scale_shift_table.unsqueeze(0) + temb.float()
|
||||
).chunk(6, dim=2)
|
||||
# batch_size, seq_len, 1, inner_dim
|
||||
shift_msa = shift_msa.squeeze(2)
|
||||
scale_msa = scale_msa.squeeze(2)
|
||||
gate_msa = gate_msa.squeeze(2)
|
||||
c_shift_msa = c_shift_msa.squeeze(2)
|
||||
c_scale_msa = c_scale_msa.squeeze(2)
|
||||
c_gate_msa = c_gate_msa.squeeze(2)
|
||||
else:
|
||||
# temb: batch_size, 6, inner_dim (wan2.1/wan2.2 14B)
|
||||
e = self.scale_shift_table + temb.float()
|
||||
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
|
||||
6, dim=1)
|
||||
assert shift_msa.dtype == torch.float32
|
||||
|
||||
# 1. Self-attention
|
||||
@@ -637,9 +653,21 @@ class WanTransformer3DModel(CachableDiT):
|
||||
hidden_states = self.patch_embedding(hidden_states)
|
||||
hidden_states = hidden_states.flatten(2).transpose(1, 2)
|
||||
|
||||
# timestep shape: batch_size, or batch_size, seq_len (wan 2.2 ti2v)
|
||||
if timestep.dim() == 2:
|
||||
ts_seq_len = timestep.shape[1]
|
||||
timestep = timestep.flatten() # batch_size * seq_len
|
||||
else:
|
||||
ts_seq_len = None
|
||||
|
||||
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
|
||||
timestep, encoder_hidden_states, encoder_hidden_states_image)
|
||||
timestep_proj = timestep_proj.unflatten(1, (6, -1))
|
||||
timestep, encoder_hidden_states, encoder_hidden_states_image, timestep_seq_len=ts_seq_len)
|
||||
if ts_seq_len is not None:
|
||||
# batch_size, seq_len, 6, inner_dim
|
||||
timestep_proj = timestep_proj.unflatten(2, (6, -1))
|
||||
else:
|
||||
# batch_size, 6, inner_dim
|
||||
timestep_proj = timestep_proj.unflatten(1, (6, -1))
|
||||
|
||||
if encoder_hidden_states_image is not None:
|
||||
encoder_hidden_states = torch.concat(
|
||||
@@ -678,8 +706,15 @@ class WanTransformer3DModel(CachableDiT):
|
||||
self.maybe_cache_states(hidden_states, original_hidden_states)
|
||||
|
||||
# 5. Output norm, projection & unpatchify
|
||||
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2,
|
||||
dim=1)
|
||||
if temb.dim() == 3:
|
||||
# batch_size, seq_len, inner_dim (wan 2.2 ti2v)
|
||||
shift, scale = (self.scale_shift_table.unsqueeze(0) + temb.unsqueeze(2)).chunk(2, dim=2)
|
||||
shift = shift.squeeze(2)
|
||||
scale = scale.squeeze(2)
|
||||
else:
|
||||
# batch_size, inner_dim
|
||||
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2, dim=1)
|
||||
|
||||
hidden_states = self.norm_out(hidden_states, shift, scale)
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
|
||||
@@ -782,4 +817,4 @@ class WanTransformer3DModel(CachableDiT):
|
||||
if self.is_even:
|
||||
return hidden_states + self.previous_residual_even
|
||||
else:
|
||||
return hidden_states + self.previous_residual_odd
|
||||
return hidden_states + self.previous_residual_odd
|
||||
@@ -62,6 +62,7 @@ class WanPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
stage=DenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
scheduler=self.get_module("scheduler"),
|
||||
vae=self.get_module("vae"),
|
||||
pipeline=self))
|
||||
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
|
||||
@@ -21,6 +21,7 @@ _PIPELINE_NAME_TO_ARCHITECTURE_NAME: dict[str, str] = {
|
||||
"WanPipeline": "wan",
|
||||
"WanDMDPipeline": "wan",
|
||||
"WanImageToVideoPipeline": "wan",
|
||||
"WanDMDPipeline": "wan",
|
||||
"StepVideoPipeline": "stepvideo",
|
||||
"HunyuanVideoPipeline": "hunyuan",
|
||||
}
|
||||
|
||||
@@ -30,7 +30,7 @@ from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.pipelines.stages.validators import VerificationResult
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
from fastvideo.utils import dict_to_3d_list
|
||||
from fastvideo.utils import dict_to_3d_list, masks_like
|
||||
|
||||
try:
|
||||
from fastvideo.attention.backends.sliding_tile_attn import (
|
||||
@@ -57,10 +57,11 @@ class DenoisingStage(PipelineStage):
|
||||
the initial noise into the final output.
|
||||
"""
|
||||
|
||||
def __init__(self, transformer, scheduler, pipeline=None) -> None:
|
||||
def __init__(self, transformer, scheduler, vae=None, pipeline=None) -> None:
|
||||
super().__init__()
|
||||
self.transformer = transformer
|
||||
self.scheduler = scheduler
|
||||
self.vae = vae
|
||||
self.pipeline = weakref.ref(pipeline) if pipeline else None
|
||||
attn_head_size = self.transformer.hidden_size // self.transformer.num_attention_heads
|
||||
self.attn_backend = get_attn_backend(
|
||||
@@ -184,8 +185,43 @@ class DenoisingStage(PipelineStage):
|
||||
assert neg_prompt_embeds is not None
|
||||
assert torch.isnan(neg_prompt_embeds[0]).sum() == 0
|
||||
|
||||
latent_model_input = latents.to(target_dtype)
|
||||
assert latent_model_input.shape[0] == 1, "only support batch size 1"
|
||||
|
||||
if fastvideo_args.pipeline_config.ti2v_task and batch.pil_image is not None:
|
||||
logger.info("===========Using TI2V task===========")
|
||||
# TI2V directly replaces the first frame of the latent with
|
||||
# the image latent instead of appending along the channel dim
|
||||
assert batch.image_latent is None, "TI2V task should not have image latents"
|
||||
assert self.vae is not None, "VAE is not provided for TI2V task"
|
||||
z = self.vae.encode(batch.pil_image).mean.float()
|
||||
logger.info(f"z shape: {z.shape}")
|
||||
logger.info(f"latent_model_input shape: {latent_model_input.shape}")
|
||||
latent_model_input = latent_model_input.squeeze(0)
|
||||
mask1, mask2 = masks_like([latent_model_input], zero=True)
|
||||
# logger.info(f"mask1 shape: {mask1.shape}")
|
||||
# logger.info(f"mask2 shape: {mask2.shape}")
|
||||
latent_model_input = (1. -
|
||||
mask2[0]) * z + mask2[0] * latent_model_input
|
||||
# latent_model_input = latent_model_input.unsqueeze(0)
|
||||
latent_model_input = latent_model_input.to(get_local_torch_device())
|
||||
latents = latent_model_input
|
||||
F = batch.num_frames
|
||||
temporal_scale = fastvideo_args.pipeline_config.vae_config.arch_config.scale_factor_temporal
|
||||
spatial_scale = fastvideo_args.pipeline_config.vae_config.arch_config.scale_factor_spatial
|
||||
patch_size = fastvideo_args.pipeline_config.dit_config.arch_config.patch_size
|
||||
seq_len = ((F - 1) // temporal_scale +
|
||||
1) * (batch.height // spatial_scale) * (
|
||||
batch.width // spatial_scale) // (patch_size[1] *
|
||||
patch_size[2])
|
||||
import math
|
||||
seq_len = int(math.ceil(seq_len / sp_world_size)) * sp_world_size
|
||||
logger.info("latents shape: %s", latents.shape)
|
||||
|
||||
# Run denoising loop
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
# logger.info(f"seq_len: {seq_len}")
|
||||
logger.info(f"init timesteps: {timesteps}")
|
||||
for i, t in enumerate(timesteps):
|
||||
# Skip if interrupted
|
||||
if hasattr(self, 'interrupt') and self.interrupt:
|
||||
@@ -194,15 +230,43 @@ class DenoisingStage(PipelineStage):
|
||||
# Expand latents for I2V
|
||||
latent_model_input = latents.to(target_dtype)
|
||||
if batch.image_latent is not None:
|
||||
assert not fastvideo_args.pipeline_config.ti2v_task, "image latents should not be provided for TI2V task"
|
||||
latent_model_input = torch.cat(
|
||||
[latent_model_input, batch.image_latent],
|
||||
dim=1).to(target_dtype)
|
||||
if fastvideo_args.pipeline_config.ti2v_task and batch.pil_image is not None:
|
||||
logger.info(f"before ti2v timestep: {t}")
|
||||
timestep = [t]
|
||||
timestep = torch.stack(timestep).to(
|
||||
get_local_torch_device())
|
||||
|
||||
logger.info(f"mask2 shape: {mask2[0].shape}")
|
||||
logger.info(f"mask[0][0] shape: {mask2[0][0].shape}")
|
||||
logger.info(
|
||||
f"mask[0][0][:, ::2, ::2] shape: {mask2[0][0][:, ::2, ::2].shape}"
|
||||
)
|
||||
temp_ts = (mask2[0][0][:, ::2, ::2] * timestep)
|
||||
logger.info(f"temp_ts shape before flatten: {temp_ts.shape}")
|
||||
temp_ts = temp_ts.flatten()
|
||||
logger.info(f"temp_ts: {temp_ts}")
|
||||
logger.info(f"temp_ts shape: {temp_ts.shape}")
|
||||
temp_ts = torch.cat([
|
||||
temp_ts,
|
||||
temp_ts.new_ones(seq_len - temp_ts.size(0)) * timestep
|
||||
])
|
||||
timestep = temp_ts.unsqueeze(0)
|
||||
logger.info(f"after ti2v timestep: {timestep}")
|
||||
t_expand = timestep.repeat(latent_model_input.shape[0], 1)
|
||||
else:
|
||||
t_expand = t.repeat(latent_model_input.shape[0])
|
||||
logger.info(f"t_expand shape: {t_expand.shape}")
|
||||
# logger.info(f"t_expand: {t_expand}")
|
||||
|
||||
assert torch.isnan(latent_model_input).sum() == 0
|
||||
latent_model_input = self.scheduler.scale_model_input(
|
||||
latent_model_input, t)
|
||||
|
||||
# Prepare inputs for transformer
|
||||
t_expand = t.repeat(latent_model_input.shape[0])
|
||||
guidance_expand = (
|
||||
torch.tensor(
|
||||
[fastvideo_args.pipeline_config.embedded_cfg_scale] *
|
||||
@@ -257,6 +321,8 @@ class DenoisingStage(PipelineStage):
|
||||
# fastvideo_args=fastvideo_args
|
||||
):
|
||||
# Run transformer
|
||||
cuda_memory_before = torch.cuda.memory_allocated()
|
||||
logger.info(f"cuda memory before transformer: {cuda_memory_before / 1024 / 1024 / 1024} GB")
|
||||
noise_pred = self.transformer(
|
||||
latent_model_input,
|
||||
prompt_embeds,
|
||||
@@ -303,6 +369,10 @@ class DenoisingStage(PipelineStage):
|
||||
latents,
|
||||
**extra_step_kwargs,
|
||||
return_dict=False)[0]
|
||||
if fastvideo_args.pipeline_config.ti2v_task and batch.pil_image is not None:
|
||||
latents = latents.squeeze(0)
|
||||
latents = (1. - mask2[0]) * z + mask2[0] * latents
|
||||
# latents = latents.unsqueeze(0)
|
||||
|
||||
# Update progress bar
|
||||
if i == len(timesteps) - 1 or (
|
||||
|
||||
@@ -12,6 +12,9 @@ from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.validators import (StageValidators,
|
||||
VerificationResult)
|
||||
from fastvideo.utils import best_output_size
|
||||
from PIL import Image
|
||||
import torchvision.transforms.functional as TF
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -100,6 +103,35 @@ class InputValidationStage(PipelineStage):
|
||||
image = load_image(batch.image_path)
|
||||
batch.pil_image = image
|
||||
|
||||
img = batch.pil_image
|
||||
ih, iw = img.height, img.width
|
||||
logger.info(f"img height: {ih}, img width: {iw}")
|
||||
patch_size = fastvideo_args.pipeline_config.dit_config.arch_config.patch_size
|
||||
vae_stride = fastvideo_args.pipeline_config.vae_config.arch_config.scale_factor_spatial
|
||||
logger.info(f"patch_size: {patch_size}, vae_stride: {vae_stride}")
|
||||
dh, dw = patch_size[1] * vae_stride, patch_size[2] * vae_stride
|
||||
max_area = 704 * 1280
|
||||
ow, oh = best_output_size(iw, ih, dw, dh, max_area)
|
||||
|
||||
scale = max(ow / iw, oh / ih)
|
||||
img = img.resize((round(iw * scale), round(ih * scale)), Image.LANCZOS)
|
||||
logger.info(f"resized img height: {img.height}, img width: {img.width}")
|
||||
|
||||
# center-crop
|
||||
x1 = (img.width - ow) // 2
|
||||
y1 = (img.height - oh) // 2
|
||||
img = img.crop((x1, y1, x1 + ow, y1 + oh))
|
||||
assert img.width == ow and img.height == oh
|
||||
# logger.info(f"img shape: {img.shape}")
|
||||
|
||||
# to tensor
|
||||
img = TF.to_tensor(img).sub_(0.5).div_(0.5).to(self.device).unsqueeze(1)
|
||||
logger.info(f"img shape: {img.shape}")
|
||||
img = img.unsqueeze(0)
|
||||
batch.height = oh
|
||||
batch.width = ow
|
||||
batch.pil_image = img
|
||||
|
||||
return batch
|
||||
|
||||
def verify_input(self, batch: ForwardBatch,
|
||||
|
||||
@@ -812,3 +812,66 @@ def set_random_seed(seed: int) -> None:
|
||||
@lru_cache(maxsize=1)
|
||||
def is_vsa_available() -> bool:
|
||||
return importlib.util.find_spec("vsa") is not None
|
||||
|
||||
|
||||
# adapted from: https://github.com/Wan-Video/Wan2.2/blob/main/wan/utils/utils.py
|
||||
def masks_like(tensor,
|
||||
zero=False,
|
||||
generator=None,
|
||||
p=0.2) -> tuple[list[torch.Tensor], list[torch.Tensor]]:
|
||||
assert isinstance(tensor, list)
|
||||
out1 = [torch.ones(u.shape, dtype=u.dtype, device=u.device) for u in tensor]
|
||||
|
||||
out2 = [torch.ones(u.shape, dtype=u.dtype, device=u.device) for u in tensor]
|
||||
|
||||
if zero:
|
||||
if generator is not None:
|
||||
for u, v in zip(out1, out2, strict=False):
|
||||
random_num = torch.rand(1,
|
||||
generator=generator,
|
||||
device=generator.device).item()
|
||||
if random_num < p:
|
||||
u[:, 0] = torch.normal(mean=-3.5,
|
||||
std=0.5,
|
||||
size=(1, ),
|
||||
device=u.device,
|
||||
generator=generator).expand_as(
|
||||
u[:, 0]).exp()
|
||||
v[:, 0] = torch.zeros_like(v[:, 0])
|
||||
else:
|
||||
u[:, 0] = u[:, 0]
|
||||
v[:, 0] = v[:, 0]
|
||||
|
||||
else:
|
||||
for u, v in zip(out1, out2, strict=False):
|
||||
u[:, 0] = torch.zeros_like(u[:, 0])
|
||||
v[:, 0] = torch.zeros_like(v[:, 0])
|
||||
|
||||
return out1, out2
|
||||
|
||||
|
||||
# adapted from: https://github.com/Wan-Video/Wan2.2/blob/main/wan/utils/utils.py
|
||||
def best_output_size(w, h, dw, dh, expected_area):
|
||||
# float output size
|
||||
ratio = w / h
|
||||
ow = (expected_area * ratio)**0.5
|
||||
oh = expected_area / ow
|
||||
|
||||
# process width first
|
||||
ow1 = int(ow // dw * dw)
|
||||
oh1 = int(expected_area / ow1 // dh * dh)
|
||||
assert ow1 % dw == 0 and oh1 % dh == 0 and ow1 * oh1 <= expected_area
|
||||
ratio1 = ow1 / oh1
|
||||
|
||||
# process height first
|
||||
oh2 = int(oh // dh * dh)
|
||||
ow2 = int(expected_area / oh2 // dw * dw)
|
||||
assert oh2 % dh == 0 and ow2 % dw == 0 and ow2 * oh2 <= expected_area
|
||||
ratio2 = ow2 / oh2
|
||||
|
||||
# compare ratios
|
||||
if max(ratio / ratio1, ratio1 / ratio) < max(ratio / ratio2,
|
||||
ratio2 / ratio):
|
||||
return ow1, oh1
|
||||
else:
|
||||
return ow2, oh2
|
||||
|
||||
@@ -40,7 +40,8 @@ class MultiprocExecutor(Executor):
|
||||
logger.info("Using provided master port: %s", self.master_port)
|
||||
else:
|
||||
# Auto-find available port
|
||||
for port in range(29503, 65535):
|
||||
import random
|
||||
for port in range(29503 + random.randint(0, 10000), 65535):
|
||||
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
|
||||
if s.connect_ex(('localhost', port)) != 0:
|
||||
self.master_port = port
|
||||
|
||||
Reference in New Issue
Block a user