Compare commits

...
Author SHA1 Message Date
SolitaryThinker e6548c62b1 metrics and fixes 2025-08-04 05:32:13 +00:00
SolitaryThinker a3c1e34f71 fix 2025-08-04 00:24:59 +00:00
SolitaryThinker 38939bb4e7 Merge remote-tracking branch 'refs/remotes/origin/matthew/demo' into matthew/demo 2025-08-04 00:18:42 +00:00
SolitaryThinker 4e7484e09e update start 2025-08-04 00:12:45 +00:00
SolitaryThinker 1f96c53506 update 2025-08-04 00:11:26 +00:00
SolitaryThinker 8a3352d722 update ui 2025-08-03 23:37:28 +00:00
SolitaryThinker 0d71c0c875 update 2025-08-03 18:06:56 +00:00
SolitaryThinker 3a1073a7b2 fix 2025-08-03 03:12:57 +00:00
SolitaryThinker 2294254c4d attempt at perf fix 2025-08-03 02:10:25 +00:00
SolitaryThinker 017e68badb prompts 2025-08-03 00:45:45 +00:00
SolitaryThinker f248434368 extract per stage logging info 2025-08-03 00:43:57 +00:00
SolitaryThinker a00b539c5f update ui 2025-08-03 00:22:44 +00:00
SolitaryThinker 5d7beeb5a0 checkpoint 2025-08-02 01:22:21 +00:00
SolitaryThinker 7b3c4fbdd1 fix 2025-07-31 19:34:08 +00:00
SolitaryThinker f041bd3f49 fix 2025-07-31 08:53:50 +00:00
SolitaryThinker fbb9a73ab4 checkpotint 2025-07-31 08:34:39 +00:00
SolitaryThinker daea023383 checkpoint 2025-07-31 08:10:19 +00:00
SolitaryThinker 0422d18377 checkpoint 2025-07-31 06:39:14 +00:00
SolitaryThinker f46c6d923f live demo checkpoint 2025-07-31 03:15:37 +00:00
SolitaryThinker 79a31b0b45 copy over i2v changes for ti2v 5B 2025-07-31 03:15:37 +00:00
SolitaryThinker ef98769b46 reverse proxy in progress 2025-07-31 03:15:37 +00:00
SolitaryThinker 869c5f7370 checkpoint 2025-07-31 03:15:37 +00:00
SolitaryThinker 74529f22a4 can start replica across nodes 2025-07-31 03:15:37 +00:00
SolitaryThinker 2a2d67e792 move init 2025-07-31 03:15:37 +00:00
SolitaryThinker bb3367c9ac add demo 2025-07-31 03:15:37 +00:00
42 changed files with 4479 additions and 27 deletions
+200
View File
@@ -0,0 +1,200 @@
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.
+28 -13
View File
@@ -1,6 +1,10 @@
from fastvideo import VideoGenerator
# from fastvideo.configs.sample import SamplingParam
from fastvideo.configs.sample import SamplingParam
import os
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
OUTPUT_PATH = "video_samples"
def main():
@@ -8,30 +12,41 @@ def main():
# model.
# If a local path is provided, FastVideo will make a best effort
# attempt to identify the optimal arguments.
model_name = "FastVideo/FastWan2.1-T2V-14B-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,
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
text_encoder_cpu_offload=False,
# 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."
# )
results = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, sampling_param=sampling_param)
stage_names = results["stage_names"]
stage_execution_times = results["stage_execution_times"]
# print(logging_info)
print(stage_names)
print(stage_execution_times)
# 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!
+2 -1
View File
@@ -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
+38
View File
@@ -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()
+501
View File
@@ -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,936 @@
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
import time
from fastvideo.configs.sample.base import SamplingParam
class RayServeClient:
def __init__(self, backend_url: str):
self.backend_url = backend_url
self.session = requests.Session()
def check_health(self) -> bool:
"""Check if the backend is healthy"""
try:
response = self.session.get(f"{self.backend_url}/health", timeout=5)
return response.status_code == 200
except requests.exceptions.RequestException:
return False
def generate_video(self, request_data: dict) -> dict:
"""Generate video using the backend API"""
start_time = time.time()
try:
headers = {"Content-Type": "application/json"}
response = self.session.post(
f"{self.backend_url}/generate_video",
json=request_data,
headers=headers,
timeout=300 # 5 minutes timeout
)
end_time = time.time()
round_trip_time = end_time - start_time
if response.status_code == 200:
result = response.json()
# Calculate network time by subtracting backend total time
backend_total = result.get("total_time", 0)
network_time = round_trip_time - backend_total
result["network_time"] = network_time
return result
else:
return {"success": False, "error_message": f"HTTP {response.status_code}: {response.text}"}
except requests.exceptions.RequestException as e:
return {"success": False, "error_message": f"Request failed: {str(e)}"}
def save_video_from_base64(video_data: str, output_dir: str, prompt: str) -> str:
"""Save base64-encoded video data to a file"""
if not video_data:
return "No video data to save"
try:
# Remove the data URL prefix if present
if video_data.startswith('data:video/'):
video_data = video_data.split(',')[1]
# Decode base64 to bytes
video_bytes = base64.b64decode(video_data)
# Create safe filename from prompt
safe_prompt = prompt[:50].replace(' ', '_').replace('/', '_').replace('\\', '_')
video_filename = f"{safe_prompt}.mp4"
video_path = os.path.join(output_dir, video_filename)
# Ensure output directory exists
os.makedirs(output_dir, exist_ok=True)
# Save video bytes to file
with open(video_path, 'wb') as f:
f.write(video_bytes)
return f"Saved video to: {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: dict[str, 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,
randomize_seed=False,
input_image=None,
model_selection="FastVideo/FastWan2.1-T2V-1.3B-Diffusers (Text-to-Video)",
progress=None,
request: gr.Request = 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}", ""
# Validate video dimensions
max_pixels = 720 * 1280
if height * width > max_pixels:
return None, f"Video dimensions too large. Maximum allowed: 720x1280 pixels. Current: {height}x{width} = {height*width} pixels", ""
# Update progress
if progress:
progress(0.1, desc="Checking backend health...")
# Determine if this is I2V based on model selection - I2V functionality commented out
# is_i2v = "I2V" in model_selection or "Image-to-Video" in model_selection
# Handle input image for I2V - I2V functionality commented out
# image_path = None
# if is_i2v and input_image is not None:
# if progress:
# progress(0.2, desc="Processing input image...")
# 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
if progress:
progress(0.3, desc="Preparing request...")
# Map clean model names to full paths
model_path_mapping = {
"FastWan2.1-T2V-1.3B": "FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
"FastWan2.1-T2V-14B": "FastVideo/FastWan2.1-T2V-14B-Diffusers",
"FastWan2.2-TI2V-5B": "FastVideo/FastWan2.2-TI2V-5B-Diffusers"
}
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": 20, # Use default value
"randomize_seed": randomize_seed,
"return_frames": False, # We'll get video data directly
"image_path": None, # For T2V, we pass None as input_image
# "model_type": "i2v" if "I2V" in model_selection or "Image-to-Video" in model_selection else "t2v", # Use model type selection
"model_type": model_path_mapping[model_selection],
"model_path": model_selection.split(" (")[0] if model_selection else None # Extract model path from selection
}
# Send request to backend
if progress:
progress(0.4, desc="Sending request to backend...")
response = client.generate_video(request_data)
if progress:
progress(0.8, desc="Processing response...")
# 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):
video_data = response.get("video_data", "")
used_seed = response.get("seed", seed)
generation_time = response.get("generation_time", 0.0)
inference_time = response.get("inference_time", 0.0)
encoding_time = response.get("encoding_time", 0.0)
total_time = response.get("total_time", 0.0)
network_time = response.get("network_time", 0.0)
stage_names = response.get("stage_names", [])
stage_execution_times = response.get("stage_execution_times", [])
print(f"Used seed: {used_seed}")
print(f"Inference time: {inference_time:.2f}s")
print(f"Encoding time: {encoding_time:.2f}s")
print(f"Network transfer: {network_time:.2f}s")
print(f"Total time: {total_time:.2f}s")
print(f"Stage names: {stage_names}")
print(f"Stage execution times: {stage_execution_times}")
# Create detailed timing message with all cards in a single row
timing_details = f"""
<div style="margin: 10px 0;">
<h3 style="text-align: center; margin-bottom: 10px;">⏱️ Timing Breakdown</h3>
<div style="display: grid; grid-template-columns: repeat(4, 1fr); gap: 10px; margin-bottom: 10px;">
<div class="timing-card timing-card-highlight">
<div style="font-size: 20px;">🧠</div>
<div style="font-weight: bold; margin: 3px 0; font-size: 14px;">Model Inference</div>
<div style="font-size: 16px; color: #2563eb; font-weight: bold;">{inference_time:.2f}s</div>
</div>
<div class="timing-card">
<div style="font-size: 20px;">🎬</div>
<div style="font-weight: bold; margin: 3px 0; font-size: 14px;">Video Encoding</div>
<div style="font-size: 16px; color: #dc2626;">{encoding_time:.2f}s</div>
</div>
<div class="timing-card">
<div style="font-size: 20px;">🌐</div>
<div style="font-weight: bold; margin: 3px 0; font-size: 14px;">Network Transfer</div>
<div style="font-size: 16px; color: #059669;">{network_time:.2f}s</div>
</div>
<div class="timing-card">
<div style="font-size: 20px;">📊</div>
<div style="font-weight: bold; margin: 3px 0; font-size: 14px;">Total Processing</div>
<div style="font-size: 18px; color: #0277bd;">{total_time:.2f}s</div>
</div>
</div>"""
# timing_details += f"""
# <div style="margin-top: 15px;">
# <h4 style="text-align: center; margin-bottom: 10px;">🔄 Processing Stages</h4>
# <div style="display: grid; grid-template-columns: repeat(auto-fit, minmax(200px, 1fr)); gap: 8px;">
# """
#
# # Add individual stage timing cards
# for stage_name, stage_time in zip(stage_names, stage_execution_times):
# if stage_name.strip() and stage_time > 0: # Only show non-empty stages with valid times
# timing_details += f"""
# <div class="stage-card">
# <div style="font-weight: bold; font-size: 14px; margin-bottom: 5px;">{stage_name.strip()}</div>
# <div style="font-size: 16px; color: #7c3aed; font-weight: bold;">{stage_time:.2f}s</div>
# </div>
# """
#
# timing_details += """
# </div>
# </div>
# """
# Add performance insights
if inference_time > 0:
fps = num_frames / inference_time
timing_details += f"""
<div class="performance-card" style="margin-top: 15px;">
<span style="font-weight: bold;">Generation Speed: </span>
<span style="font-size: 18px; color: #6366f1; font-weight: bold;">{fps:.1f} frames/second</span>
</div>
"""
timing_details += "</div>"
# Save video data to file for Gradio to display
if video_data:
try:
if progress:
progress(0.9, desc="Saving video...")
output_dir = "outputs"
save_status, video_path = save_video_from_base64(video_data, output_dir, prompt)
print(f"Video save status: {save_status}")
if progress:
progress(1.0, desc="Generation complete!")
if video_path and os.path.exists(video_path):
return video_path, used_seed, timing_details
else:
return None, f"Video generated but failed to save: {save_status}", ""
except Exception as e:
return None, f"Failed to save video: {str(e)}", ""
else:
return None, "No video data received from backend", ""
else:
error_msg = response.get("error_message", "Unknown error occurred")
return None, f"Generation failed: {error_msg}", ""
# Example prompts
examples = []
example_labels = []
def contains_chinese(text):
"""Check if text contains Chinese characters"""
for char in text:
if '\u4e00' <= char <= '\u9fff': # CJK Unified Ideographs (Chinese characters)
return True
return False
# Load prompts from all text files in the prompts directory
prompts_dir = "prompts"
if os.path.exists(prompts_dir):
for filename in os.listdir(prompts_dir):
if filename.endswith('.txt'):
filepath = os.path.join(prompts_dir, filename)
try:
with open(filepath, "r", encoding='utf-8') as f:
for line_num, line in enumerate(f, 1):
line = line.strip()
if line and not contains_chinese(line): # Skip empty lines and lines with Chinese text
# Create a label from the first 100 characters
label = line[:100] + "..." if len(line) > 100 else line
example_labels.append(label)
examples.append(line)
except Exception as e:
print(f"Warning: Could not read {filepath}: {e}")
# Fallback to example_prompts.txt if prompts directory is empty or doesn't exist
if not examples:
try:
with open("example_prompts.txt", "r") as f:
for line in f:
line = line.strip()
if line:
example_labels.append(line[:100])
examples.append(line)
except Exception as e:
print(f"Warning: Could not read example_prompts.txt: {e}")
# Add a default example if all else fails
examples = ["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_labels = ["Crowded rooftop bar at night"]
# Create a custom theme with blue styling to match the logo
theme = gr.themes.Base().set(
button_primary_background_fill="#2563eb", # Blue color
button_primary_background_fill_hover="#1d4ed8", # Darker blue on hover
button_primary_text_color="white",
slider_color="#2563eb", # Blue slider
checkbox_background_color_selected="#2563eb", # Blue checkbox when selected
)
def get_default_values_for_model(model_selection_value):
"""Get default parameter values for the specified model"""
model_path = model_selection_value.split(" (")[0] if model_selection_value else None
if model_path and model_path in default_params:
params = default_params[model_path]
return {
'height': params.height,
'width': params.width,
'num_frames': params.num_frames,
'guidance_scale': params.guidance_scale,
'seed': params.seed,
}
else:
# Fallback defaults if model not found
return {
'height': 448,
'width': 832,
'num_frames': 61,
'guidance_scale': 3.0,
'seed': 1024,
}
# Get initial values for the default model
default_model = "FastVideo/FastWan2.1-T2V-1.3B-Diffusers (Text-to-Video)"
initial_values = get_default_values_for_model(default_model)
# Create Gradio interface
with gr.Blocks(title="FastWan", theme=theme) as demo:
# Logo using Gradio's Image component
gr.Image("fastvideo-logos/main/png/full.png", show_label=False, container=False, height=80)
gr.HTML("""
<div style="text-align: center; margin-bottom: 10px;">
<p style="font-size: 18px;"> Make Video Generation Go Blurrrrrrr </p>
<p style="font-size: 18px;"> Twitter | <a href="https://github.com/hao-ai-lab/FastVideo/tree/main" target="_blank">Code</a> | Blog | <a href="https://hao-ai-lab.github.io/FastVideo/" target="_blank">Docs</a> </p>
</div>
""")
# What is FastVideo accordion
with gr.Accordion("🎥 What Is FastVideo?", open=False):
gr.HTML("""
<div style="padding: 20px; line-height: 1.6;">
<p style="font-size: 16px; margin-bottom: 15px;">
It features a clean, consistent API that works across popular video models, making it easier for developers to author new models and incorporate system- or kernel-level optimizations. With FastVideo's optimizations, you can achieve more than 3x inference improvement compared to other systems.
</p>
</div>
""")
# 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"
# Model selection dropdown
with gr.Row():
model_selection = gr.Dropdown(
choices=[
"FastWan2.1-T2V-1.3B",
"FastWan2.1-T2V-14B",
"FastWan2.2-TI2V-5B",
# "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers (Image-to-Video)" # I2V functionality commented out
],
value="FastWan2.1-T2V-1.3B",
label="Select Model",
interactive=True
)
# Examples dropdown
with gr.Row():
example_dropdown = gr.Dropdown(
choices=example_labels,
label="Example Prompts",
value=None,
interactive=True,
allow_custom_value=False
)
# Main interface
with gr.Row():
with gr.Column(scale=6):
prompt = gr.Text(
label="Prompt",
show_label=False,
max_lines=3,
placeholder="Describe your scene...",
container=False,
lines=3,
autofocus=True,
)
with gr.Column(scale=1, min_width=120, elem_classes="center-button"):
run_button = gr.Button("Run", variant="primary", size="lg")
# Status and timing information
with gr.Row():
with gr.Column():
error_output = gr.Text(label="Error", visible=False)
# frames_output = gr.Text(label="Generation Status", visible=False)
timing_display = gr.Markdown(label="Timing Breakdown", visible=False)
# Two-column layout: Advanced options on left, Video on right
with gr.Row(equal_height=True, elem_classes="main-content-row"):
# Left column - Advanced options
with gr.Column(scale=1, elem_classes="advanced-options-column"):
with gr.Group():
gr.HTML("<div style='margin: 0 0 15px 0; text-align: center; font-size: 16px;'>Advanced Options</div>")
with gr.Row():
height = gr.Slider(
label="Height",
minimum=256,
maximum=1280,
step=32,
value=initial_values['height'],
)
width = gr.Slider(
label="Width",
minimum=256,
maximum=1280,
step=32,
value=initial_values['width']
)
with gr.Row():
num_frames = gr.Slider(
label="Number of Frames",
minimum=16,
maximum=121,
step=16,
value=initial_values['num_frames'],
)
guidance_scale = gr.Slider(
label="Guidance Scale",
minimum=1,
maximum=12,
value=initial_values['guidance_scale'],
)
with gr.Row():
use_negative_prompt = gr.Checkbox(
label="Use negative prompt", value=False)
negative_prompt = gr.Text(
label="Negative prompt",
max_lines=3,
lines=3,
placeholder="Enter a negative prompt",
visible=False,
)
seed = gr.Slider(
label="Seed",
minimum=0,
maximum=1000000,
step=1,
value=initial_values['seed'],
)
randomize_seed = gr.Checkbox(label="Randomize seed", value=False)
seed_output = gr.Number(label="Used Seed")
# Right column - Video result
with gr.Column(scale=1, elem_classes="video-column"):
result = gr.Video(
label="Generated Video",
show_label=True,
height=436, # Adjusted height for better vertical alignment
width=600, # Limit video width
container=True
)
# Add CSS to position the button and constrain width
gr.HTML("""
<style>
.center-button {
display: flex !important;
justify-content: center !important;
height: 100% !important;
padding-top: 1.4em !important;
}
/* Constrain overall width */
.gradio-container {
max-width: 1200px !important;
margin: 0 auto !important;
}
.main {
max-width: 1200px !important;
margin: 0 auto !important;
}
/* Constrain individual components */
.gr-form, .gr-box, .gr-group {
max-width: 1200px !important;
}
/* Make video component smaller */
.gr-video {
max-width: 500px !important;
margin: 0 auto !important;
}
/* Ensure equal height columns */
.main-content-row {
display: flex !important;
align-items: flex-start !important;
min-height: 500px !important;
}
.advanced-options-column,
.video-column {
display: flex !important;
flex-direction: column !important;
flex: 1 !important;
min-height: 400px !important;
}
/* Force equal heights regardless of content */
.advanced-options-column > *:last-child,
.video-column > *:last-child {
flex-grow: 0 !important;
}
/* Responsive alignment for split screen */
@media (max-width: 1400px) {
.main-content-row {
min-height: 600px !important;
}
.advanced-options-column,
.video-column {
min-height: 600px !important;
}
}
@media (max-width: 1200px) {
.main-content-row {
flex-direction: column !important;
align-items: stretch !important;
}
.advanced-options-column,
.video-column {
min-height: auto !important;
width: 100% !important;
}
}
/* Theme-agnostic timing cards */
.timing-card {
background: var(--background-fill-secondary) !important;
border: 1px solid var(--border-color-primary) !important;
color: var(--body-text-color) !important;
padding: 10px;
border-radius: 8px;
text-align: center;
min-height: 80px;
display: flex;
flex-direction: column;
justify-content: center;
}
.timing-card-highlight {
background: var(--background-fill-primary) !important;
border: 2px solid var(--color-accent) !important;
}
.stage-card {
background: var(--background-fill-secondary) !important;
border: 1px solid var(--border-color-primary) !important;
color: var(--body-text-color) !important;
padding: 10px;
border-radius: 6px;
text-align: center;
}
.performance-card {
background: var(--background-fill-secondary) !important;
border: 1px solid var(--border-color-primary) !important;
color: var(--body-text-color) !important;
padding: 10px;
border-radius: 6px;
text-align: center;
}
/* Dark mode support */
.dark .timing-card {
background: var(--background-fill-secondary) !important;
border-color: var(--border-color-primary) !important;
color: var(--body-text-color) !important;
}
.dark .timing-card-highlight {
background: var(--background-fill-primary) !important;
border-color: var(--color-accent) !important;
}
.dark .stage-card,
.dark .performance-card {
background: var(--background-fill-secondary) !important;
border-color: var(--border-color-primary) !important;
color: var(--body-text-color) !important;
}
</style>
""")
# Function to update prompt when example is selected
def on_example_select(example_label):
if example_label and example_label in example_labels:
index = example_labels.index(example_label)
return examples[index]
return ""
example_dropdown.change(
fn=on_example_select,
inputs=example_dropdown,
outputs=prompt,
)
# Disclaimer text
gr.HTML("""
<div style="text-align: center; margin-top: 10px; margin-bottom: 15px;">
<p style="font-size: 16px; margin: 0;">The compute for this demo is generously provided by <a href="https://www.gmicloud.ai/" target="_blank">GMI Cloud</a>. Note that this demo is meant to showcase FastWan2.1's quality and that under a large number of requests, generation speed may be affected.</p>
</div>
""")
# Event handlers
use_negative_prompt.change(
fn=lambda x: gr.update(visible=x),
inputs=use_negative_prompt,
outputs=negative_prompt,
)
# Model selection change handler - I2V functionality commented out
# def on_model_selection_change(model_selection):
# is_i2v = "I2V" in model_selection or "Image-to-Video" in model_selection
# prompt_placeholder = "Describe how the image should animate" if is_i2v else "Enter your prompt"
# return gr.update(visible=is_i2v), gr.update(placeholder=prompt_placeholder)
#
# model_selection.change(
# fn=on_model_selection_change,
# inputs=model_selection,
# outputs=[input_image, prompt],
# )
def on_model_selection_change(selected_model):
"""Update advanced options based on selected model's default parameters"""
if not selected_model:
return {}, {}, {}, {}, {} # Return empty updates if no model selected
# Extract model path from selection (remove the description part)
model_path = selected_model.split(" ")[0] if selected_model else None
if model_path and model_path in default_params:
params = default_params[model_path]
# Update each component with the model's default values
return (
gr.update(value=params.height), # height
gr.update(value=params.width), # width
gr.update(value=params.num_frames), # num_frames
gr.update(value=params.guidance_scale), # guidance_scale
gr.update(value=params.seed), # seed
)
else:
# If model not found in default_params, return current values (no change)
return {}, {}, {}, {}, {}
model_selection.change(
fn=on_model_selection_change,
inputs=model_selection,
outputs=[height, width, num_frames, guidance_scale, seed],
)
def handle_generation(*args, progress=None, request: gr.Request = None):
# Extract model selection and input image from args - I2V functionality commented out
model_selection, prompt, negative_prompt, use_negative_prompt, seed, guidance_scale, num_frames, height, width, randomize_seed = args
# Determine if this is I2V based on model selection - I2V functionality commented out
# is_i2v = "I2V" in model_selection or "Image-to-Video" in model_selection
# For T2V, we pass None as input_image - I2V functionality commented out
# if not is_i2v:
# input_image = None
# Call the generate_video function with progress tracking
result_path, seed_or_error, timing_details = generate_video(
prompt, negative_prompt, use_negative_prompt, seed, guidance_scale,
num_frames, height, width, randomize_seed, None, model_selection, progress, request
)
if result_path and os.path.exists(result_path):
return (
result_path,
seed_or_error,
gr.update(visible=False), # error_output
# gr.update(visible=True, value="Generation completed successfully!"), # frames_output
gr.update(visible=True, value=timing_details), # timing_display
)
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), # timing_display
)
# Unified event handler
run_button.click(
fn=handle_generation,
inputs=[
model_selection,
prompt,
negative_prompt,
use_negative_prompt,
seed,
guidance_scale,
num_frames,
height,
width,
randomize_seed,
# input_image, # Removed input_image from inputs
],
outputs=[result, seed_output, error_output, timing_display],
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_paths",
type=str,
default="FastVideo/FastWan2.1-T2V-1.3B-Diffusers,FastVideo/FastWan2.1-T2V-14B-Diffusers",
help="Comma separated list of paths to the T2V model(s)")
# parser.add_argument("--t2v_model_replicas", type=int,
# default="4,4",
# help="Comma separated list of number of replicas for the T2V model(s)")
# parser.add_argument("--i2v_model_path", # I2V functionality commented out
# type=str,
# default="Wan-AI/Wan2.1-I2V-14B-480P-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 = {}
model_paths = args.t2v_model_paths.split(",")
for model_path in model_paths:
default_params[model_path] = SamplingParam.from_pretrained(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 Models: {args.t2v_model_paths}")
# print(f"T2V Model Replicas: {args.t2v_model_replicas}")
# print(f"I2V Model: {args.i2v_model_path}") # I2V functionality commented out
# Use FastAPI to serve custom HTML with proper Open Graph metadata
from fastapi import FastAPI
from fastapi.responses import HTMLResponse, FileResponse
import uvicorn
import os
app = FastAPI()
@app.get("/logo.png")
async def get_logo():
return FileResponse("fastvideo-logos/main/png/full.png", media_type="image/png")
@app.get("/", response_class=HTMLResponse)
def index():
return """
<!DOCTYPE html>
<html lang="en">
<head>
<meta charset="UTF-8" />
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
<meta property="og:title" content="FastWan" />
<meta property="og:description" content="Make video generation go blurrrrrrr" />
<meta property="og:image" content="https://fastwan.fastvideo.org/logo.png" />
<meta property="og:url" content="https://fastwan.fastvideo.org/" />
<meta property="og:type" content="website" />
<meta name="twitter:card" content="summary_large_image" />
<meta name="twitter:title" content="FastWan" />
<meta name="twitter:description" content="Make video generation go blurrrrrrr" />
<meta name="twitter:image" content="https://fastwan.fastvideo.org/logo.png" />
<title>FastWan</title>
<link rel="icon" type="image/png" href="/gradio/file/fastvideo-logos/main/png/icon-simple.png">
<style>
body, html {
margin: 0;
padding: 0;
height: 100%;
overflow: hidden;
}
iframe {
width: 100%;
height: 100vh;
border: none;
}
</style>
</head>
<body>
<iframe src="/gradio" width="100%" height="100%" style="border: none;"></iframe>
</body>
</html>
"""
# Mount Gradio app under /gradio
app = gr.mount_gradio_app(
app,
demo,
path="/gradio",
allowed_paths=[os.path.abspath("outputs"), os.path.abspath("temp_images"), os.path.abspath("fastvideo-logos")]
)
# Run the FastAPI server
uvicorn.run(app, host=args.host, port=args.port)
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()
+181
View File
@@ -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";
}
}
}
+173
View File
@@ -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,621 @@
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
import imageio
from ray.serve.handle import DeploymentHandle
NUM_GPUS = 8
SUPPORTED_MODELS = [
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
"FastVideo/FastWan2.1-T2V-14B-Diffusers",
"FastVideo/FastWan2.2-TI2V-5B-Diffusers",
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
"Wan-AI/Wan2.1-T2V-14B-Diffusers",
"Wan-AI/Wan2.2-TI2V-5B-Diffusers",
]
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
model_path: Optional[str] = None # Specific model path to use
class VideoGenerationResponse(BaseModel):
video_data: Optional[str] = None # Base64 encoded video data
seed: int
success: bool
error_message: Optional[str] = None
generation_time: Optional[float] = None
# Detailed timing information
model_load_time: Optional[float] = None
inference_time: Optional[float] = None
encoding_time: Optional[float] = None
total_time: Optional[float] = None
stage_names: Optional[List[str]] = None
stage_execution_times: Optional[List[float]] = None
def encode_video_to_base64(frames: List[np.ndarray], fps: int = 24) -> str:
"""Convert numpy frames to base64-encoded MP4 video"""
if not frames:
return ""
try:
# Save frames to bytes buffer as MP4
buffer = io.BytesIO()
imageio.mimsave(buffer, frames, fps=fps, format="mp4")
buffer.seek(0)
# Encode to base64
video_base64 = base64.b64encode(buffer.getvalue()).decode('utf-8')
return f"data:video/mp4;base64,{video_base64}"
except Exception as e:
print(f"Warning: Failed to encode video: {e}")
return ""
# ---------------------------------------------------------------------------
# Model-Specific Deployments (one per model type)
# These deployments each load a single model and expose a `generate_video`
# method that can be invoked via a DeploymentHandle. Each deployment can be
# scaled independently by configuring `num_replicas` when binding.
# ---------------------------------------------------------------------------
@serve.deployment( # T2V 1.3B model deployment
ray_actor_options={"num_cpus": 10, "num_gpus": 1, "runtime_env": {"conda": "fv"}},
)
class T2VModelDeployment:
"""Serve deployment wrapping the 1.3 B text-to-video model."""
def __init__(self, t2v_model_path: str, output_path: str = "outputs"):
self.model_path = t2v_model_path
self.output_path = output_path
# Ensure output directory exists
os.makedirs(self.output_path, exist_ok=True)
# Delay helps avoid GPU contention when many replicas start at once
time.sleep(5)
# Ensure correct attention backend for FastVideo
if "FastVideo" in self.model_path:
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
else:
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "FLASH_ATTN"
os.environ["FASTVIDEO_STAGE_LOGGING"] = "1"
# Lazy import to keep the deployment import-safe on head node
from fastvideo.entrypoints.video_generator import VideoGenerator
from fastvideo.configs.sample.base import SamplingParam
print(f"Initializing T2V model: {self.model_path}")
self.generator = VideoGenerator.from_pretrained(
model_path=self.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,
enable_stage_verification=False,
)
self.default_params = SamplingParam.from_pretrained(self.model_path)
print("✅ T2V model initialized successfully")
def generate_video(self, video_request: "VideoGenerationRequest") -> "VideoGenerationResponse":
"""Generate a video for the given request and return an encoded response."""
import time
total_start_time = time.time()
# Deep-copy default sampling params and override with request values
params = deepcopy(self.default_params)
params.prompt = video_request.prompt
if video_request.use_negative_prompt:
params.negative_prompt = video_request.negative_prompt
params.seed = video_request.seed if not video_request.randomize_seed else torch.randint(0, 1_000_000, (1,)).item()
# Ensure the generator respects the chosen seed strategy
params.randomize_seed = video_request.randomize_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
# Do not write to disk when called via API
params.save_video = False
params.return_frames = False
# Track inference time
inference_start_time = time.time()
# Generate video frames
result = self.generator.generate_video(
prompt=video_request.prompt,
sampling_param=params,
save_video=False,
return_frames=False,
)
inference_end_time = time.time()
inference_time = inference_end_time - inference_start_time
frames = result if isinstance(result, list) else result.get("frames", [])
generation_time = result.get("generation_time", 0.0) if isinstance(result, dict) else 0.0
logging_info = result.get("logging_info", None)
if logging_info:
stage_names = logging_info.get_execution_order()
stage_execution_times = [logging_info.get_stage_info(stage_name).get("execution_time", 0.0) for stage_name in stage_names]
else:
stage_names = []
stage_execution_times = []
# Track encoding time
encoding_start_time = time.time()
# Encode outputs
video_data = encode_video_to_base64(frames, fps=16)
encoding_end_time = time.time()
encoding_time = encoding_end_time - encoding_start_time
total_end_time = time.time()
total_time = total_end_time - total_start_time
return VideoGenerationResponse(
video_data=video_data,
seed=params.seed,
success=True,
generation_time=generation_time,
inference_time=inference_time,
encoding_time=encoding_time,
total_time=total_time,
stage_names=stage_names,
stage_execution_times=stage_execution_times,
)
# @serve.deployment( # I2V 14B model deployment - I2V functionality commented out
# ray_actor_options={"num_cpus": 10, "num_gpus": 1, "runtime_env": {"conda": "fv"}},
# )
# class I2VModelDeployment:
# """Serve deployment wrapping the 14 B image-to-video model."""
#
# def __init__(self, i2v_model_path: str, output_path: str = "outputs"):
# self.model_path = i2v_model_path
# self.output_path = output_path
#
# os.makedirs(self.output_path, exist_ok=True)
# time.sleep(10)
#
# os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
#
# from fastvideo.entrypoints.video_generator import VideoGenerator
# from fastvideo.configs.sample.base import SamplingParam
#
# print(f"Initializing I2V model: {self.model_path}")
# self.generator = VideoGenerator.from_pretrained(
# model_path=self.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.default_params = SamplingParam.from_pretrained(self.model_path)
# print("✅ I2V model initialized successfully")
#
# def generate_video(self, video_request: "VideoGenerationRequest") -> "VideoGenerationResponse":
# import time
#
# total_start_time = time.time()
#
# params = deepcopy(self.default_params)
#
# params.prompt = video_request.prompt
# if video_request.use_negative_prompt:
# params.negative_prompt = video_request.negative_prompt
#
# params.seed = video_request.seed if not video_request.randomize_seed else torch.randint(0, 1_000_000, (1,)).item()
# 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
#
# if video_request.image_path:
# params.image_path = video_request.image_path
#
# params.save_video = False
# params.return_frames = True
#
# # Track inference time
# inference_start_time = time.time()
#
# result = self.generator.generate_video(
# prompt=video_request.prompt,
# sampling_param=params,
# save_video=False,
# return_frames=True,
# )
#
# inference_end_time = time.time()
# inference_time = inference_end_time - inference_start_time
#
# frames = result if isinstance(result, list) else result.get("frames", [])
# generation_time = result.get("generation_time", 0.0) if isinstance(result, dict) else 0.0
#
# # Track encoding time
# encoding_start_time = time.time()
#
# video_data = encode_video_to_base64(frames, fps=24)
# encoded_frames = encode_frames_to_base64(frames) if video_request.return_frames and frames else None
#
# encoding_end_time = time.time()
# encoding_time = encoding_end_time - encoding_start_time
#
# total_end_time = time.time()
# total_time = total_end_time - total_start_time
#
# return VideoGenerationResponse(
# video_data=video_data,
# frames=encoded_frames,
# seed=params.seed,
# success=True,
# generation_time=generation_time,
# inference_time=inference_time,
# encoding_time=encoding_time,
# total_time=total_time,
# )
@serve.deployment( # T2V 14B model deployment with optimized settings
ray_actor_options={"num_cpus": 16, "num_gpus": 1, "runtime_env": {"conda": "fv"}},
)
class T2V14BModelDeployment:
"""Serve deployment wrapping the 14B text-to-video model with optimized settings."""
def __init__(self, t2v_14b_model_path: str, output_path: str = "outputs"):
self.model_path = t2v_14b_model_path
self.output_path = output_path
# Ensure output directory exists
os.makedirs(self.output_path, exist_ok=True)
# Delay helps avoid GPU contention when many replicas start at once
time.sleep(10)
# Ensure correct attention backend for FastVideo
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
os.environ["FASTVIDEO_STAGE_LOGGING"] = "1"
# Lazy import to keep the deployment import-safe on head node
from fastvideo.entrypoints.video_generator import VideoGenerator
from fastvideo.configs.sample.base import SamplingParam
print(f"Initializing T2V 14B model: {self.model_path}")
self.generator = VideoGenerator.from_pretrained(
model_path=self.model_path,
num_gpus=1,
use_fsdp_inference=True,
text_encoder_cpu_offload=True, # Enable CPU offload for 14B model
dit_cpu_offload=True, # Enable CPU offload for 14B model
vae_cpu_offload=False,
VSA_sparsity=0.9, # Higher sparsity for 14B model
enable_stage_verification=False,
)
self.default_params = SamplingParam.from_pretrained(self.model_path)
print("✅ T2V 14B model initialized successfully")
def generate_video(self, video_request: "VideoGenerationRequest") -> "VideoGenerationResponse":
"""Generate a video for the given request and return an encoded response."""
import time
total_start_time = time.time()
# Deep-copy default sampling params and override with request values
params = deepcopy(self.default_params)
params.prompt = video_request.prompt
if video_request.use_negative_prompt:
params.negative_prompt = video_request.negative_prompt
params.seed = video_request.seed if not video_request.randomize_seed else torch.randint(0, 1_000_000, (1,)).item()
# Ensure the generator respects the chosen seed strategy
params.randomize_seed = video_request.randomize_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
# Do not write to disk when called via API
params.save_video = False
params.return_frames = False
# Track inference time
inference_start_time = time.time()
# Generate video frames
result = self.generator.generate_video(
prompt=video_request.prompt,
sampling_param=params,
save_video=False,
return_frames=False,
)
inference_end_time = time.time()
inference_time = inference_end_time - inference_start_time
frames = result if isinstance(result, list) else result.get("frames", [])
generation_time = result.get("generation_time", 0.0) if isinstance(result, dict) else 0.0
logging_info = result.get("logging_info", None)
if logging_info:
stage_names = logging_info.get_execution_order()
stage_execution_times = [logging_info.get_stage_info(stage_name).get("execution_time", 0.0) for stage_name in stage_names]
else:
stage_names = []
stage_execution_times = []
# Track encoding time
encoding_start_time = time.time()
# Encode outputs
video_data = encode_video_to_base64(frames, fps=16)
encoding_end_time = time.time()
encoding_time = encoding_end_time - encoding_start_time
total_end_time = time.time()
total_time = total_end_time - total_start_time
return VideoGenerationResponse(
video_data=video_data,
seed=params.seed,
success=True,
generation_time=generation_time,
inference_time=inference_time,
encoding_time=encoding_time,
total_time=total_time,
stage_names=stage_names,
stage_execution_times=stage_execution_times,
)
# 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)
# Import Prometheus metrics (but don't instantiate at module level)
from prometheus_client import Counter, Histogram, generate_latest, CONTENT_TYPE_LATEST
import time
@serve.deployment(ray_actor_options={"num_cpus": 2})
@serve.ingress(app)
class FastVideoAPI:
"""Ingress deployment that routes requests to either the T2V or I2V model deployment."""
def __init__(self, t2v_deployments: Dict[str, DeploymentHandle]): # Removed i2v_deployment
self.t2v_deployments = t2v_deployments
# self.t2v_14b_handle = t2v_14b_deployment
# self.i2v_handle = i2v_deployment # I2V functionality commented out
# Initialize Prometheus metrics inside the deployment
self.request_count = Counter('fastvideo_requests_total', 'Total FastVideo requests', ['model_type', 'status'])
self.request_duration = Histogram('fastvideo_request_duration_seconds', 'FastVideo request duration', ['model_type'])
self.video_generation_time = Histogram('fastvideo_video_generation_seconds', 'Video generation time', ['model_type'])
@app.post("/generate_video", response_model=VideoGenerationResponse)
@limiter.limit("50/minute") # Allow 50 requests per minute per IP
async def generate_video(self, request: Request, video_request: VideoGenerationRequest) -> VideoGenerationResponse:
"""Route the request to the appropriate model deployment based on `model_type` and `model_path`."""
start_time = time.time()
model_type = video_request.model_path.split('/')[-1] if video_request.model_path else "unknown"
try:
# assert video_request.model_type in self.t2v_deployments, f"Model {video_request.model_type} not found"
if video_request.model_type not in self.t2v_deployments:
raise ValueError(f"Model {video_request.model_type} not found")
response_ref = self.t2v_deployments[video_request.model_type].generate_video.remote(video_request)
# if video_request.model_type.lower() == "i2v": # I2V functionality commented out
# response_ref = self.i2v_handle.generate_video.remote(video_request)
# if video_request.model_type.lower() == "t2v":
# # Route T2V requests based on model path
# if video_request.model_path and "14b" in video_request.model_path.lower():
# response_ref = self.t2v_14b_handle.generate_video.remote(video_request)
# else:
# # Default to 1.3B model
# response_ref = self.t2v_handle.generate_video.remote(video_request)
# else:
# # Default to 1.3B T2V model
# response_ref = self.t2v_handle.generate_video.remote(video_request)
# Await the remote response
response = await response_ref
# Record success metrics
self.request_count.labels(model_type=model_type, status="success").inc()
self.request_duration.labels(model_type=model_type).observe(time.time() - start_time)
if hasattr(response, 'generation_time_seconds') and response.generation_time_seconds:
self.video_generation_time.labels(model_type=model_type).observe(response.generation_time_seconds)
return response
except Exception as e:
# Record error metrics
self.request_count.labels(model_type=model_type, status="error").inc()
self.request_duration.labels(model_type=model_type).observe(time.time() - start_time)
return VideoGenerationResponse(
video_data=None,
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"}
@app.get("/metrics")
async def metrics(self):
"""Expose Prometheus metrics"""
from fastapi import Response
return Response(generate_latest(), media_type="text/plain")
def start_ray_serve(
*,
t2v_model_paths: str,
t2v_model_replicas: str,
# t2v_14b_model_path: str = "FastVideo/FastWan2.1-T2V-14B-Diffusers",
# i2v_model_path: str = "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers", # I2V functionality commented out
output_path: str = "outputs",
host: str = "0.0.0.0",
port: int = 8000,
# t2v_replicas: int = 4,
# t2v_14b_replicas: int = 4, # Reduced replicas for 14B model due to higher resource requirements
# i2v_replicas: int = 4, # I2V functionality commented out
):
"""Start Ray Serve with independently scalable deployments for each model."""
if not ray.is_initialized():
ray.init()
# Bind model deployments with configurable replica counts
t2v_deps = {}
for model_path, replicas in zip(t2v_model_paths.split(","), t2v_model_replicas.split(",")):
t2v_dep = T2VModelDeployment.options(num_replicas=int(replicas)).bind(model_path, output_path)
t2v_deps[model_path] = t2v_dep
# i2v_dep = I2VModelDeployment.options(num_replicas=i2v_replicas).bind(i2v_model_path, output_path) # I2V functionality commented out
# Ingress
api = FastVideoAPI.bind(t2v_deps) # Removed i2v_dep
serve.run(api, route_prefix="/", name="fast_video")
print(f"Ray Serve backend started at http://{host}:{port}")
for model_path, replicas in zip(t2v_model_paths.split(","), t2v_model_replicas.split(",")):
print(f"T2V Model: {model_path} | Replicas: {replicas}")
# print(f"T2V 14B Model: {t2v_14b_model_path} | Replicas: {t2v_14b_replicas}")
# print(f"I2V Model: {i2v_model_path} | Replicas: {i2v_replicas}") # I2V functionality commented out
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_paths",
type=str,
default="FastVideo/FastWan2.1-T2V-1.3B-Diffusers,FastVideo/FastWan2.1-T2V-14B-Diffusers",
help="Comma separated list of paths to the T2V model(s)")
parser.add_argument("--t2v_model_replicas",
type=str,
default="4,4",
help="Comma separated list of number of replicas for the T2V model(s)")
# parser.add_argument("--t2v_14b_model_path",
# type=str,
# default="FastVideo/FastWan2.1-T2V-14B-Diffusers",
# help="Path to the T2V 14B model")
# parser.add_argument("--i2v_model_path", # I2V functionality commented out
# 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("--t2v_replicas",
# type=int,
# default=4,
# help="Number of replicas for the T2V model deployment")
# parser.add_argument("--t2v_14b_replicas",
# type=int,
# default=4, # Reduced default for 14B model due to higher resource requirements
# help="Number of replicas for the T2V 14B model deployment")
# parser.add_argument("--i2v_replicas", # I2V functionality commented out
# type=int,
# default=4, # Reduced default for 14B model due to higher resource requirements
# help="Number of replicas for the I2V model deployment")
args = parser.parse_args()
split_models = args.t2v_model_paths.split(",")
split_replicas = [int(replica) for replica in args.t2v_model_replicas.split(",")]
assert len(split_models) == len(split_replicas), "Number of models and replicas must match"
assert sum(split_replicas) <= NUM_GPUS, "Total number of replicas must be less than or equal to 16"
for model, replicas in zip(split_models, split_replicas):
assert model in SUPPORTED_MODELS, f"Model {model} not supported"
assert replicas > 0, f"Replicas must be greater than 0"
start_ray_serve(
t2v_model_paths=args.t2v_model_paths,
t2v_model_replicas=args.t2v_model_replicas,
# t2v_14b_model_path=args.t2v_14b_model_path,
# i2v_model_path=args.i2v_model_path, # I2V functionality commented out
output_path=args.output_path,
host=args.host,
port=args.port,
# t2v_replicas=args.t2v_replicas,
# t2v_14b_replicas=args.t2v_14b_replicas,
# i2v_replicas=args.i2v_replicas, # I2V functionality commented out
)
# ---- 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,243 @@
"""
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_paths", args.t2v_model_paths,
"--t2v_model_replicas", args.t2v_model_replicas,
# "--i2v_model_path", args.i2v_model_path, # I2V functionality commented out
"--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_paths", args.t2v_model_paths,
# "--t2v_model_replicas", args.t2v_model_replicas,
# "--i2v_model_path", args.i2v_model_path, # I2V functionality commented out
"--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_paths",
type=str,
default="FastVideo/FastWan2.1-T2V-1.3B-Diffusers,FastVideo/FastWan2.1-T2V-14B-Diffusers",
help="Comma separated list of paths to the T2V model(s)")
parser.add_argument("--t2v_model_replicas",
type=str,
default="4,4",
help="Comma separated list of number of replicas for the T2V model(s)")
# parser.add_argument("--t2v_14b_model_path",
# type=str,
# default="FastVideo/FastWan2.1-T2V-14B-Diffusers",
# help="Path to the T2V 14B model")
# parser.add_argument("--i2v_model_path", # I2V functionality commented out
# 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_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=7861,
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 Models: {args.t2v_model_paths}")
print(f"T2V Model Replicas: {args.t2v_model_replicas}")
# print(f"I2V Model: {args.i2v_model_path}") # I2V functionality commented out
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()
+147
View File
@@ -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!")
@@ -0,0 +1,14 @@
<svg width="200" height="201" viewBox="0 0 200 201" fill="none" xmlns="http://www.w3.org/2000/svg">
<rect width="200" height="200" transform="translate(0 0.691925)" fill="white"/>
<path d="M86.4467 27.9646L64.2778 97.1314H84.6732L91.7672 71.4156L124.445 71.4156L134.898 57.2275L96.201 57.2275L100.635 43.9262H146.844L159.278 27.9646H86.4467Z" fill="#0C0D6F"/>
<path d="M158.423 107.852H131.821L91.6798 160.17L84.5858 107.852H64.4273L73.2948 174.359H95.4637L158.423 107.852Z" fill="#0C0D6F"/>
<path d="M86.4467 27.9646L64.2778 97.1314H84.6732L91.7672 71.4156L124.445 71.4156L134.898 57.2275L96.201 57.2275L100.635 43.9262H146.844L159.278 27.9646H86.4467Z" stroke="#0C0D6F" stroke-width="1.77351"/>
<path d="M158.423 107.852H131.821L91.6798 160.17L84.5858 107.852H64.4273L73.2948 174.359H95.4637L158.423 107.852Z" stroke="#0C0D6F" stroke-width="1.77351"/>
<path d="M158.297 107.923H131.694L166.713 34.6841L97.1109 126.545H118.393L95.3374 174.429L158.297 107.923Z" fill="#FDC717" stroke="#FDC717" stroke-width="1.77351" stroke-miterlimit="16"/>
<path d="M53.6055 107.773L62.473 174.279L67.7935 174.279L58.926 107.772L53.6055 107.773Z" fill="#0C0D6F" stroke="#0C0D6F" stroke-width="1.77351"/>
<path d="M42.9648 107.773L51.8324 174.279L53.6059 174.279L44.7384 107.772L42.9648 107.773Z" fill="#0C0D6F" stroke="#0C0D6F" stroke-width="1.77351"/>
<path d="M32.3232 107.772L41.1908 174.279L42.0775 174.279L33.21 107.772L32.3232 107.772Z" fill="#0C0D6F" stroke="#0C0D6F" stroke-width="0.886754"/>
<path d="M53.6055 97.1314L75.7743 27.9646H81.0949L58.926 97.1314H53.6055Z" fill="#0C0D6F" stroke="#0C0D6F" stroke-width="1.77351"/>
<path d="M42.9648 97.1314L65.1337 27.9646H66.9072L44.7384 97.1314H42.9648Z" fill="#0C0D6F" stroke="#0C0D6F" stroke-width="1.77351"/>
<path d="M32.3232 97.1315L54.4921 27.9646H55.3789L33.21 97.1315H32.3232Z" fill="#0C0D6F" stroke="#0C0D6F" stroke-width="0.886754"/>
</svg>

After

Width:  |  Height:  |  Size: 1.8 KiB

@@ -0,0 +1,7 @@
<svg width="200" height="201" viewBox="0 0 200 201" fill="none" xmlns="http://www.w3.org/2000/svg">
<rect width="200" height="200" transform="translate(0 0.308105)" fill="white"/>
<path d="M49.8511 145.66L78.6319 55.8637H85.5394L56.7585 145.66H49.8511Z" fill="#0C0D6F" stroke="#0C0D6F" stroke-width="2.30244"/>
<path d="M36.0376 145.66L64.8185 55.8637H67.1209L38.3401 145.66H36.0376Z" fill="#0C0D6F" stroke="#0C0D6F" stroke-width="2.30244"/>
<path d="M22.2222 145.66L51.003 55.8637H52.1543L23.3734 145.66H22.2222Z" fill="#0C0D6F" stroke="#0C0D6F" stroke-width="1.15122"/>
<path d="M92.4465 55.8648L63.666 145.66H90.144L99.3538 112.275H144.251L150.007 93.855H105.11L110.866 76.5868H173.032L178.788 55.8648H92.4465Z" fill="#0C0D6F" stroke="#0C0D6F" stroke-width="2.30244"/>
</svg>

After

Width:  |  Height:  |  Size: 779 B

@@ -0,0 +1,8 @@
<svg width="200" height="201" viewBox="0 0 200 201" fill="none" xmlns="http://www.w3.org/2000/svg">
<rect width="200" height="200" transform="translate(0 0.308105)" fill="white"/>
<path d="M49.8511 145.66L78.6319 55.8637H85.5394L56.7585 145.66H49.8511Z" fill="#0C0D6F" stroke="#0C0D6F" stroke-width="2.30244"/>
<path d="M36.0376 145.66L64.8185 55.8637H67.1209L38.3401 145.66H36.0376Z" fill="#0C0D6F" stroke="#0C0D6F" stroke-width="2.30244"/>
<path d="M22.2222 145.66L51.003 55.8637H52.1543L23.3734 145.66H22.2222Z" fill="#0C0D6F" stroke="#0C0D6F" stroke-width="1.15122"/>
<path d="M92.4465 55.8648L63.666 145.66H90.144L99.3538 112.275H144.251L150.007 93.855H105.11L110.866 76.5868H173.032L178.788 55.8648H92.4465Z" fill="#0C0D6F" stroke="#0C0D6F" stroke-width="2.30244"/>
<path d="M131.799 93.8999H104.705L134.343 30.1061L69.483 112.866H91.1583L67.6768 161.635L131.799 93.8999Z" fill="#FDC717" stroke="#FDC717" stroke-width="1.80627" stroke-miterlimit="16"/>
</svg>

After

Width:  |  Height:  |  Size: 966 B

@@ -0,0 +1,14 @@
<svg width="200" height="201" viewBox="0 0 200 201" fill="none" xmlns="http://www.w3.org/2000/svg">
<rect width="200" height="200" transform="translate(0 0.691925)" fill="white"/>
<path d="M89.2368 40.0858L70.8901 97.3273H87.769L93.64 76.0452L120.684 76.0453L129.334 64.3034L97.3093 64.3034L100.979 53.2954H139.22L149.511 40.0858H89.2368Z" fill="#0C0D6F"/>
<path d="M148.804 106.199H126.788L93.5676 149.498L87.6967 106.199H71.0138L78.3525 161.239H96.6991L148.804 106.199Z" fill="#0C0D6F"/>
<path d="M89.2368 40.0858L70.8901 97.3273H87.769L93.64 76.0452L120.684 76.0453L129.334 64.3034L97.3093 64.3034L100.979 53.2954H139.22L149.511 40.0858H89.2368Z" stroke="#0C0D6F" stroke-width="1.46773"/>
<path d="M148.804 106.199H126.788L93.5676 149.498L87.6967 106.199H71.0138L78.3525 161.239H96.6991L148.804 106.199Z" stroke="#0C0D6F" stroke-width="1.46773"/>
<path d="M148.699 106.258H126.683L155.664 45.6468L98.062 121.669H115.675L96.5942 161.298L148.699 106.258Z" fill="#FDC717" stroke="#FDC717" stroke-width="1.46773" stroke-miterlimit="16"/>
<path d="M62.0576 106.134L69.3963 161.174L73.7995 161.174L66.4608 106.134L62.0576 106.134Z" fill="#0C0D6F" stroke="#0C0D6F" stroke-width="1.46773"/>
<path d="M53.2515 106.134L60.5901 161.174L62.0579 161.174L54.7192 106.134L53.2515 106.134Z" fill="#0C0D6F" stroke="#0C0D6F" stroke-width="1.46773"/>
<path d="M44.4443 106.134L51.783 161.174L52.5169 161.174L45.1782 106.134L44.4443 106.134Z" fill="#0C0D6F" stroke="#0C0D6F" stroke-width="0.733866"/>
<path d="M62.0576 97.3273L80.4043 40.0858H84.8075L66.4608 97.3273H62.0576Z" fill="#0C0D6F" stroke="#0C0D6F" stroke-width="1.46773"/>
<path d="M53.2515 97.3273L71.5981 40.0858H73.0658L54.7192 97.3273H53.2515Z" fill="#0C0D6F" stroke="#0C0D6F" stroke-width="1.46773"/>
<path d="M44.4443 97.3274L62.791 40.0859H63.5248L45.1782 97.3274H44.4443Z" fill="#0C0D6F" stroke="#0C0D6F" stroke-width="0.733866"/>
</svg>

After

Width:  |  Height:  |  Size: 1.8 KiB

+18
View File
@@ -0,0 +1,18 @@
<svg width="252" height="105" viewBox="0 0 252 105" fill="none" xmlns="http://www.w3.org/2000/svg">
<path d="M89.4843 55.5457H101.361L87.7028 101H74.638L89.4843 55.5457Z" fill="#356CFF"/>
<path fill-rule="evenodd" clip-rule="evenodd" d="M96.0167 1.00057H112.645L118.583 48.273H104.924L103.737 39.7882H79.9827L67.5117 55.5457H85.3273L43.1638 101H28.3174L22.3789 55.5457H33.6621L38.4129 91.3031L58.604 68.2729H44.3515L96.0167 1.00057ZM100.768 13.1217L87.7028 29.4852H103.143L100.768 13.1217Z" fill="#356CFF"/>
<path d="M37.2252 1.00057L22.3789 48.273H36.0375L40.7884 30.6974L62.6727 30.6974L69.6727 21.0005L43.7576 21.0004L46.7269 11.9096L77.6727 11.9096L86 1.00057L37.2252 1.00057Z" fill="#356CFF"/>
<path fill-rule="evenodd" clip-rule="evenodd" d="M108.488 55.5457L94.2351 101C94.2351 101 105.518 101 120.959 101C136.399 101 144.078 93.0133 148.276 79.788C152.432 68.0157 153.027 55.5457 136.399 55.5457C119.771 55.5457 108.488 55.5457 108.488 55.5457ZM109.081 90.697L116.802 65.8487C116.802 65.8487 120.959 65.8487 132.242 65.8487C143.525 65.8487 137.586 78.5759 135.211 84.0304C133.307 88.4021 127.491 90.697 122.74 90.697C117.989 90.697 109.081 90.697 109.081 90.697Z" fill="#356CFF"/>
<path d="M173.188 1.00056L168.625 11.9096C168.625 11.9096 149.386 11.9092 142.525 11.9095C135.664 11.9098 136.586 20.3944 141.337 20.3944H159.747C168.654 20.3944 166.961 33.6899 163.904 38.5761C160.467 44.0675 157.371 48.273 148.463 48.273L125.188 48.273L124 37.97L147.87 37.97C153.808 37.97 156.184 29.4852 151.433 29.4852H131.836C120.142 29.4852 125.897 1.00043 141.337 1.00043L173.188 1.00056Z" fill="#356CFF"/>
<path d="M179.938 1.00056L175.688 11.9096L191.221 11.9096L179.938 48.273H192.409L203.692 11.9096L219.132 11.9095L223.289 1.00043L179.938 1.00056Z" fill="#356CFF"/>
<path d="M161.341 55.5457H202.845L198.5 65.8487H169.654L167.279 73.7268H188.5L184.749 82.8177H164.31L161.934 90.697H190.251L186.624 101H146.494L161.341 55.5457Z" fill="#356CFF"/>
<path fill-rule="evenodd" clip-rule="evenodd" d="M230.821 54.9391C255.169 54.9391 251.776 67.0602 249.231 77.9692C246.686 88.8783 240.917 101 217.757 101C194.596 101 195.606 88.8783 199.347 77.9692C203.089 67.0602 206.473 54.9391 230.821 54.9391ZM237.948 77.9692C239.984 70.6965 240.917 65.242 228.446 65.242C215.975 65.242 211.818 71.9087 210.037 77.9692C208.255 84.0298 208.255 91.3025 219.538 91.3025C230.821 91.3025 235.911 85.2419 237.948 77.9692Z" fill="#356CFF"/>
<path d="M173.188 1.00056L168.625 11.9096C168.625 11.9096 149.386 11.9092 142.525 11.9095C135.664 11.9098 136.586 20.3944 141.337 20.3944M173.188 1.00056C173.188 1.00056 156.777 1.00043 141.337 1.00043M173.188 1.00056L141.337 1.00043M141.337 20.3944C146.088 20.3944 150.839 20.3944 159.747 20.3944M141.337 20.3944H159.747M159.747 20.3944C168.654 20.3944 166.961 33.6899 163.904 38.5761C160.467 44.0675 157.371 48.273 148.463 48.273M148.463 48.273C139.556 48.273 125.188 48.273 125.188 48.273M148.463 48.273L125.188 48.273M125.188 48.273L124 37.97M124 37.97C124 37.97 141.931 37.97 147.87 37.97M124 37.97L147.87 37.97M147.87 37.97C153.808 37.97 156.184 29.4852 151.433 29.4852M151.433 29.4852C146.682 29.4852 138.962 29.4852 131.836 29.4852M151.433 29.4852H131.836M131.836 29.4852C120.142 29.4852 125.897 1.00043 141.337 1.00043M37.2252 1.00057L22.3789 48.273H36.0375L40.7884 30.6974L62.6727 30.6974L69.6727 21.0005L43.7576 21.0004L46.7269 11.9096L77.6727 11.9096L86 1.00057L37.2252 1.00057ZM96.0167 1.00057H112.645L118.583 48.273H104.924L103.737 39.7882H79.9827L67.5117 55.5457H85.3273L43.1638 101H28.3174L22.3789 55.5457H33.6621L38.4129 91.3031L58.604 68.2729H44.3515L96.0167 1.00057ZM87.7028 29.4852L100.768 13.1217L103.143 29.4852H87.7028ZM89.4843 55.5457H101.361L87.7028 101H74.638L89.4843 55.5457ZM108.488 55.5457L94.2351 101C94.2351 101 105.518 101 120.959 101C136.399 101 144.078 93.0133 148.276 79.788C152.432 68.0157 153.027 55.5457 136.399 55.5457C119.771 55.5457 108.488 55.5457 108.488 55.5457ZM116.802 65.8487L109.081 90.697C109.081 90.697 117.989 90.697 122.74 90.697C127.491 90.697 133.307 88.4021 135.211 84.0304C137.586 78.5759 143.525 65.8487 132.242 65.8487C120.959 65.8487 116.802 65.8487 116.802 65.8487ZM179.938 1.00056L175.688 11.9096L191.221 11.9096L179.938 48.273H192.409L203.692 11.9096L219.132 11.9095L223.289 1.00043L179.938 1.00056ZM161.341 55.5457H202.845L198.5 65.8487H169.654L167.279 73.7268H188.5L184.749 82.8177H164.31L161.934 90.697H190.251L186.624 101H146.494L161.341 55.5457ZM230.821 54.9391C255.169 54.9391 251.776 67.0602 249.231 77.9692C246.686 88.8783 240.917 101 217.757 101C194.596 101 195.606 88.8783 199.347 77.9692C203.089 67.0602 206.473 54.9391 230.821 54.9391ZM228.446 65.242C240.917 65.242 239.984 70.6965 237.948 77.9692C235.911 85.2419 230.821 91.3025 219.538 91.3025C208.255 91.3025 208.255 84.0298 210.037 77.9692C211.818 71.9087 215.975 65.242 228.446 65.242Z" stroke="#356CFF" stroke-width="1.18771"/>
<path d="M15.2524 55.5451L21.191 100.999L24.7541 100.999L18.8156 55.5451L15.2524 55.5451Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.18771"/>
<path d="M8.12646 55.5451L14.065 100.999L15.2527 100.999L9.31417 55.5451L8.12646 55.5451Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.18771"/>
<path d="M1 55.5451L6.93853 100.999L7.53239 100.999L1.59385 55.5451L1 55.5451Z" fill="#356CFF" stroke="#356CFF" stroke-width="0.593853"/>
<path d="M15.2524 48.2724L30.0988 1H33.6619L18.8156 48.2724H15.2524Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.18771"/>
<path d="M8.12646 48.2724L22.9728 1H24.1605L9.31417 48.2724H8.12646Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.18771"/>
<path d="M1 48.2724L15.8463 1H16.4402L1.59385 48.2724H1Z" fill="#356CFF" stroke="#356CFF" stroke-width="0.593853"/>
<path d="M85.3271 55.5457H67.5116L87 12.7363L44.3513 68.2729H58.6038L43.1636 101L85.3271 55.5457Z" fill="#FDC717" stroke="#FDC717" stroke-width="1.18771" stroke-miterlimit="16"/>
</svg>

After

Width:  |  Height:  |  Size: 5.7 KiB

+14
View File
@@ -0,0 +1,14 @@
<svg width="200" height="201" viewBox="0 0 200 201" fill="none" xmlns="http://www.w3.org/2000/svg">
<rect width="200" height="200" transform="translate(0 0.691925)" fill="black"/>
<path d="M86.4467 27.9646L64.2778 97.1314H84.6732L91.7672 71.4156L124.445 71.4156L134.898 57.2275L96.201 57.2275L100.635 43.9262H146.844L159.278 27.9646H86.4467Z" fill="#356CFF"/>
<path d="M158.423 107.852H131.821L91.6798 160.17L84.5858 107.852H64.4273L73.2948 174.359H95.4637L158.423 107.852Z" fill="#356CFF"/>
<path d="M86.4467 27.9646L64.2778 97.1314H84.6732L91.7672 71.4156L124.445 71.4156L134.898 57.2275L96.201 57.2275L100.635 43.9262H146.844L159.278 27.9646H86.4467Z" stroke="#356CFF" stroke-width="1.77351"/>
<path d="M158.423 107.852H131.821L91.6798 160.17L84.5858 107.852H64.4273L73.2948 174.359H95.4637L158.423 107.852Z" stroke="#356CFF" stroke-width="1.77351"/>
<path d="M158.297 107.923H131.694L166.713 34.6841L97.1109 126.545H118.393L95.3374 174.429L158.297 107.923Z" fill="#FDC717" stroke="#FDC717" stroke-width="1.77351" stroke-miterlimit="16"/>
<path d="M53.6055 107.773L62.473 174.279L67.7935 174.279L58.926 107.772L53.6055 107.773Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.77351"/>
<path d="M42.9648 107.773L51.8324 174.279L53.6059 174.279L44.7384 107.772L42.9648 107.773Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.77351"/>
<path d="M32.3232 107.772L41.1908 174.279L42.0775 174.279L33.21 107.772L32.3232 107.772Z" fill="#356CFF" stroke="#356CFF" stroke-width="0.886754"/>
<path d="M53.6055 97.1314L75.7743 27.9646H81.0949L58.926 97.1314H53.6055Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.77351"/>
<path d="M42.9648 97.1314L65.1337 27.9646H66.9072L44.7384 97.1314H42.9648Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.77351"/>
<path d="M32.3232 97.1315L54.4921 27.9646H55.3789L33.21 97.1315H32.3232Z" fill="#356CFF" stroke="#356CFF" stroke-width="0.886754"/>
</svg>

After

Width:  |  Height:  |  Size: 1.8 KiB

@@ -0,0 +1,7 @@
<svg width="200" height="201" viewBox="0 0 200 201" fill="none" xmlns="http://www.w3.org/2000/svg">
<rect width="200" height="200" transform="translate(0 0.308105)" fill="black"/>
<path d="M49.8511 145.66L78.6319 55.8637H85.5394L56.7585 145.66H49.8511Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
<path d="M36.0376 145.66L64.8185 55.8637H67.1209L38.3401 145.66H36.0376Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
<path d="M22.2222 145.66L51.003 55.8637H52.1543L23.3734 145.66H22.2222Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.15122"/>
<path d="M92.4465 55.8648L63.666 145.66H90.144L99.3538 112.275H144.251L150.007 93.855H105.11L110.866 76.5868H173.032L178.788 55.8648H92.4465Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
</svg>

After

Width:  |  Height:  |  Size: 779 B

@@ -0,0 +1,8 @@
<svg width="200" height="201" viewBox="0 0 200 201" fill="none" xmlns="http://www.w3.org/2000/svg">
<rect width="200" height="200" transform="translate(0 0.308105)" fill="black"/>
<path d="M49.8511 145.66L78.6319 55.8637H85.5394L56.7585 145.66H49.8511Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
<path d="M36.0376 145.66L64.8185 55.8637H67.1209L38.3401 145.66H36.0376Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
<path d="M22.2222 145.66L51.003 55.8637H52.1543L23.3734 145.66H22.2222Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.15122"/>
<path d="M92.4465 55.8648L63.666 145.66H90.144L99.3538 112.275H144.251L150.007 93.855H105.11L110.866 76.5868H173.032L178.788 55.8648H92.4465Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
<path d="M131.799 93.8999H104.705L134.343 30.1061L69.483 112.866H91.1583L67.6768 161.635L131.799 93.8999Z" fill="#FDC717" stroke="#FDC717" stroke-width="1.80627" stroke-miterlimit="16"/>
</svg>

After

Width:  |  Height:  |  Size: 966 B

@@ -0,0 +1,7 @@
<svg width="160" height="144" viewBox="0 0 160 144" fill="none" xmlns="http://www.w3.org/2000/svg">
<path d="M28.8511 122.66L57.6319 32.8637H64.5394L35.7585 122.66H28.8511Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
<path d="M15.0376 122.66L43.8185 32.8637H46.1209L17.3401 122.66H15.0376Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
<path d="M1.22217 122.66L30.003 32.8637H31.1543L2.3734 122.66H1.22217Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.15122"/>
<path d="M71.4465 32.8648L42.666 122.66H69.144L78.3538 89.2745H123.251L129.007 70.855H84.1099L89.866 53.5868H152.032L157.788 32.8648H71.4465Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
<path d="M110.799 70.8999H83.7053L113.343 7.10608L48.483 89.8657H70.1583L46.6768 138.635L110.799 70.8999Z" fill="#FDC717" stroke="#FDC717" stroke-width="1.80627" stroke-miterlimit="16"/>
</svg>

After

Width:  |  Height:  |  Size: 885 B

+6
View File
@@ -0,0 +1,6 @@
<svg width="160" height="93" viewBox="0 0 160 93" fill="none" xmlns="http://www.w3.org/2000/svg">
<path d="M28.8511 91.66L57.6319 1.86368H64.5394L35.7585 91.66H28.8511Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
<path d="M15.0376 91.66L43.8185 1.86368H46.1209L17.3401 91.66H15.0376Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
<path d="M1.22217 91.66L30.003 1.86366H31.1543L2.3734 91.66H1.22217Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.15122"/>
<path d="M71.4465 1.86483L42.666 91.6599H69.144L78.3538 58.2746H123.251L129.007 39.855H84.1099L89.866 22.5868H152.032L157.788 1.86483H71.4465Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
</svg>

After

Width:  |  Height:  |  Size: 691 B

+14
View File
@@ -0,0 +1,14 @@
<svg width="200" height="201" viewBox="0 0 200 201" fill="none" xmlns="http://www.w3.org/2000/svg">
<rect width="200" height="200" transform="translate(0 0.691925)" fill="black"/>
<path d="M89.2368 40.0858L70.8901 97.3273H87.769L93.64 76.0452L120.684 76.0453L129.334 64.3034L97.3093 64.3034L100.979 53.2954H139.22L149.511 40.0858H89.2368Z" fill="#356CFF"/>
<path d="M148.804 106.199H126.788L93.5676 149.498L87.6967 106.199H71.0138L78.3525 161.239H96.6991L148.804 106.199Z" fill="#356CFF"/>
<path d="M89.2368 40.0858L70.8901 97.3273H87.769L93.64 76.0452L120.684 76.0453L129.334 64.3034L97.3093 64.3034L100.979 53.2954H139.22L149.511 40.0858H89.2368Z" stroke="#356CFF" stroke-width="1.46773"/>
<path d="M148.804 106.199H126.788L93.5676 149.498L87.6967 106.199H71.0138L78.3525 161.239H96.6991L148.804 106.199Z" stroke="#356CFF" stroke-width="1.46773"/>
<path d="M148.699 106.258H126.683L155.664 45.6468L98.062 121.669H115.675L96.5942 161.298L148.699 106.258Z" fill="#FDC717" stroke="#FDC717" stroke-width="1.46773" stroke-miterlimit="16"/>
<path d="M62.0576 106.134L69.3963 161.174L73.7995 161.174L66.4608 106.134L62.0576 106.134Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.46773"/>
<path d="M53.2515 106.134L60.5901 161.174L62.0579 161.174L54.7192 106.134L53.2515 106.134Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.46773"/>
<path d="M44.4443 106.134L51.783 161.174L52.5169 161.174L45.1782 106.134L44.4443 106.134Z" fill="#356CFF" stroke="#356CFF" stroke-width="0.733866"/>
<path d="M62.0576 97.3273L80.4043 40.0858H84.8075L66.4608 97.3273H62.0576Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.46773"/>
<path d="M53.2515 97.3273L71.5981 40.0858H73.0658L54.7192 97.3273H53.2515Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.46773"/>
<path d="M44.4443 97.3274L62.791 40.0859H63.5248L45.1782 97.3274H44.4443Z" fill="#356CFF" stroke="#356CFF" stroke-width="0.733866"/>
</svg>

After

Width:  |  Height:  |  Size: 1.8 KiB

+13
View File
@@ -0,0 +1,13 @@
<svg width="141" height="153" viewBox="0 0 141 153" fill="none" xmlns="http://www.w3.org/2000/svg">
<path d="M55.4467 0.964615L33.2778 70.1314H53.6732L60.7672 44.4156L93.4453 44.4156L103.898 30.2275L65.201 30.2275L69.6347 16.9262H115.844L128.278 0.964615H55.4467Z" fill="#356CFF"/>
<path d="M127.423 80.852H100.821L60.6798 133.17L53.5858 80.852H33.4273L42.2948 147.359H64.4637L127.423 80.852Z" fill="#356CFF"/>
<path d="M55.4467 0.964615L33.2778 70.1314H53.6732L60.7672 44.4156L93.4453 44.4156L103.898 30.2275L65.201 30.2275L69.6347 16.9262H115.844L128.278 0.964615H55.4467Z" stroke="#356CFF" stroke-width="1.77351"/>
<path d="M127.423 80.852H100.821L60.6798 133.17L53.5858 80.852H33.4273L42.2948 147.359H64.4637L127.423 80.852Z" stroke="#356CFF" stroke-width="1.77351"/>
<path d="M127.297 80.9227H100.694L135.713 7.68411L66.1109 99.5445H87.393L64.3374 147.429L127.297 80.9227Z" fill="#FDC717" stroke="#FDC717" stroke-width="1.77351" stroke-miterlimit="16"/>
<path d="M22.6055 80.7725L31.473 147.279L36.7935 147.279L27.926 80.7725L22.6055 80.7725Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.77351"/>
<path d="M11.9648 80.7725L20.8324 147.279L22.6059 147.279L13.7384 80.7725L11.9648 80.7725Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.77351"/>
<path d="M1.32324 80.7725L10.1908 147.279L11.0775 147.279L2.21 80.7725L1.32324 80.7725Z" fill="#356CFF" stroke="#356CFF" stroke-width="0.886754"/>
<path d="M22.6055 70.1314L44.7743 0.964615H50.0949L27.926 70.1314H22.6055Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.77351"/>
<path d="M11.9648 70.1314L34.1337 0.964615H35.9072L13.7384 70.1314H11.9648Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.77351"/>
<path d="M1.32324 70.1315L23.4921 0.964645H24.3789L2.21 70.1315H1.32324Z" fill="#356CFF" stroke="#356CFF" stroke-width="0.886754"/>
</svg>

After

Width:  |  Height:  |  Size: 1.8 KiB

+3
View File
@@ -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
+5 -1
View File
@@ -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
}
+23
View File
@@ -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
+11 -2
View File
@@ -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
}
+16 -1
View File
@@ -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
+10 -3
View File
@@ -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}
@@ -303,8 +307,9 @@ class VideoGenerator:
# Run inference
start_time = time.perf_counter()
output_batch = self.executor.execute_forward(batch, fastvideo_args)
samples = output_batch
output_batch = self.executor.execute_forward(batch, fastvideo_args)
samples = output_batch.output
logging_info = output_batch.logging_info
gen_time = time.perf_counter() - start_time
logger.info("Generated successfully in %.2f seconds", gen_time)
@@ -334,9 +339,11 @@ class VideoGenerator:
else:
return {
"samples": samples,
"frames": frames,
"prompts": prompt,
"size": (target_height, target_width, batch.num_frames),
"generation_time": gen_time
"generation_time": gen_time,
"logging_info": logging_info,
}
def set_lora_adapter(self,
+1
View File
@@ -82,6 +82,7 @@ class WanTimeTextImageEmbedding(nn.Module):
encoder_hidden_states: torch.Tensor,
encoder_hidden_states_image: torch.Tensor | None = None,
):
logger.info(f"WTF timestep shape: {timestep.shape}")
temb = self.time_embedder(timestep)
timestep_proj = self.time_modulation(temb)
@@ -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",
@@ -16,6 +16,41 @@ import torch
from fastvideo.attention import AttentionMetadata
from fastvideo.configs.sample.teacache import TeaCacheParams, WanTeaCacheParams
from collections import OrderedDict
from typing import Any, Dict, List
import time
class PipelineLoggingInfo:
"""Simple approach using OrderedDict to track stage metrics."""
def __init__(self):
# OrderedDict preserves insertion order and allows easy access
self.stages: OrderedDict[str, Dict[str, Any]] = OrderedDict()
def add_stage_execution_time(self, stage_name: str, execution_time: float):
"""Add execution time for a stage."""
if stage_name not in self.stages:
self.stages[stage_name] = {}
self.stages[stage_name]['execution_time'] = execution_time
self.stages[stage_name]['timestamp'] = time.time()
def add_stage_metric(self, stage_name: str, metric_name: str, value: Any):
"""Add any metric for a stage."""
if stage_name not in self.stages:
self.stages[stage_name] = {}
self.stages[stage_name][metric_name] = value
def get_stage_info(self, stage_name: str) -> Dict[str, Any]:
"""Get all info for a specific stage."""
return self.stages.get(stage_name, {})
def get_execution_order(self) -> List[str]:
"""Get stages in execution order."""
return list(self.stages.keys())
def get_total_execution_time(self) -> float:
"""Get total pipeline execution time."""
return sum(stage.get('execution_time', 0) for stage in self.stages.values())
@dataclass
@@ -128,6 +163,9 @@ class ForwardBatch:
# VSA parameters
VSA_sparsity: float = 0.0
# Logging info
logging_info: PipelineLoggingInfo = field(default_factory=PipelineLoggingInfo)
def __post_init__(self):
"""Initialize dependent fields after dataclass initialization."""
+1
View File
@@ -21,6 +21,7 @@ _PIPELINE_NAME_TO_ARCHITECTURE_NAME: dict[str, str] = {
"WanPipeline": "wan",
"WanDMDPipeline": "wan",
"WanImageToVideoPipeline": "wan",
"WanDMDPipeline": "wan",
"StepVideoPipeline": "stepvideo",
"HunyuanVideoPipeline": "hunyuan",
}
+1
View File
@@ -155,6 +155,7 @@ class PipelineStage(ABC):
execution_time = time.perf_counter() - start_time
logger.info("[%s] Execution completed in %s ms", stage_name,
execution_time * 1000)
batch.logging_info.add_stage_execution_time(stage_name, execution_time)
except Exception as e:
execution_time = time.perf_counter() - start_time
logger.error("[%s] Error during execution after %s ms: %s",
+70 -3
View File
@@ -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,42 @@ 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).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)
timestep = temp_ts
logger.info(f"after ti2v timestep: {timestep}")
t = timestep
# 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] *
@@ -303,6 +366,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,
+63
View File
@@ -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
+5 -1
View File
@@ -21,6 +21,7 @@ from fastvideo.pipelines import ForwardBatch, build_pipeline
from fastvideo.platforms import current_platform
from fastvideo.utils import (get_exception_traceback,
kill_itself_when_parent_died)
import fastvideo.envs as envs
logger = init_logger(__name__)
@@ -140,7 +141,10 @@ class Worker:
fastvideo_args = recv_rpc['kwargs']['fastvideo_args']
output_batch = self.execute_forward(forward_batch,
fastvideo_args)
self.pipe.send({"output_batch": output_batch.output.cpu()})
logging_info = None
if envs.FASTVIDEO_STAGE_LOGGING:
logging_info = output_batch.logging_info
self.pipe.send({"output_batch": output_batch.output.cpu(), "logging_info": logging_info})
elif method_name == 'set_lora_adapter':
lora_nickname = recv_rpc['kwargs']['lora_nickname']
lora_path = recv_rpc['kwargs']['lora_path']
+18 -2
View File
@@ -15,6 +15,7 @@ from fastvideo.logger import init_logger
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.worker.executor import Executor
from fastvideo.worker.gpu_worker import run_worker_process
import fastvideo.envs as envs
logger = init_logger(__name__)
@@ -40,7 +41,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
@@ -80,7 +82,21 @@ class MultiprocExecutor(Executor):
"forward_batch": forward_batch,
"fastvideo_args": fastvideo_args
})
return cast(ForwardBatch, responses[0]["output_batch"])
output = responses[0]["output_batch"]
logging_info = None
if envs.FASTVIDEO_STAGE_LOGGING:
logging_info = responses[0]["logging_info"]
else:
logging_info = None
result_batch = ForwardBatch(
data_type=forward_batch.data_type,
output=output,
logging_info=logging_info
)
return result_batch
def set_lora_adapter(self,
lora_nickname: str,