Compare commits
3
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
03ac81165f | ||
|
|
963fe90a57 | ||
|
|
dc9a1b9ec8 |
@@ -1,38 +1,46 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=t2v
|
||||
#SBATCH --job-name=wl_t2v
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --nodes=1
|
||||
#SBATCH --ntasks=1
|
||||
#SBATCH --nodes=4
|
||||
#SBATCH --ntasks=4
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:1
|
||||
#SBATCH --gres=gpu:8
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=dmd_t2v_output/t2v_%j.out
|
||||
#SBATCH --error=dmd_t2v_output/t2v_%j.err
|
||||
#SBATCH --output=dmd_t2v_output/sf.out
|
||||
#SBATCH --error=dmd_t2v_output/sf.err
|
||||
#SBATCH --exclusive
|
||||
set -e -x
|
||||
|
||||
|
||||
# Environment Setup
|
||||
source ~/conda/miniconda/bin/activate
|
||||
conda activate will-fv
|
||||
|
||||
# Basic Info
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export TORCH_NCCL_ENABLE_MONITORING=0
|
||||
# different cache dir for different processes
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
|
||||
export MASTER_PORT=29501
|
||||
export MASTER_PORT=29503
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
export WANDB_API_KEY="50632ebd88ffd970521cec9ab4a1a2d7e85bfc45"
|
||||
export WANDB_API_KEY="8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde"
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=offline
|
||||
export WANDB_MODE=online
|
||||
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
|
||||
# Configs
|
||||
NUM_GPUS=4
|
||||
NUM_GPUS=8
|
||||
|
||||
# Model paths for Self-Forcing DMD distillation:
|
||||
GENERATOR_MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
|
||||
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-14B-Diffusers" # Teacher model
|
||||
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Teacher model
|
||||
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Critic model
|
||||
|
||||
DATA_DIR="data/mixkit-64_processed/Node_0_GPU_1_File_1/combined_parquet_dataset"
|
||||
VALIDATION_DATASET_FILE="data/mixkit-64_processed/validation.json"
|
||||
# DATA_DIR="data/test-text-preprocessing/Node_0_GPU_1_File_1/combined_parquet_dataset/"
|
||||
DATA_DIR="/mnt/weka/home/hao.zhang/matthew/FastVideo/data/test-text-preprocessing"
|
||||
VALIDATION_DATASET_FILE="data/crush-smol-single_processed_t2v/validation.json"
|
||||
VALIDATION_DATASET_FILE="/mnt/weka/home/hao.zhang/wl/FastVideo/examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free/validation_64.json"
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
@@ -40,7 +48,9 @@ VALIDATION_DATASET_FILE="data/mixkit-64_processed/validation.json"
|
||||
training_args=(
|
||||
--tracker_project_name SFwan_t2v_distill_self_forcing_dmd # Updated for self-forcing DMD
|
||||
--output_dir "/mnt/sharefs/users/hao.zhang/SFwan_t2v_finetune"
|
||||
--max_train_steps 4000
|
||||
# --use_sf_wan
|
||||
# --sf_ode_init_path "checkpoints/ode_init.pt"
|
||||
--max_train_steps 6000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
@@ -58,11 +68,11 @@ training_args=(
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS # 64
|
||||
--num_gpus 32 # 64
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1 # 64
|
||||
--hsdp_shard_dim $NUM_GPUS
|
||||
--hsdp_replicate_dim 32 # 64
|
||||
--hsdp_shard_dim 1
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
@@ -84,7 +94,7 @@ dataset_args=(
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 100
|
||||
--validation_steps 10
|
||||
--validation_sampling_steps "4"
|
||||
--validation_guidance_scale "6.0" # not used for dmd inference
|
||||
)
|
||||
@@ -112,6 +122,7 @@ miscellaneous_args=(
|
||||
--ema_decay 0.99
|
||||
--ema_start_step 100
|
||||
--init_weights_from_safetensors "/mnt/weka/home/hao.zhang/wl/Self-Forcing/diffusers_ode_init/model.safetensors"
|
||||
# --init_weights_from_safetensors "/mnt/weka/home/hao.zhang/wl/FastVideo2/warp_vidprom_8b16k_test_warp_1e-5/checkpoint-2000/transformer/diffusion_pytorch_model.safetensors"
|
||||
)
|
||||
|
||||
# Self-forcing DMD arguments
|
||||
@@ -129,15 +140,17 @@ dmd_args=(
|
||||
# Self-forcing specific arguments
|
||||
self_forcing_args=(
|
||||
--independent_first_frame False # Whether to treat first frame independently
|
||||
--same_step_across_blocks True # Whether to use same denoising step across all blocks
|
||||
--same_step_across_blocks False # Whether to use same denoising step across all blocks
|
||||
--last_step_only False # Whether to only use the last denoising step
|
||||
--context_noise 0 # Amount of noise to add during context caching (0 = no noise)
|
||||
--validate_cache_structure False # Set to True for debugging KV cache issues
|
||||
)
|
||||
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--master_port $MASTER_PORT \
|
||||
srun torchrun \
|
||||
--nnodes $SLURM_JOB_NUM_NODES \
|
||||
--node_rank $SLURM_PROCID \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
fastvideo/training/wan_self_forcing_distillation_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
|
||||
@@ -55,462 +55,6 @@
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In an industrial setting, a person leans casually against a railing, exuding a sense of confidence and composure. They are wearing a striking outfit, consisting of a vibrant, patterned jacket over a simple white crop top, creating a bold contrast. The atmosphere is infused with warm, ambient lighting that casts soft shadows on the concrete walls and metallic surfaces. Intricate wiring and pipes form an intricate backdrop, enhancing the urban aesthetic. Their relaxed posture and direct, engaging gaze suggest a sense of ease in this industrial environment. This scene encapsulates a blend of modern fashion and gritty, urban architecture, creating a visually compelling narrative.",
|
||||
"video_path": "Fashion/mixkit-portrait-of-a-hipster-woman-walking-down-a-stairs-1297_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A man is energetically stretching in an open-air setting, surrounded by rows of vibrant red seats that suggest an amphitheater or outdoor venue. He wears a sleeveless black shirt layered with a hooded vest, emphasizing his athletic build as he engages in a warm-up routine. Behind him, the striking modern architecture of the building features geometric panels, with large sections of glass and overlapping metallic beams creating a dynamic backdrop. The scene captures the contrast between his focused movements and the static, bold design of the structure, while the surrounding greenery adds a touch of nature to the environment. The overall atmosphere is one of preparation and anticipation, with the man appearing determined and ready for an upcoming event or performance.",
|
||||
"video_path": "Sport/mixkit-man-doing-arm-stretches-595_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A young woman is seated on the floor in front of a plush, beige tufted couch, fully engrossed in sorting through a stack of papers. Her dark hair falls loosely past her shoulders, and she wears a green plaid shirt, contributing to the casual yet focused atmosphere. She gently places the papers onto a small round white table, occasionally lifting individual sheets to examine them more closely. Her expression shifts subtly, reflecting concentration and contemplation as she processes the information on the pages. Two small, round nested tables hold her documents, along with a small plant in a gray pot, adding a touch of greenery to the scene. The background features a dark paneled wall, creating a contrasting backdrop for the light-colored furniture. The setting is tranquil and organized, the couch and tables arranged symmetrically, conveying a sense of harmony. A calculator rests on the smaller table, hinting at a task involving calculations or budgeting.",
|
||||
"video_path": "Woman/mixkit-frustrated-woman-throws-paperwork-on-the-floor-4526_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A heavily rusted metal gate stands firmly locked, with two vertical bars joined by a thick, old chain that loops elegantly around them. The chain's texture is coarse and rugged, its surface reflecting varying shades of orange and brown, indicative of years exposed to the elements. At the heart of the chain, a black iron padlock, slightly worn yet imposing, secures the gate, its curves and edges smooth against the aged links. The gate's metalwork is outlined by a backdrop of soft, blurred greenery, suggesting a serene and isolated location beyond the barrier. Tall trees rise in the distance, their trunks and leaves creating a lush, forest-like setting that contrasts with the gate's severe rust. A pathway leads away from the gate, its surface uneven with patches of moss and weathered stone visible in the soft focus, inviting yet inaccessible. The ambiance is quiet and mysterious, with a sense of abandonment hanging subtly in the air, evoking curiosity about what lies beyond. Shadows play across the gate, cast by branches swaying gently in the breeze, adding to the dynamic interaction of light and texture. This scene, rich in detail and atmosphere, captures the viewer's imagination, evoking both the allure of the forbidden and the beauty of decay.",
|
||||
"video_path": "forest/mixkit-rusty-fence-with-a-chain-of-a-property-in-nature-5294_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In a serene and softly lit yoga studio, three individuals engage in a yoga session, each performing an upward-facing stretch. The central figure is a woman with shoulder-length brown hair, dressed in a light cropped top and green leggings, her posture reflecting grace and concentration. To her right, another participant, a woman in a purple outfit, mirrors the pose with equal poise. On her left, a person with a bun focuses intently, supported slightly by yoga blocks beneath their hands. The warm-colored wooden floor contrasts soothingly with the soft pastel mural on the back wall, featuring an abstract design and partial visage of a serene face. Natural light floods the space from a large window on the right, where lush greens peek through, adding an element of tranquility. In the corner of the room, a collection of meditation instruments, including a gong and a Buddha statue, subtly frame the peaceful setting. The mood is calm yet focused, as all three participants are deeply engaged in their practice. The scene combines elements of balance, harmony, and a shared journey towards mindfulness. This depiction captures the essence of a yoga session that blends personal growth with collective experience.",
|
||||
"video_path": "People/mixkit-small-group-of-people-doing-yoga-together-43730_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In the deep blue expanse of the ocean, two dolphins glide effortlessly, their sleek bodies reflecting the sunlight filtering through the water. The prominent shadows and caustics create a shimmering effect on their skin, capturing the beauty of their natural habitat. Each dolphin moves with a fluid grace, occasionally interacting with gentle nudges, showcasing their playful and social nature. The scene is vibrant and dynamic, with the clear blue background accentuating the dolphins' movements, making it an ideal subject for AI recreation.",
|
||||
"video_path": "sea/mixkit-dolphins-underwater-4133_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In the video, a young woman stands against a vibrant graffiti-covered wall, deeply engrossed in her smartphone. Her expression reflects a mix of focus and subtle satisfaction as she interacts with the screen. She wears a black floral-patterned top, which contrasts with the bright, abstract shapes and bold colors of the mural behind her. As she continues to engage with her phone, a series of like count notifications appear on the screen, indicating a growing online appreciation. The wall behind her features a striking mix of geometric and organic shapes, including swirls of teal, orange, and black, with large humanoid figures in a pop-art style. Her long, light-brown hair frames her face, adding a calm, composed aura amidst the lively backdrop. The video captures a blend of contemporary digital interaction and expressive urban art, creating a dynamic yet harmonious scene.",
|
||||
"video_path": "Girl/mixkit-girl-looking-at-the-likes-in-her-post-4914_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A young mother and her baby sit comfortably on a bed, surrounded by an inviting, cozy atmosphere. The woman, wearing a sleeveless top and jeans, is gently engaging with the baby, who is dressed in an adorable animal-print onesie. The child is seated on the bed with colorful toys scattered around, including a plush toy and a board book. The warm glow from a hanging lamp casts a soft light on them, enhancing the serene environment. Pillows are propped up against the headboard, providing a cushioned backdrop as the mother leans slightly over to interact with the baby. A small bottle is visible beside her, suggesting a nurturing setting. Her hand gestures animatedly as she holds up a soft, white cushion with red and blue accents, likely stimulating the baby\u2019s curiosity. Their shared moment is filled with affection and joy, a perfect snapshot of familial bonding.",
|
||||
"video_path": "Baby/mixkit-loving-mother-and-her-baby-playing-with-soft-toys-49966_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A young girl with long brown hair sits at a round wooden table, engrossed in working on her laptop. The laptop screen is a vivid green, suggesting a green screen effect is in use. To her left, a doll dressed in a yellow and white outfit is casually laid on top of some books, adding a playful and innocent touch to the scene. The setting is cozy, with sheer curtains in the background allowing soft natural light to spill into the room. The girl's posture and focused attention on the laptop suggest she is either playing a game or learning something new. This serene and domestic atmosphere is complemented by the slight blur of a dark couch in the foreground, framing the focused activity of the child.",
|
||||
"video_path": "Girl/mixkit-little-girl-doing-homework-on-a-laptop-4757_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "An expansive view of a calm bay reveals a fleet of sailboats, each anchored in a regimented line stretching toward the horizon. The water is a serene blue, reflecting the soft hues of the early morning sky. A gentle breeze is indicated by the subtle ripples trailing behind the boats, while a single, larger vessel cuts a distinct path, leaving a graceful wake in its journey to the open sea. On one side, a cluster of modern high-rise buildings stands, contrasting against the natural simplicity of the water, suggesting a blend of urban and marine life. The distant shoreline is barely visible, softened by the atmospheric perspective, giving a sense of endless waters meeting the sky. The overall mood is peaceful and orderly, with the boats appearing almost as sentinels guarding the expanse of the tranquil bay.",
|
||||
"video_path": "beach/mixkit-flying-backwards-over-the-sea-near-a-coast-50187_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In the video, a person is standing in the center of a dark, featureless space, illuminated by a spotlight that emphasizes their presence. The individual is dressed in a traditional martial arts uniform, known as a gi, which is predominantly white with a black belt tied around the waist, indicating a high level of expertise. The background remains pitch black, creating a stark contrast with the brightly lit figure, ensuring complete focus on them. The person's expression is serious and focused, reflecting a deep sense of discipline and concentration. Their hands move gracefully, transitioning through various martial arts stances, demonstrating practiced skill and fluidity. The uniform's crisp fabric folds and subtly reflects the light, further highlighting each precise movement. Despite the simplicity of the environment, the scene is dynamic, with each motion capturing the essence of martial arts practice. The video effectively conveys a sense of calm strength and mastery, making it ideal for an AI to recreate with attention to posture, lighting, and attire.",
|
||||
"video_path": "Sport/mixkit-karate-fighter-bowing-to-the-front-49706_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In a dimly lit room bathed in a mix of neon purple and blue lights, a focused individual is seated in a gaming chair. She wears a white hoodie and large headphones with cat ears that glow softly, creating a striking silhouette. Her hands rest on a keyboard, typing swiftly as she concentrates intently on the screen in front of her. The atmosphere exudes a sense of intensity and immersion, with the soft-colored lighting enhancing the futuristic vibe. Her long hair cascades down her shoulders, adding a touch of elegance to the otherwise tech-centric setting. The overall scene captures the essence of a dedicated gamer deeply engaged in her virtual world.",
|
||||
"video_path": "earth/mixkit-a-young-woman-wearing-headphones-with-rgb-lights-suddenly-gets-51621_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "Inside a dimly-lit bus, five individuals are seated along the rows of worn seats, each subtly illuminated by the colorful lights emanating from overhead. On the left, a woman sits with a relaxed posture, her curly hair accented by a patterned scarf, wearing a plaid outfit paired with bright neon socks. Next to her, a person clad in a denim jacket appears deep in thought, resting their head on a hand. Further back, another figure in a bucket hat and oversized yellow attire gazes across the aisle, evoking a sense of introspection. The atmosphere is enriched by the soft glow of red and green lights, bathing the bus interior in an almost surreal ambiance, creating a compelling tableau of urban life.",
|
||||
"video_path": "Music/mixkit-conceptual-urban-fashion-42581_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "An aerial view captures two tennis players on a court, with one dressed in white on the left and another in red on the right. They are mid-game, each poised for action with rackets in hand, accentuated by their strategic positioning at opposite baselines. The court itself is a stark, deep blue, bordered by the vibrant green of the surrounding area, with a dark central net dividing the space. Long shadows stretch dramatically across the ground, suggesting a late afternoon setting. The subtly textured surface of the court contrasts with the crisp, white lines marking its boundaries and sections. This scene creates a vivid, balanced composition, highlighting both the competitive tension and serene atmosphere of the game.",
|
||||
"video_path": "People/mixkit-two-people-playing-tennis-aerial-view-880_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In a vibrant, dreamlike setting, a lone figure moves energetically against a backdrop of deep blue and purple hues, casting emotive shadows that ripple with dynamic motion. The figure, almost obscured by a smeared effect, suggests a rhythmic dance or a passionate performance, arms blurred as they sweep through colorful, streaked lighting. A neon glow accentuates their form, particularly highlighting the face which is abstractly illuminated in bursts of orange and red, suggesting intense emotional expression. The scene is dominated by two primary elements \u2013 the figure\u2019s motion and the dramatic lighting, creating a synergy of human emotion and visual spectacle. Swirling trails of light seem to intertwine with the figure, like a visual symphony of movement and color that floods the space. The lighting changes, casting intricate patterns on the figure and the surrounding space, giving the impression of a kaleidoscope in motion. Despite the blurred and abstract portrayal, there is a sense of focus conveyed through the figure\u2019s intent movements, akin to a conductor orchestrating a visual and auditory performance. The environment resonates with an electric energy, suggesting a seamless fusion of art and technology. As the visual drama unfolds, the scene invites viewers to lose themselves in the abstract dance and the play of vivid luminance.",
|
||||
"video_path": "Music/mixkit-dancer-dancing-with-a-light-bar-in-his-hands-42221_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In a brightly lit studio, a photographer wearing a denim jacket focuses intently, capturing shots with a professional camera. Facing him, a model stands gracefully, adjusting her long, flowing hair with delicate movements. The scene is characterized by strong contrasts; the model's soft pink attire and gentle gestures complement the rugged, precise demeanor of the photographer. Positioned against a minimalist backdrop, the pair work seamlessly, with the camera\u2019s lens pointed directly at the model, capturing her elegance. The soft, diffused lighting casts a gentle glow on both subjects, creating an airy and ethereal atmosphere perfect for a high-fashion photo shoot.",
|
||||
"video_path": "Fashion/mixkit-professional-photo-session-with-a-young-female-model-41621_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "The video showcases a serene, expansive landscape covered with a variety of trees dotting the hills. The hills gently slope across the frame, with patches of dry grass contrasting against the lush green foliage. Tall trees with dense canopies stand elegantly, casting soft shadows on the ground below. The sunlight bathes the entire scene, highlighting the varied textures of the leaves and terrain. Gaps between the trees reveal a narrow dirt path meandering through the hills, suggesting a sense of quiet solitude. The undulating hills extend into the distance, creating depth and a calming sense of vast space. The verdant hues of the leaves contrast with the earthy tones of the hills, enhancing the visual richness. In the background, a faint outline of distant hills can be seen, blurred softly by the atmospheric perspective. This tranquil setting could be efficiently recreated in a virtual environment by focusing on its layered composition, color palette, and natural textures.",
|
||||
"video_path": "forest/mixkit-aerial-panorama-of-a-sunny-mountain-landscape-40846_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A bustling ski slope comes alive with skiers descending a pristine, snow-covered hill, surrounded by towering, snow-draped evergreens. Several figures stand atop the slope, silhouetted against a clear blue sky, preparing to embark on their ski run. The chair lift on the right continuously drops off eager adventurers, adding to the excitement at the hilltop. Each skier, clad in colorful winter gear, carves distinct paths into the textured snow as they weave their way down. The interplay of sunlight and shadows accentuates the myriad tracks etched into the slope, creating a dynamic visual rhythm. The scene captures a vibrant winter wonderland, full of action and the thrill of a perfect ski day.",
|
||||
"video_path": "Car/mixkit-skiers-on-a-snowy-slope-3327_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "The scene unfolds within a dimly lit bus, where three young individuals are seated, each absorbed in their unique world. To the left, a person with tied-back hair rests their head on their hand, dressed casually in a jacket and jeans, projecting a relaxed demeanor. Central to the frame is another individual, sitting upright with intense focus, donning a plaid blazer and oversize hoops, enhancing their confident presence. The muted green and red lighting casts an atmospheric glow, adding depth and intrigue to the setting. On the right, a person in a bucket hat and striped shirt leans back, appearing contemplative as they adjust their hat with a nonchalant gesture. The interplay of light and shadow highlights their expressions, creating an intimate and cinematic ambiance. Together, these figures form a cohesive tableau, capturing a moment of introspection amid a bustling yet serene urban environment.",
|
||||
"video_path": "City/mixkit-three-models-posing-to-the-lens-while-on-board-a-42575_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A determined climber is scaling a massive rock face, showcasing exceptional strength and skill. The person, clad in a teal shirt and dark pants, climbs with precision, their movements measured and deliberate. They are secured by climbing gear, which includes ropes and a harness, emphasizing their commitment to safety. The rugged texture of the sandy-colored rock provides an imposing backdrop, adding drama and scale to the climb. In the distance, other large rock formations and sparse vegetation can be seen under a bright, overcast sky, contributing to the natural and adventurous atmosphere. The scene captures a moment of focus and challenge, highlighting the climber's tenacity and the breathtaking environment.",
|
||||
"video_path": "Sport/mixkit-alpinist-climbing-a-huge-rock-in-a-desert-43306_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A woman stands confidently in front of a large array of solar panels, her navy blue jumpsuit contrasting against the lush green grass beneath her feet. Her expression is calm and focused, eyes facing directly ahead, suggesting a deep connection to the subject matter\u2014renewable energy. The sunlight bathes the scene in warm hues, casting gentle shadows and highlighting the geometric precision of the solar panels' grid-like structure. The background reveals a blend of nature and technology, as the panels are anchored on a grassy slope with foliage on the left side of the frame. This composition captures a harmonious blend of human innovation and environmental consciousness, accentuated by the serene outdoor setting.",
|
||||
"video_path": "Business/mixkit-woman-standing-in-front-of-a-solar-panel-4880_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In the video, two people are working at a wooden desk, using an iMac computer. One person, wearing a white knit sweater, is using the apple wireless mouse with their right hand, while their left hand rests on the sleek white keyboard. Their movements are smooth yet intentional, suggesting they are focused on a task on the computer screen. The monitor displays a well-organized array of files and folders, hinting at a task that involves detailed organization or detailed data navigation. The second person, only subtly visible, sits closely by and appears to observe or assist, creating a collaborative atmosphere. Their presence adds a quiet dynamic to the scene, as if they are ready to provide input or guidance. Sticky notes with handwritten notes are attached to the monitor\u2019s stand, adding a touch of personal organization amidst the digital workspace. The focus on the keyboard and mouse emphasizes a streamlined workflow, indicative of a productive work environment. The overall ambiance is calm and focuses on teamwork, technology, and efficient workspace management.",
|
||||
"video_path": "People/mixkit-person-with-glasses-working-on-a-desktop-computer-3248_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A man stands in front of a modern glass facade, taking off a dark hoodie to reveal his gray tank top underneath. His arms are lifted high as he maneuvers the hoodie over his head, showcasing a fluid motion that conveys a sense of calm and routine. The lighting highlights the contours of his muscles, emphasizing a combination of strength and quiet determination. Behind him, the reflective surface of the glass panels provides a subtle backdrop, enhancing the focus on his focused and serene demeanor.",
|
||||
"video_path": "Sport/mixkit-man-puts-on-sleeveless-hoodie-603_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "The video displays a captivating dance of fiery orange flames against a stark black background, creating an intense visual contrast. The flames twist and intertwine, forming symmetrical, swirling patterns that expand and contract rhythmically across the frame. Each fiery tendril seems to be alive, moving with an almost hypnotic fluidity that captures the viewer's attention. The illumination from the flames casts subtle shadows, enhancing the depth and texture of the scene. Overall, the dynamic movement and vibrant color palette create an atmosphere of both beauty and power.",
|
||||
"video_path": "fire/mixkit-two-orange-flames-on-black-background-685_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In this scene, a person is seated in a dimly lit room, possibly a recording studio, holding several drumsticks in their hands. The individual's face is partially obscured by sunglasses, adding a touch of mystery to their demeanor. They are wearing a colorful, patterned shirt with a mix of orange and blue tones that stands out against the darker background. The person appears focused and engaged with the drumsticks, their hands prominently displayed. The ambient light casts warm, soft shadows, emphasizing the texture and colors of their shirt and the wooden drumsticks. The room features wooden paneling, which complements the overall cozy, music-centric setting of the scene. The use of perspective centers on the drumsticks, highlighting the importance of rhythm and music in the captured moment.",
|
||||
"video_path": "Music/mixkit-drummer-stretching-before-playing-42783_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A man is casually sitting on a sofa, engrossed in his meal and entertainment. He is holding a TV remote in one hand while reaching for food with the other, indicating a laid-back, comfortable evening. The table before him is filled with takeout containers, revealing a variety of appetizers and dishes, suggestive of a casual dining experience at home. The background is defined by colorful patterned cushions, adding a cozy, homey feel to the scene. Warm, ambient lighting highlights the relaxed atmosphere, casting soft shadows that contribute to the intimate setting. In this moment, he takes a bite of a sandwich, comfortably balancing his attention between food and whatever is playing on the screen.",
|
||||
"video_path": "Man/mixkit-man-watching-tv-and-eating-fast-food-26089_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "The scene opens to a breathtaking view of a tranquil ocean horizon at dusk, displaying a vibrant tapestry of oranges, pinks, and purples as the sun sets. In the foreground, tall, swaying palm trees frame the scene, their silhouettes stark against the colorful sky. The ocean itself shimmers with reflections of the sunset, creating a peaceful, almost ethereal atmosphere. A small boat can be seen in the distance, centered on the horizon, adding a sense of scale and solitude to the scene. The waves gently lap the shore, creating faint patterns on the sandy beach, which stretches across the foreground. Above, the sky is dotted with scattered clouds that catch the last light of the day, enhancing the drama and beauty of the scene. The overall mood is serene and contemplative, capturing a perfect moment of nature\u2019s grandeur.",
|
||||
"video_path": "beach/mixkit-sunset-with-sailing-boats-2166_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A man sits hunched on a couch, the weight of emotions clearly visible on his posture. He wears a simple, gray t-shirt, and his head is bowed, resting in his hands, which cover most of his face, obscuring his features. The gentle light filtering through sheer curtains in the background casts a soft glow upon him, emphasizing the contrast between his static form and the hazy brightness behind. His elbows rest upon his knees, suggesting a posture of deep contemplation or distress. The simplicity of the room, with its muted colors, highlights the focus on the man's internal struggle. Delicate detailing on the fabric of his shirt adds texture, enhancing the scene's realism. Subtle changes in the natural light indicate the passage of time, as the man remains unmoving, absorbed in thought. This intimate moment captures a profound vulnerability, making the scene universally relatable and poignant.",
|
||||
"video_path": "Man/mixkit-worried-and-sad-man-with-his-head-down-4701_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A pair of hands, belonging to an unseen figure, carefully unrolls a large sheet of crisp, white paper on a dark wooden table. The lighting is warm, casting a gentle glow that highlights the textures of the paper and the wood grain of the table. As the paper unfurls, the edges reveal the faint beginnings of a colorful map printed on its surface. The arms, clad in a casual gray T-shirt, suggest a relaxed and focused task at hand. Each motion is deliberate, with fingers deftly guiding the paper, ensuring it lays flat without creases. In the background, a hint of a red curtain can be seen, adding a touch of color and depth to the setting. The composition of the scene emphasizes the contrast between the bright paper and the rich tones of the surroundings. This serene and methodical action evokes a sense of exploration and preparation.",
|
||||
"video_path": "Man/mixkit-unrolling-a-world-map-on-a-table-21626_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A young woman sits on a vibrant green seat inside a bus, illuminated by the soft glow of pink and blue lights. Her outfit is a striking mix of colors: a neon pink top paired with a jacket featuring dark sleeves, and jeans that provide a neutral contrast. She wears large, hoop earrings that catch the light as she moves slightly, exuding an air of cool confidence. Her gaze is directed thoughtfully to the side, suggesting contemplation or daydreaming during her commute. The metallic pole beside her adds a geometric element to the composition, reflecting the kaleidoscope of neon hues. The background is a clean, futuristic white, serving as a blank canvas that amplifies the neon atmosphere. Her relaxed posture and the modern bus setting create a scene that captures a blend of urban life and personal introspection.",
|
||||
"video_path": "City/mixkit-fashion-model-posing-on-a-bus-42578_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A silver SUV drives along a winding, snow-covered mountain road, with dense pine trees blanketed in snow lining both sides. The scene is serene, with the vehicle moving smoothly, possibly on a winter journey or vacation. As the SUV disappears around the bend, another, darker SUV follows, creating a sense of motion and perspective on the snow-dusted asphalt. The towering, snow-laden rock formation to the right contrasts with the dark green of the pines, highlighting the peacefulness of the wintry landscape.",
|
||||
"video_path": "Car/mixkit-curve-on-a-snowy-forest-road-3317_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "The video showcases a vibrant urban skyline during twilight, with towering buildings reflecting the warm hues of the setting sun. A series of tall, cylindrical structures dominate the foreground, adjacent to a complex of industrial equipment and grids. The scene includes modern high-rise buildings with glass exteriors, capturing the evolving architecture of a bustling cityscape. A prominent structure labeled \"CITY OF AUSTIN POWER PLANT\" stands out, highlighting the industrial theme amidst the urban backdrop. The soft glow of city lights begins to pierce the approaching dusk, creating an inviting yet dynamic atmosphere. Shadows cast by the buildings add depth and contrast, emphasizing their massive scale and intricate designs. The overall composition is balanced between the natural light of the sunset and the artificial illumination of the city, offering a compelling visual narrative.",
|
||||
"video_path": "Car/mixkit-slow-air-travel-in-reverse-over-a-big-city-49841_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In the scene, a striking architectural structure dominates the view, bathed in a soft, ambient light. The enormous yellow arches serve as the centerpiece, drawing the eye upwards with their majestic curves and towering presence. The smooth, clean surfaces of the structure reflect the light, highlighting the texture and depth of the architecture. In the foreground, blurred streaks of headlights and taillights suggest the motion of vehicles passing by, adding dynamic energy to the otherwise still scene. The contrast between the fast-moving lights and the static arches creates a balanced composition. To the left, a lone streetlamp and a small tree provide a touch of nature and urban elements against the monumental backdrop. The night sky subtly peeks through the gaps in the structure, hinting at a clear, calm evening. Shadows from the arches create patterns on the ground, adding an intricate detail to the scene. Overall, the combination of light, shadow, and movement makes for a dramatic and visually captivating moment.",
|
||||
"video_path": "Car/mixkit-a-fast-timelapse-of-the-street-with-a-monumental-yellow-50993_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A tranquil marina comes into full view under the golden hues of a setting sun. A collection of gleaming yachts and boats are neatly moored, their reflections shimmering softly on the gentle water. The sun's low position casts elongated shadows over the bustling harbor scene, while rolling hillsides surround the distant cityscape. The skyline is interspersed with modern buildings and clusters of residences, adding layers to the vibrant community. At the center, a broad wooden pier juts confidently into the harbor, extending an invitation for leisurely strolls. To the left, various shops and colorful structures line the waterfront, indicating a vibrant coastal economy. The entire atmosphere exudes a serene yet lively charm, balancing the hustle of maritime activity with the peacefulness of the encroaching dusk. It's a scene of calm anticipation, as if the whole place holds its breath before the night's events unfold.",
|
||||
"video_path": "beach/mixkit-harbor-on-a-tourist-coast-with-many-boats-and-yachts-40077_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "The video features a confident individual standing atop a structure against a clear blue sky, exuding a sense of freedom and style. The person is clad in a striking yellow button-up shirt tied at the waist, and beneath it, they wear a simple white top that adds to their relaxed yet stylish appearance. Completing the ensemble are high-waisted white jeans paired with a black belt, adding a touch of contrast. Around their neck is a bold red scarf, providing a splash of color and an air of vintage flair. The person's sunglasses, tinted in yellow, reflect the sunlight and contribute to the overall cool and composed demeanor. Their hair is styled elegantly, pulled back with headphones resting over the ears, suggesting they are immersed in music. One hand casually grazes the headphones, while the other rests gently on the railing, grounding the individual in the moment. The scene is an effortless blend of fashion and tranquility, capturing the spirit of sunny, carefree days.",
|
||||
"video_path": "Music/mixkit-standing-woman-listening-to-music-460_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A ballerina gracefully spins and moves across a pink-hued studio, her poised figure accentuated by a shimmering white tutu and bodice. The background, a continuous wash of soft pink, provides a serene and ethereal atmosphere, emphasizing her fluid movements. Her arms extend with elegance, highlighting the delicacy and precision of her ballet pose, while her focused expression adds intensity to the scene. The subtle details of her costume, combined with the pink monochromatic ambiance, create a dreamlike spectacle, ideal for an AI to envision a oneiric dance setting.",
|
||||
"video_path": "Dance/mixkit-portrait-of-a-ballerina-spinning-with-pink-background-40163_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "The scene unfolds with two human figures in the distance, making their way through a serene meadow, thick with tall golden grass swaying gently in the breeze. The sun hangs low in the sky, casting a soft, diffused glow that illuminates the landscape with a warm, ethereal light. These figures, clad in hiking gear, move deliberately, suggesting they're either embarking on or concluding a journey. Their silhouettes contrast against the lush greenery of the surrounding trees, whose branches reach out, framing the horizon. The play of light and shadow among the trees creates a quilt of textures, with each leaf catching a hint of the sun's dying rays. This tranquil setting evokes a sense of calm and adventure, capturing the quintessential beauty of nature\u2019s landscape.",
|
||||
"video_path": "People/mixkit-landscape-in-nature-while-two-people-are-jogging-44348_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A large cargo ship is docked at an industrial port, its white superstructure contrasting with the deep green and yellow of its deck. The foreground is dominated by the calm, deep blue waters of the harbor, which reflect the vessel\u2019s imposing presence. Surrounding the ship, a series of industrial buildings and storage facilities are visible, hinting at the bustling activity of the port. The deck is intricately detailed, featuring an array of pipes, equipment, and railings, showcasing the ship's functionality and purpose. In the background, a paved area with green patches and a few parked vehicles adds to the busy, industrious atmosphere of the scene.",
|
||||
"video_path": "sea/mixkit-empty-cargo-ship-waiting-at-the-port-4209_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A lone climber ascends a towering rock face, clad in a pink shirt and gray pants, displaying a determined and focused expression. The climber navigates the rugged surface, where the texture of the rock is peppered with natural pockets and crevices that offer handholds and footholds. Sunlight casts soft shadows across the cliff, highlighting the intricate patterns and the climber\u2019s strategic movements. The cliff looms high, with sparse vegetation breaking the monotony of the stone, while distant rocky formations form a dramatic backdrop against the clear blue sky. The climber\u2019s gear, including a harness and chalk bag, underscores the adventure and challenge woven into this majestic, vertical journey.",
|
||||
"video_path": "Sport/mixkit-mountaineer-girl-climbing-a-steep-rocky-mountain-41089_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A person is seen in a close-up shot, skillfully adjusting the tuning pegs of a guitar, showcasing a focused and practiced hand. The image is in black and white, highlighting the contrast between the textures of the instrument and the clothing. The individual's shirt, visible in the background, adds a soft, subtle texture, while the dark tones of the guitar neck create depth in the scene. This composition captures a moment of concentration and finesse, perfect for recreating an intimate musical setting.",
|
||||
"video_path": "Music/mixkit-guitarist-playing-so-inspired-black-and-white-shot-44178_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A musician is playing a large brass instrument with the words \"Brass Band\" clearly visible on its bell. The scene is set against a vibrant yellow backdrop, casting a warm glow on the subject. The musician wears a dark cap and a matching suit, adding a formal touch to his attire. He is deeply focused on his performance, with the instrument's intricate tubing adding complexity to the visual composition. The lighting creates dramatic shadows and highlights, emphasizing the musician's expression and the instrument's metallic sheen. This harmonious blend of color and form captures the essence of a live brass band performance.",
|
||||
"video_path": "Music/mixkit-musician-playing-the-trombone-while-dancing-43752_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In the video, a lone musician stands gracefully in front of a grand cathedral, playing an accordion while surrounded by the lively water display of a central fountain. Dressed in a casual ensemble, he wears a light-colored shirt, dark pants, and a flat cap that gives him a vintage charm. His posture is relaxed, yet engaged, as he sways gently in rhythm with the music, casting soft shadows on the cobblestone steps beneath him. The backdrop features the cathedral's towering twin spires, with intricate stonework that casts a rich, historical aura around the scene. Sunlight bathes the entire setting, enhancing the golden hues of the cathedral facade and creating a halo-like effect around the musician. The fountain's water jets splash playfully, catching glimmers of light and adding a dynamic element to the tranquil atmosphere. The scene captures a harmonious blend of architectural majesty and human creativity, framed by the clear, azure sky that extends infinitely above. It's a vivid depiction of solitude and artistry, set against a timeless urban landscape.",
|
||||
"video_path": "Music/mixkit-man-plays-an-accordion-in-front-of-a-fountain-630_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In the tranquil video, a person sits in a meditative pose on a gentle hillside, silhouetted against the dawning sky. The person is facing the breathtaking sunrise, with their back slightly turned to the viewer, wearing a simple, light-colored shirt. Their right hand rests on their knee, fingers relaxed in a common meditation mudra, symbolizing calmness and peace. The sky, a stunning blend of soft oranges and deep purples, gradually brightens, casting a warm glow over the lush, green landscape. To the left, the outlines of distant urban buildings can be seen against the horizon, adding a contrast between nature and city life. A river reflecting the sky's colors meanders through the scene, lending a serene, flowing dynamic to the landscape. Trees rise and fall gently across the terrain, their leaves rustling only faintly in the morning breeze. The person remains still and focused, embodying a moment of mindfulness and connection with nature. This visual captures a harmonious balance, evoking a sense of tranquility and introspection.",
|
||||
"video_path": "City/mixkit-girl-meditating-in-yoga-pose-at-sunset-4803_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A serene landscape video captures a breathtaking panoramic view of a vast valley covered in a gentle mist. The undulating hills are lush with dense greenery, their rich foliage creating a vibrant border on the left side of the frame. The mist weaves through the landscape like a soft, ethereal blanket, lending a dream-like quality to the scene. In the distance, several mountain peaks emerge, their dark outlines contrasting against the pale blue sky. A few faint, wispy clouds drift lazily across the horizon, complementing the tranquil atmosphere. The sunlight filters through the haze, casting a warm glow and highlighting different textures of the flora. The overall mood is calm and contemplative, inviting the viewer to pause and appreciate nature's untouched beauty. The composition emphasizes depth and expansiveness, drawing attention to the harmony between earth and sky. This captivating scene embodies tranquility, offering a perfect backdrop for meditation or relaxation.",
|
||||
"video_path": "forest/mixkit-flying-over-a-hill-with-a-view-of-the-surrounding-49743_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In this scene, a bearded individual is intently focused on their smartphone, with the sun setting in the background, casting a warm glow across the cityscape. The person, partially visible, is wearing a dark, buttoned shirt that contrasts with the golden hue of the sunset. Their hands are holding the smartphone delicately but purposefully, reflecting a sense of engagement and focus on the screen. The sunlight creates a striking lens flare effect, enhancing the dramatic atmosphere of the moment as it glimmers off the phone\u2019s surface. The surrounding environment hints at an elevated vantage point, providing a panoramic view of the urban landscape below.",
|
||||
"video_path": "City/mixkit-guy-texting-at-sunset-265_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In an expansive, industrial space defined by towering columns and high ceilings, a solitary figure takes center stage. The person, dressed in dark, fitted clothing, assumes a powerful, dynamic stance with one leg bent forward and both arms outstretched in a horizontal arc. Framing this pose are intense flames that engulf their arms, creating a striking visual contrast against the muted tones of the room. The fire forms a brilliant halo of orange and yellow, casting flickering shadows on the weathered walls and worn, tiled floor. This interplay between light and dark showcases the dancer's poise and agility, as they maintain balance amidst the intense heat. Windows line the background, their panes dimly illuminated by the daylight filtering in, adding depth and perspective to the scene. The entire performance evokes a sense of raw energy and elemental mastery, as the figure continues to manipulate the fire in a seamless, mesmerizing display.",
|
||||
"video_path": "fire/mixkit-expert-juggler-doing-tricks-with-a-stick-with-fire-43663_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A man is playing the violin, focused intently on his music. His fingers gracefully dance along the strings, flawlessly executing each note. He holds the violin close to his chin with a sense of familiarity and expertise. The rich, warm tones of the violin reflect in the soft lighting of the room. He wears a dark shirt, and a subtle necklace rests against his chest, adding a personal touch to his attire. The bow moves smoothly across the strings, producing a melody that seems to fill the space with emotion. His expression is one of concentration and passion, immersing himself fully in the performance. The background is softly blurred, bringing the violin's intricate craftsmanship and his precise movements into sharp focus. This serene and intimate moment captures the essence of his musical artistry.",
|
||||
"video_path": "Music/mixkit-fiddler-playing-a-song-639_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In the dimly lit parking garage, two figures engage in an impromptu game of soccer. The first person, wearing a light grey shirt and black pants with three white stripes, skillfully maneuvers the ball with precise footwork. The ground is slick with patches of water, reflecting the vibrant neon lights above. A second figure, clad in dark clothing, stands poised in the background, ready to intercept. The space is defined by stark yellow lines and orange safety bollards, adding structure to the chaotic energy of the scene. The soccer ball glides smoothly across the wet floor, kicking up droplets as it passes. Despite the muted colors of the environment, the players' movements are dynamic and full of life. Their shadowy silhouettes dance with the reflecting light, creating a mesmerizing visual interplay. The atmosphere is charged with focus and camaraderie, encapsulating the essence of a late-night urban soccer experience.",
|
||||
"video_path": "Sport/mixkit-player-making-skillful-play-in-a-street-soccer-game-43504_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A lone climber is seen scaling a towering vertical rock face, demonstrating remarkable strength and focus. Dressed in a light-colored shirt and jeans, the climber grips the stone tightly, navigating the rough textures and crevices with precision. The sheer cliff is massive, exhibiting a range of natural hues from light tan to deep gray, accentuating the climber's figure against the vast rocky backdrop. Surrounding the cliff, scattered greenery and rugged terrain provide a sense of wilderness and isolation. The scene portrays a daring ascension requiring concentration and skill, capturing the essence of human endeavor against nature's formidable beauty.",
|
||||
"video_path": "Sport/mixkit-skilled-mountaineer-climbing-a-gigantic-mountain-41083_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In this serene landscape, a lush meadow stretches across the foreground, dotted with vibrant yellow wildflowers swaying gently in the breeze. A towering tree stands majestically on the right side, its branches reaching wide under the bright blue sky filled with fluffy white clouds. On the left, dense trees form a natural corridor leading to the horizon, suggesting a sense of journey and possibility. The richness of the green grass contrasts beautifully with the golden hue of the distant fields, creating a harmonious palette of nature\u2019s colors. The play of light and shadow adds depth and dimension, evoking a tranquil, inviting atmosphere. It's a scene where nature\u2019s beauty simply commands attention, offering a perfect escape into tranquility.",
|
||||
"video_path": "sky/mixkit-countryside-meadow-4075_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A solitary boat glides across the expansive, tranquil expanse of a serene lake. The vessel leaves a gentle wake behind, creating delicate ripples across the mirror-like surface. The water appears a rich shade of teal, seamlessly blending with the sky at the horizon. Silhouettes of distant trees are faintly visible, creating a picturesque backdrop that enhances the solitary journey of the boat. The sky is a calm gradient, shifting from soft oranges near the shore to the pale blues above. In the distance, a few slender poles emerge from the water, remnants of an old structure or natural formation. The mood of the scene is one of peace and solitude, with the boat journeying steadily through the quiet landscape. There is a sense of endless possibilities as the boat moves toward the unseen beyond the frame. The simplicity and stillness of the scene invite contemplation and reflection, encapsulating a perfect moment of quietude on the water.",
|
||||
"video_path": "mountain/mixkit-motorboat-on-a-large-lake-with-turquoise-blue-waters-4996_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In a cozy, dimly lit caf\u00e9, a woman sits alone at a rustic wooden table, fully engrossed in her reading. Her dark, wavy hair frames her face as she leans forward over an open book, suggesting deep focus and contemplation. The caf\u00e9\u2019s ambiance is warm, with hanging pendant lights casting a soft glow over the wooden shelves lined with jars and coffee paraphernalia in the background. A small cup of coffee rests just within her reach, alongside a glass dome encasing a solitary pastry, adding a touch of tranquility to the scene. Her casual attire, a denim jacket over a simple shirt, complements the laid-back, comfortable setting of the caf\u00e9. The contrast between her concentrated expression and the bustling, yet subdued caf\u00e9 atmosphere creates a harmonious, serene visual. The overall composition captures a quiet moment of introspection amidst the gentle hum of caf\u00e9 life.",
|
||||
"video_path": "Woman/mixkit-woman-drinking-coffee-in-a-cafe-223_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In a vast, deserted landscape under the night sky, a solitary figure stands at a small music setup, illuminated by strategically placed lights. The person is engrossed in playing a keyboard, with various electronic equipment surrounding them, casting soft glows of orange and blue hues across the scene. To the left, a large circular light adds a dramatic focal point, highlighting the intense contrast between the darkness and the lit performance area. This setup, with its minimalistic design and strategic lighting, creates a captivating and easily recognizable scene that merges the serene, expansive backdrop with an intimate, focused music performance.",
|
||||
"video_path": "Music/mixkit-talented-dj-playing-in-a-lonely-desert-42414_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In a bustling urban scene, cars zoom past a weathered building, their blurred motion a testament to the city\u2019s lively pace. The building, with its faded yellow and brown facade, boasts graffiti that speaks of both art and decay, framing the scene with an air of urban grit. A solitary figure stands slightly to the side, clad casually in a gray top and mustard trousers, gazing into the street, seemingly detached from the surrounding flurry. The motion of the traffic creates a dynamic contrast against the static backdrop, emphasizing the relentless movement of the city. As the video progresses, a bright yellow taxi appears, slowing down as it approaches the figure, adding a pop of color to the desaturated hues of the environment. The interaction suggests a routine, a possibly daily exchange between the driver and the pedestrian, hinting at the rhythms of city life. Overhead, a soft, overcast sky casts a diffused light, lending the scene a subdued, timeless quality. Small elements, like the vertical pole cutting through the frame and the distant chatter of urban sounds, complete this vivid tableau of urban existence.",
|
||||
"video_path": "Car/mixkit-morning-in-the-street-time-lapse-1648_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A young woman sits on a curb in a tranquil park, basking in the golden hue of the setting sun. Beside her, a collie dog rests calmly, its fur illuminated by the warm sunlight, creating a serene glow. The woman's hand gently strokes the dog's back, highlighting the bond and affection between them. Tall trees surround the pair, casting elongated shadows on the leaf-laden ground, adding to the peaceful and intimate ambiance of the scene.",
|
||||
"video_path": "Pets/mixkit-a-woman-pets-a-dog-in-a-park-1562_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In the video, a grand, majestic elephant stands in an open, sunlit field, its massive form dominating the scene. The elephant's skin is a tapestry of earthy tones, with rough, textured wrinkles that add character to its already imposing presence. Its trunk, a powerful and flexible appendage, moves gently, swaying as the elephant possibly enjoys the warmth of the day. The background is a blur of greenery, suggesting a lively environment filled with trees and shrubs that provide a natural habitat. Light plays on the elephant's skin, highlighting patches of dust and dirt that give it an authentic wilderness look. The scene captures the tranquility and majesty of this gentle giant in its natural surroundings.",
|
||||
"video_path": "Zoo/mixkit-wet-elephant-in-the-savanna-3663_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In the video, a fluffy dog with brown patches is intently engaged with a bright red toy shaped like a fire hydrant, which has a yellow and orange rope attached. The dog's body is relaxed as it lies on a plain white background, concentrating on nudging and playfully biting the toy. Its ears perk up slightly with curiosity, and its eyes are fixated on the toy, suggesting a scene of focused playfulness. The neutral tones of the dog's fur contrast starkly against the vivid red of the toy, creating a visually striking moment.",
|
||||
"video_path": "Pets/mixkit-a-cute-border-collie-dog-play-with-a-fire-street-50662_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
@@ -9,9 +9,9 @@ def main():
|
||||
# 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-T2V-1.3B-Diffusers",
|
||||
"Wan-AI/Wan2.1-T2V-14B-Diffusers",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
num_gpus=4,
|
||||
use_fsdp_inference=True,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
@@ -25,9 +25,7 @@ def main():
|
||||
# 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 = (
|
||||
"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."
|
||||
"A watermelon wearing a helmet is crushed by a hydraulic press, causing it to flatten and burst open."
|
||||
)
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True)
|
||||
# video = generator.generate_video(prompt, sampling_param=sampling_param, output_path="wan_t2v_videos/")
|
||||
@@ -35,13 +33,9 @@ def main():
|
||||
# Generate another video with a different prompt, without reloading the
|
||||
# model!
|
||||
prompt2 = (
|
||||
"A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
"The video shows a green and orange object being flattened as if it were under a hydraulic press, with the press moving down and compressing the object")
|
||||
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
main()
|
||||
|
||||
@@ -16,6 +16,7 @@ def main():
|
||||
use_fsdp_inference=True,
|
||||
text_encoder_cpu_offload=False,
|
||||
dit_cpu_offload=False,
|
||||
use_sf_wan=True,
|
||||
)
|
||||
|
||||
sampling_param = SamplingParam.from_pretrained(model_name)
|
||||
|
||||
@@ -165,6 +165,9 @@ class FastVideoArgs:
|
||||
# MoE parameters used by Wan2.2
|
||||
boundary_ratio: float | None = None
|
||||
|
||||
# XXX
|
||||
use_sf_wan: bool = False # force self-forcing Wan model for both distillation and validation
|
||||
|
||||
@property
|
||||
def training_mode(self) -> bool:
|
||||
return not self.inference_mode
|
||||
@@ -191,6 +194,12 @@ class FastVideoArgs:
|
||||
help=
|
||||
"The path of the model weights. This can be a local folder or a Hugging Face repo ID.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--use-sf-wan",
|
||||
action=StoreBoolean,
|
||||
default=FastVideoArgs.use_sf_wan,
|
||||
help="Use self-forcing Wan model for both distillation and validation",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--model-dir",
|
||||
type=str,
|
||||
@@ -708,6 +717,7 @@ class TrainingArgs(FastVideoArgs):
|
||||
same_step_across_blocks: bool = False # Use same exit timestep for all blocks
|
||||
last_step_only: bool = False # Only use the last timestep for training
|
||||
context_noise: int = 0 # Context noise level for cache updates
|
||||
sf_ode_init_path: str = "" # Path to ODE init weights for self-forcing model
|
||||
|
||||
@classmethod
|
||||
def from_cli_args(cls, args: argparse.Namespace) -> "TrainingArgs":
|
||||
@@ -1138,6 +1148,11 @@ class TrainingArgs(FastVideoArgs):
|
||||
type=int,
|
||||
default=TrainingArgs.context_noise,
|
||||
help="Context noise level for cache updates")
|
||||
parser.add_argument(
|
||||
"--sf-ode-init-path",
|
||||
type=str,
|
||||
default=TrainingArgs.sf_ode_init_path,
|
||||
help="Path to ODE init weights for self-forcing model")
|
||||
|
||||
return parser
|
||||
|
||||
|
||||
@@ -212,9 +212,9 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
|
||||
frame_seqlen = normalized.shape[1] // num_frames
|
||||
modulated = (
|
||||
normalized.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
|
||||
(1 + scale) + shift).flatten(1, 2)
|
||||
(1.0 + scale) + shift).flatten(1, 2)
|
||||
else:
|
||||
modulated = normalized * (1 + scale) + shift
|
||||
modulated = normalized * (1.0 + scale) + shift
|
||||
return modulated, residual_output
|
||||
|
||||
|
||||
@@ -267,11 +267,11 @@ class LayerNormScaleShift(nn.Module):
|
||||
frame_seqlen = normalized.shape[1] // num_frames
|
||||
output = (
|
||||
normalized.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
|
||||
(1 + scale) + shift).flatten(1, 2)
|
||||
(1.0 + scale) + shift).flatten(1, 2)
|
||||
else:
|
||||
# scale.shape: [batch_size, 1, inner_dim]
|
||||
# shift.shape: [batch_size, 1, inner_dim]
|
||||
output = normalized * (1 + scale) + shift
|
||||
output = normalized * (1.0 + scale) + shift
|
||||
|
||||
if self.compute_dtype == torch.float32:
|
||||
output = output.to(x.dtype)
|
||||
|
||||
@@ -149,9 +149,15 @@ class CausalWanSelfAttention(nn.Module):
|
||||
local_start_index = local_end_index - num_new_tokens
|
||||
# kv_cache["k"] = kv_cache["k"].detach()
|
||||
# kv_cache["v"] = kv_cache["v"].detach()
|
||||
# logger.info("kv_cache['k'] is in comp graph: %s", kv_cache["k"].requires_grad or kv_cache["k"].grad_fn is not None)
|
||||
kv_cache["k"][:, local_start_index:local_end_index] = roped_key
|
||||
kv_cache["v"][:, local_start_index:local_end_index] = v
|
||||
logger.info("kv_cache['k'] is in comp graph: %s", kv_cache["k"].requires_grad or kv_cache["k"].grad_fn is not None)
|
||||
s, e = local_start_index, local_end_index
|
||||
k = kv_cache["k"]
|
||||
kv_cache["k"] = torch.cat([k[:, :s], roped_key, k[:, e:]], dim=1)
|
||||
|
||||
v0 = kv_cache["v"]
|
||||
kv_cache["v"] = torch.cat([v0[:, :s], v, v0[:, e:]], dim=1)
|
||||
# kv_cache["k"][:, local_start_index:local_end_index] = roped_key
|
||||
# kv_cache["v"][:, local_start_index:local_end_index] = v
|
||||
x = self.attn(
|
||||
roped_query,
|
||||
kv_cache["k"][:, max(0, local_end_index - self.max_attention_size):local_end_index],
|
||||
@@ -179,7 +185,7 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
super().__init__()
|
||||
|
||||
# 1. Self-attention
|
||||
self.norm1 = nn.LayerNorm(dim, eps, elementwise_affine=False)
|
||||
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
|
||||
self.to_q = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_k = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_v = ReplicatedLinear(dim, dim, bias=True)
|
||||
@@ -212,7 +218,8 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=True,
|
||||
dtype=torch.float32)
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
|
||||
# 2. Cross-attention
|
||||
# Only T2V for now
|
||||
@@ -225,7 +232,8 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=False,
|
||||
dtype=torch.float32)
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
|
||||
# 3. Feed-forward
|
||||
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
|
||||
@@ -250,34 +258,29 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
if hidden_states.dim() == 4:
|
||||
hidden_states = hidden_states.squeeze(1)
|
||||
num_frames = temb.shape[1]
|
||||
frame_seqlen = hidden_states.shape[1] // num_frames
|
||||
frame_seqlen = hidden_states.shape[1] // num_frames
|
||||
bs, seq_length, _ = hidden_states.shape
|
||||
orig_dtype = hidden_states.dtype
|
||||
# assert orig_dtype != torch.float32
|
||||
e = self.scale_shift_table + temb
|
||||
e = self.scale_shift_table + temb.float()
|
||||
# e.shape: [batch_size, num_frames, 6, inner_dim]
|
||||
assert e.shape == (bs, num_frames, 6, self.hidden_dim)
|
||||
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
|
||||
6, dim=2)
|
||||
# *_msa.shape: [batch_size, num_frames, 1, inner_dim]
|
||||
# assert shift_msa.dtype == torch.float32
|
||||
|
||||
# logger.info("temb sum: %s, dtype: %s", temb.float().sum().item(), temb.dtype)
|
||||
# logger.info("scale_msa sum: %s, dtype: %s", scale_msa.float().sum().item(), scale_msa.dtype)
|
||||
# logger.info("shift_msa sum: %s, dtype: %s", shift_msa.float().sum().item(), shift_msa.dtype)
|
||||
assert shift_msa.dtype == torch.float32
|
||||
|
||||
# 1. Self-attention
|
||||
norm_hidden_states = (self.norm1(hidden_states).unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
|
||||
(1 + scale_msa) + shift_msa).flatten(1, 2)
|
||||
# logger.info("norm_hidden_states sum: %s, shape: %s", norm_hidden_states.float().sum().item(), norm_hidden_states.shape)
|
||||
norm_hidden_states = (self.norm1(hidden_states.float()).unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
|
||||
(1 + scale_msa) + shift_msa).flatten(1, 2).to(orig_dtype)
|
||||
query, _ = self.to_q(norm_hidden_states)
|
||||
key, _ = self.to_k(norm_hidden_states)
|
||||
value, _ = self.to_v(norm_hidden_states)
|
||||
|
||||
if self.norm_q is not None:
|
||||
query = self.norm_q.forward_native(query)
|
||||
query = self.norm_q(query)
|
||||
if self.norm_k is not None:
|
||||
key = self.norm_k.forward_native(key)
|
||||
key = self.norm_k(key)
|
||||
|
||||
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
|
||||
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
|
||||
@@ -291,6 +294,8 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
|
||||
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
|
||||
hidden_states, attn_output, gate_msa, null_shift, null_scale)
|
||||
norm_hidden_states, hidden_states = norm_hidden_states.to(
|
||||
orig_dtype), hidden_states.to(orig_dtype)
|
||||
|
||||
# 2. Cross-attention
|
||||
attn_output = self.attn2(norm_hidden_states,
|
||||
@@ -299,10 +304,13 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
crossattn_cache=crossattn_cache)
|
||||
norm_hidden_states, hidden_states = self.cross_attn_residual_norm(
|
||||
hidden_states, attn_output, 1, c_shift_msa, c_scale_msa)
|
||||
norm_hidden_states, hidden_states = norm_hidden_states.to(
|
||||
orig_dtype), hidden_states.to(orig_dtype)
|
||||
|
||||
# 3. Feed-forward
|
||||
ff_output = self.ffn(norm_hidden_states)
|
||||
hidden_states = self.mlp_residual(hidden_states, ff_output, c_gate_msa)
|
||||
hidden_states = hidden_states.to(orig_dtype)
|
||||
|
||||
return hidden_states
|
||||
|
||||
@@ -365,7 +373,8 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
norm_type="layer",
|
||||
eps=config.eps,
|
||||
elementwise_affine=False,
|
||||
dtype=torch.float32)
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
self.proj_out = nn.Linear(
|
||||
inner_dim, config.out_channels * math.prod(config.patch_size))
|
||||
self.scale_shift_table = nn.Parameter(
|
||||
@@ -487,16 +496,12 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
)
|
||||
freqs_cos = freqs_cos.to(hidden_states.device)
|
||||
freqs_sin = freqs_sin.to(hidden_states.device)
|
||||
freqs_cis = (freqs_cos,
|
||||
freqs_sin) if freqs_cos is not None else None
|
||||
freqs_cis = (freqs_cos.float(),
|
||||
freqs_sin.float()) if freqs_cos is not None else None
|
||||
|
||||
hidden_states = self.patch_embedding(hidden_states)
|
||||
grid_sizes = torch.stack(
|
||||
[torch.tensor(hidden_states[0].shape[1:], dtype=torch.long)])
|
||||
hidden_states = hidden_states.flatten(2).transpose(1, 2)
|
||||
|
||||
encoder_hidden_states = torch.cat([encoder_hidden_states, encoder_hidden_states.new_zeros(1, self.text_len - encoder_hidden_states.size(1), encoder_hidden_states.size(2))], dim=1)
|
||||
|
||||
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
|
||||
timestep.flatten(), encoder_hidden_states, encoder_hidden_states_image)
|
||||
timestep_proj = timestep_proj.unflatten(1, (6, self.hidden_size)).unflatten(dim=0, sizes=timestep.shape)
|
||||
@@ -543,9 +548,14 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
hidden_states = self.norm_out(hidden_states, shift, scale)
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
|
||||
output = self.unpatchify(hidden_states, grid_sizes)
|
||||
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
|
||||
post_patch_height,
|
||||
post_patch_width, p_t, p_h, p_w,
|
||||
-1)
|
||||
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
|
||||
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
|
||||
|
||||
return torch.stack(output)
|
||||
return output
|
||||
|
||||
def _forward_train(self,
|
||||
hidden_states: torch.Tensor,
|
||||
@@ -586,8 +596,8 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
)
|
||||
freqs_cos = freqs_cos.to(hidden_states.device)
|
||||
freqs_sin = freqs_sin.to(hidden_states.device)
|
||||
freqs_cis = (freqs_cos,
|
||||
freqs_sin) if freqs_cos is not None else None
|
||||
freqs_cis = (freqs_cos.float(),
|
||||
freqs_sin.float()) if freqs_cos is not None else None
|
||||
|
||||
# Construct blockwise causal attn mask
|
||||
if self.block_mask is None:
|
||||
@@ -600,12 +610,8 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
)
|
||||
|
||||
hidden_states = self.patch_embedding(hidden_states)
|
||||
grid_sizes = torch.stack(
|
||||
[torch.tensor(hidden_states[0].shape[1:], dtype=torch.long)])
|
||||
hidden_states = hidden_states.flatten(2).transpose(1, 2)
|
||||
|
||||
encoder_hidden_states = torch.cat([encoder_hidden_states, encoder_hidden_states.new_zeros(1, self.text_len - encoder_hidden_states.size(1), encoder_hidden_states.size(2))], dim=1)
|
||||
|
||||
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
|
||||
timestep.flatten(), encoder_hidden_states, encoder_hidden_states_image)
|
||||
timestep_proj = timestep_proj.unflatten(1, (6, self.hidden_size)).unflatten(dim=0, sizes=timestep.shape)
|
||||
@@ -640,9 +646,14 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
hidden_states = self.norm_out(hidden_states, shift, scale)
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
|
||||
output = self.unpatchify(hidden_states, grid_sizes)
|
||||
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
|
||||
post_patch_height,
|
||||
post_patch_width, p_t, p_h, p_w,
|
||||
-1)
|
||||
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
|
||||
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
|
||||
|
||||
return torch.stack(output)
|
||||
return output
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -653,30 +664,3 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
return self._forward_inference(*args, **kwargs)
|
||||
else:
|
||||
return self._forward_train(*args, **kwargs)
|
||||
|
||||
|
||||
def unpatchify(self, x, grid_sizes):
|
||||
r"""
|
||||
|
||||
|
||||
Args:
|
||||
x (List[Tensor]):
|
||||
List of patchified features, each with shape [L, C_out * prod(patch_size)]
|
||||
grid_sizes (Tensor):
|
||||
Original spatial-temporal grid dimensions before patching,
|
||||
|
||||
|
||||
Returns:
|
||||
Tensor:
|
||||
Reconstructed video tensors with shape [B, C_out, F, H / 8, W / 8]
|
||||
"""
|
||||
|
||||
c = self.out_channels
|
||||
out = []
|
||||
for u, v in zip(x, grid_sizes.tolist()):
|
||||
u = u[:math.prod(v)].view(*v, *self.patch_size, c)
|
||||
u = u.permute(6, 0, 3, 1, 4, 2, 5)
|
||||
# u = torch.einsum('fhwpqrc->cfphqwr', u.contiguous())
|
||||
u = u.reshape(c, *[i * j for i, j in zip(v, self.patch_size)])
|
||||
out.append(u)
|
||||
return out
|
||||
|
||||
@@ -1,5 +1,3 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import math
|
||||
@@ -39,14 +37,16 @@ class WanImageEmbedding(torch.nn.Module):
|
||||
def __init__(self, in_features: int, out_features: int):
|
||||
super().__init__()
|
||||
|
||||
self.norm1 = nn.LayerNorm(in_features)
|
||||
self.norm1 = FP32LayerNorm(in_features)
|
||||
self.ff = MLP(in_features, in_features, out_features, act_type="gelu")
|
||||
self.norm2 = nn.LayerNorm(out_features)
|
||||
self.norm2 = FP32LayerNorm(out_features)
|
||||
|
||||
def forward(self, encoder_hidden_states_image: torch.Tensor) -> torch.Tensor:
|
||||
def forward(self,
|
||||
encoder_hidden_states_image: torch.Tensor) -> torch.Tensor:
|
||||
dtype = encoder_hidden_states_image.dtype
|
||||
hidden_states = self.norm1(encoder_hidden_states_image)
|
||||
hidden_states = self.ff(hidden_states)
|
||||
hidden_states = self.norm2(hidden_states)
|
||||
hidden_states = self.norm2(hidden_states).to(dtype)
|
||||
return hidden_states
|
||||
|
||||
|
||||
@@ -62,7 +62,7 @@ class WanTimeTextImageEmbedding(nn.Module):
|
||||
super().__init__()
|
||||
|
||||
self.time_embedder = TimestepEmbedder(
|
||||
dim, frequency_embedding_size=time_freq_dim, act_layer="silu", freq_dtype=torch.float64)
|
||||
dim, frequency_embedding_size=time_freq_dim, act_layer="silu")
|
||||
self.time_modulation = ModulateProjection(dim,
|
||||
factor=6,
|
||||
act_layer="silu")
|
||||
@@ -156,12 +156,12 @@ class WanT2VCrossAttention(WanSelfAttention):
|
||||
b, n, d = x.size(0), self.num_heads, self.head_dim
|
||||
|
||||
# compute query, key, value
|
||||
q = self.norm_q.forward_native(self.to_q(x)[0]).view(b, -1, n, d)
|
||||
q = self.norm_q(self.to_q(x)[0]).view(b, -1, n, d)
|
||||
|
||||
if crossattn_cache is not None:
|
||||
if not crossattn_cache["is_init"]:
|
||||
crossattn_cache["is_init"] = True
|
||||
k = self.norm_k.forward_native(self.to_k(context)[0]).view(b, -1, n, d)
|
||||
k = self.norm_k(self.to_k(context)[0]).view(b, -1, n, d)
|
||||
v = self.to_v(context)[0].view(b, -1, n, d)
|
||||
crossattn_cache["k"] = k
|
||||
crossattn_cache["v"] = v
|
||||
@@ -169,7 +169,7 @@ class WanT2VCrossAttention(WanSelfAttention):
|
||||
k = crossattn_cache["k"]
|
||||
v = crossattn_cache["v"]
|
||||
else:
|
||||
k = self.norm_k.forward_native(self.to_k(context)[0]).view(b, -1, n, d)
|
||||
k = self.norm_k(self.to_k(context)[0]).view(b, -1, n, d)
|
||||
v = self.to_v(context)[0].view(b, -1, n, d)
|
||||
|
||||
# compute attention
|
||||
@@ -213,10 +213,10 @@ class WanI2VCrossAttention(WanSelfAttention):
|
||||
b, n, d = x.size(0), self.num_heads, self.head_dim
|
||||
|
||||
# compute query, key, value
|
||||
q = self.norm_q.forward_native(self.to_q(x)[0]).view(b, -1, n, d)
|
||||
k = self.norm_k.forward_native(self.to_k(context)[0]).view(b, -1, n, d)
|
||||
q = self.norm_q(self.to_q(x)[0]).view(b, -1, n, d)
|
||||
k = self.norm_k(self.to_k(context)[0]).view(b, -1, n, d)
|
||||
v = self.to_v(context)[0].view(b, -1, n, d)
|
||||
k_img = self.norm_added_k.forward_native(self.add_k_proj(context_img)[0]).view(
|
||||
k_img = self.norm_added_k(self.add_k_proj(context_img)[0]).view(
|
||||
b, -1, n, d)
|
||||
v_img = self.add_v_proj(context_img)[0].view(b, -1, n, d)
|
||||
img_x = self.attn(q, k_img, v_img)
|
||||
@@ -247,7 +247,7 @@ class WanTransformerBlock(nn.Module):
|
||||
super().__init__()
|
||||
|
||||
# 1. Self-attention
|
||||
self.norm1 = nn.LayerNorm(dim, eps, elementwise_affine=False)
|
||||
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
|
||||
self.to_q = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_k = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_v = ReplicatedLinear(dim, dim, bias=True)
|
||||
@@ -278,29 +278,29 @@ class WanTransformerBlock(nn.Module):
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=True,
|
||||
dtype=torch.float32)
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
|
||||
# 2. Cross-attention
|
||||
if added_kv_proj_dim is not None:
|
||||
# I2V
|
||||
self.attn2 = WanI2VCrossAttention(dim,
|
||||
num_heads,
|
||||
qk_norm=qk_norm,
|
||||
self.attn2 = WanI2VCrossAttention(dim,
|
||||
num_heads,
|
||||
qk_norm=qk_norm,
|
||||
eps=eps)
|
||||
|
||||
else:
|
||||
# T2V
|
||||
self.attn2 = WanT2VCrossAttention(dim,
|
||||
num_heads,
|
||||
qk_norm=qk_norm,
|
||||
self.attn2 = WanT2VCrossAttention(dim,
|
||||
num_heads,
|
||||
qk_norm=qk_norm,
|
||||
eps=eps)
|
||||
|
||||
self.cross_attn_residual_norm = ScaleResidualLayerNormScaleShift(
|
||||
dim,
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=False,
|
||||
dtype=torch.float32)
|
||||
dim,
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=False,
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
|
||||
# 3. Feed-forward
|
||||
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
|
||||
@@ -319,11 +319,12 @@ class WanTransformerBlock(nn.Module):
|
||||
hidden_states = hidden_states.squeeze(1)
|
||||
bs, seq_length, _ = hidden_states.shape
|
||||
orig_dtype = hidden_states.dtype
|
||||
# assert orig_dtype != torch.float32
|
||||
|
||||
if temb.dim() == 4:
|
||||
# temb: batch_size, seq_len, 6, inner_dim (wan2.2 ti2v)
|
||||
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = (
|
||||
self.scale_shift_table.unsqueeze(0) + temb
|
||||
self.scale_shift_table.unsqueeze(0) + temb.float()
|
||||
).chunk(6, dim=2)
|
||||
# batch_size, seq_len, 1, inner_dim
|
||||
shift_msa = shift_msa.squeeze(2)
|
||||
@@ -334,20 +335,22 @@ class WanTransformerBlock(nn.Module):
|
||||
c_gate_msa = c_gate_msa.squeeze(2)
|
||||
else:
|
||||
# temb: batch_size, 6, inner_dim (wan2.1/wan2.2 14B)
|
||||
e = self.scale_shift_table + temb
|
||||
e = self.scale_shift_table + temb.float()
|
||||
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
|
||||
6, dim=1)
|
||||
assert shift_msa.dtype == torch.float32
|
||||
|
||||
# 1. Self-attention
|
||||
norm_hidden_states = self.norm1(hidden_states) * (1 + scale_msa) + shift_msa
|
||||
norm_hidden_states = (self.norm1(hidden_states.float()) *
|
||||
(1 + scale_msa) + shift_msa).to(orig_dtype)
|
||||
query, _ = self.to_q(norm_hidden_states)
|
||||
key, _ = self.to_k(norm_hidden_states)
|
||||
value, _ = self.to_v(norm_hidden_states)
|
||||
|
||||
if self.norm_q is not None:
|
||||
query = self.norm_q.forward_native(query)
|
||||
query = self.norm_q(query)
|
||||
if self.norm_k is not None:
|
||||
key = self.norm_k.forward_native(key)
|
||||
key = self.norm_k(key)
|
||||
|
||||
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
|
||||
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
|
||||
@@ -367,20 +370,26 @@ class WanTransformerBlock(nn.Module):
|
||||
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
|
||||
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
|
||||
hidden_states, attn_output, gate_msa, null_shift, null_scale)
|
||||
norm_hidden_states, hidden_states = norm_hidden_states.to(
|
||||
orig_dtype), hidden_states.to(orig_dtype)
|
||||
|
||||
# 2. Cross-attention
|
||||
attn_output = self.attn2(norm_hidden_states,
|
||||
context=encoder_hidden_states,
|
||||
attn_output = self.attn2(norm_hidden_states,
|
||||
context=encoder_hidden_states,
|
||||
context_lens=None)
|
||||
norm_hidden_states, hidden_states = self.cross_attn_residual_norm(
|
||||
hidden_states, attn_output, 1, c_shift_msa, c_scale_msa)
|
||||
norm_hidden_states, hidden_states = norm_hidden_states.to(
|
||||
orig_dtype), hidden_states.to(orig_dtype)
|
||||
|
||||
# 3. Feed-forward
|
||||
ff_output = self.ffn(norm_hidden_states)
|
||||
hidden_states = self.mlp_residual(hidden_states, ff_output, c_gate_msa)
|
||||
hidden_states = hidden_states.to(orig_dtype)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class WanTransformerBlock_VSA(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
@@ -397,7 +406,7 @@ class WanTransformerBlock_VSA(nn.Module):
|
||||
super().__init__()
|
||||
|
||||
# 1. Self-attention
|
||||
self.norm1 = nn.LayerNorm(dim, eps, elementwise_affine=False)
|
||||
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
|
||||
self.to_q = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_k = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_v = ReplicatedLinear(dim, dim, bias=True)
|
||||
@@ -429,7 +438,8 @@ class WanTransformerBlock_VSA(nn.Module):
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=True,
|
||||
dtype=torch.float32)
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
|
||||
# 2. Cross-attention
|
||||
if added_kv_proj_dim is not None:
|
||||
@@ -449,7 +459,8 @@ class WanTransformerBlock_VSA(nn.Module):
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=False,
|
||||
dtype=torch.float32)
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
|
||||
# 3. Feed-forward
|
||||
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
|
||||
@@ -469,22 +480,23 @@ class WanTransformerBlock_VSA(nn.Module):
|
||||
bs, seq_length, _ = hidden_states.shape
|
||||
orig_dtype = hidden_states.dtype
|
||||
# assert orig_dtype != torch.float32
|
||||
e = self.scale_shift_table + temb
|
||||
e = self.scale_shift_table + temb.float()
|
||||
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
|
||||
6, dim=1)
|
||||
assert shift_msa.dtype == torch.float32
|
||||
|
||||
# 1. Self-attention
|
||||
norm_hidden_states = (self.norm1(hidden_states) *
|
||||
(1 + scale_msa) + shift_msa)
|
||||
norm_hidden_states = (self.norm1(hidden_states.float()) *
|
||||
(1 + scale_msa) + shift_msa).to(orig_dtype)
|
||||
query, _ = self.to_q(norm_hidden_states)
|
||||
key, _ = self.to_k(norm_hidden_states)
|
||||
value, _ = self.to_v(norm_hidden_states)
|
||||
gate_compress, _ = self.to_gate_compress(norm_hidden_states)
|
||||
|
||||
if self.norm_q is not None:
|
||||
query = self.norm_q.forward_native(query)
|
||||
query = self.norm_q(query)
|
||||
if self.norm_k is not None:
|
||||
key = self.norm_k.forward_native(key)
|
||||
key = self.norm_k(key)
|
||||
|
||||
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
|
||||
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
|
||||
@@ -509,6 +521,8 @@ class WanTransformerBlock_VSA(nn.Module):
|
||||
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
|
||||
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
|
||||
hidden_states, attn_output, gate_msa, null_shift, null_scale)
|
||||
norm_hidden_states, hidden_states = norm_hidden_states.to(
|
||||
orig_dtype), hidden_states.to(orig_dtype)
|
||||
|
||||
# 2. Cross-attention
|
||||
attn_output = self.attn2(norm_hidden_states,
|
||||
@@ -516,15 +530,17 @@ class WanTransformerBlock_VSA(nn.Module):
|
||||
context_lens=None)
|
||||
norm_hidden_states, hidden_states = self.cross_attn_residual_norm(
|
||||
hidden_states, attn_output, 1, c_shift_msa, c_scale_msa)
|
||||
norm_hidden_states, hidden_states = norm_hidden_states.to(
|
||||
orig_dtype), hidden_states.to(orig_dtype)
|
||||
|
||||
# 3. Feed-forward
|
||||
ff_output = self.ffn(norm_hidden_states)
|
||||
hidden_states = self.mlp_residual(hidden_states, ff_output, c_gate_msa)
|
||||
hidden_states = hidden_states.to(orig_dtype)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
|
||||
class WanTransformer3DModel(CachableDiT):
|
||||
_fsdp_shard_conditions = WanVideoConfig()._fsdp_shard_conditions
|
||||
_compile_conditions = WanVideoConfig()._compile_conditions
|
||||
@@ -582,7 +598,8 @@ class WanTransformer3DModel(CachableDiT):
|
||||
norm_type="layer",
|
||||
eps=config.eps,
|
||||
elementwise_affine=False,
|
||||
dtype=torch.float32)
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
self.proj_out = nn.Linear(
|
||||
inner_dim, config.out_channels * math.prod(config.patch_size))
|
||||
self.scale_shift_table = nn.Parameter(
|
||||
@@ -642,12 +659,10 @@ class WanTransformer3DModel(CachableDiT):
|
||||
rope_theta=10000)
|
||||
freqs_cos = freqs_cos.to(hidden_states.device)
|
||||
freqs_sin = freqs_sin.to(hidden_states.device)
|
||||
freqs_cis = (freqs_cos,
|
||||
freqs_sin) if freqs_cos is not None else None
|
||||
freqs_cis = (freqs_cos.float(),
|
||||
freqs_sin.float()) if freqs_cos is not None else None
|
||||
|
||||
hidden_states = self.patch_embedding(hidden_states)
|
||||
grid_sizes = torch.stack(
|
||||
[torch.tensor(hidden_states[0].shape[1:], dtype=torch.long)])
|
||||
hidden_states = hidden_states.flatten(2).transpose(1, 2)
|
||||
|
||||
# timestep shape: batch_size, or batch_size, seq_len (wan 2.2 ti2v)
|
||||
@@ -657,8 +672,6 @@ class WanTransformer3DModel(CachableDiT):
|
||||
else:
|
||||
ts_seq_len = None
|
||||
|
||||
encoder_hidden_states = torch.cat([encoder_hidden_states, encoder_hidden_states.new_zeros(1, self.text_len - encoder_hidden_states.size(1), encoder_hidden_states.size(2))], dim=1)
|
||||
|
||||
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
|
||||
timestep, encoder_hidden_states, encoder_hidden_states_image, timestep_seq_len=ts_seq_len)
|
||||
if ts_seq_len is not None:
|
||||
@@ -715,35 +728,14 @@ class WanTransformer3DModel(CachableDiT):
|
||||
hidden_states = self.norm_out(hidden_states, shift, scale)
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
|
||||
output = self.unpatchify(hidden_states, grid_sizes)
|
||||
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
|
||||
post_patch_height,
|
||||
post_patch_width, p_t, p_h, p_w,
|
||||
-1)
|
||||
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
|
||||
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
|
||||
|
||||
return torch.stack(output)
|
||||
|
||||
def unpatchify(self, x, grid_sizes):
|
||||
r"""
|
||||
|
||||
|
||||
Args:
|
||||
x (List[Tensor]):
|
||||
List of patchified features, each with shape [L, C_out * prod(patch_size)]
|
||||
grid_sizes (Tensor):
|
||||
Original spatial-temporal grid dimensions before patching,
|
||||
|
||||
|
||||
Returns:
|
||||
Tensor:
|
||||
Reconstructed video tensors with shape [B, C_out, F, H / 8, W / 8]
|
||||
"""
|
||||
|
||||
c = self.out_channels
|
||||
out = []
|
||||
for u, v in zip(x, grid_sizes.tolist()):
|
||||
u = u[:math.prod(v)].view(*v, *self.patch_size, c)
|
||||
u = u.permute(6, 0, 3, 1, 4, 2, 5)
|
||||
# u = torch.einsum('fhwpqrc->cfphqwr', u.contiguous())
|
||||
u = u.reshape(c, *[i * j for i, j in zip(v, self.patch_size)])
|
||||
out.append(u)
|
||||
return out
|
||||
return output
|
||||
|
||||
def maybe_cache_states(self, hidden_states: torch.Tensor,
|
||||
original_hidden_states: torch.Tensor) -> None:
|
||||
@@ -835,4 +827,5 @@ class WanTransformer3DModel(CachableDiT):
|
||||
if self.is_even:
|
||||
return hidden_states + self.previous_residual_even
|
||||
else:
|
||||
return hidden_states + self.previous_residual_odd
|
||||
return hidden_states + self.previous_residual_odd
|
||||
|
||||
@@ -19,8 +19,14 @@ from fastvideo.pipelines.stages import (ConditioningStage, DecodingStage,
|
||||
TextEncodingStage)
|
||||
# isort: on
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.sf_utils.wan_wrapper import WanDiffusionWrapper
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
|
||||
|
||||
class WanCausalDMDPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
|
||||
@@ -30,6 +36,21 @@ class WanCausalDMDPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
"""Set up pipeline stages with proper dependency injection."""
|
||||
if fastvideo_args.use_sf_wan:
|
||||
# timestep shift is 5.0 for self-forcing Wan model
|
||||
# see https://github.com/guandeh17/Self-Forcing/blob/33593df3e81fa3ec10239271dd2c100facac6de1/configs/self_forcing_dmd.yaml#L50
|
||||
config = self.get_module("transformer").config
|
||||
if not isinstance(self.modules["transformer"], WanDiffusionWrapper):
|
||||
sf_transformer = WanDiffusionWrapper(
|
||||
model_name="Wan2.1-T2V-1.3B", timestep_shift=5.0, is_causal=True, config=config)
|
||||
del self.modules["transformer"]
|
||||
state_dict = torch.load('checkpoints/self_forcing_dmd.pt')
|
||||
sf_transformer.load_state_dict(state_dict['generator_ema'])
|
||||
sf_transformer.to(get_local_torch_device())
|
||||
self.modules["transformer"] = sf_transformer
|
||||
logger.info("Using self-forcing Wan model for DMD inference")
|
||||
else:
|
||||
logger.info("transformer is already a WanDiffusionWrapper")
|
||||
|
||||
self.add_stage(stage_name="input_validation_stage",
|
||||
stage=InputValidationStage())
|
||||
|
||||
@@ -205,7 +205,7 @@ def import_pipeline_classes(
|
||||
except ImportError as e:
|
||||
raise ImportError(
|
||||
f"Could not import {pipeline_type_package_name} when importing pipeline classes: {e}"
|
||||
) from None
|
||||
) from e
|
||||
|
||||
type_to_arch_to_pipeline_dict[pipeline_type_str] = arch_to_pipeline_dict
|
||||
|
||||
|
||||
@@ -150,6 +150,7 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=target_dtype,
|
||||
enabled=autocast_enabled):
|
||||
assert False, "image_first_btchw is not supported"
|
||||
_ = self.transformer(
|
||||
image_first_btchw,
|
||||
prompt_embeds,
|
||||
@@ -176,6 +177,7 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=target_dtype,
|
||||
enabled=autocast_enabled):
|
||||
assert False, "ref_btchw is not supported"
|
||||
_ = self.transformer(
|
||||
ref_btchw,
|
||||
prompt_embeds,
|
||||
@@ -273,6 +275,9 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
(latent_model_input.shape[0], 1),
|
||||
device=latent_model_input.device,
|
||||
dtype=torch.long)
|
||||
# if fastvideo_args.use_sf_wan:
|
||||
# # SF wan wrapper requires BTCHW input
|
||||
# latent_model_input = latent_model_input.permute(0, 2, 1, 3, 4)
|
||||
pred_noise_btchw = self.transformer(
|
||||
latent_model_input,
|
||||
prompt_embeds,
|
||||
@@ -285,6 +290,14 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
**image_kwargs,
|
||||
**pos_cond_kwargs,
|
||||
).permute(0, 2, 1, 3, 4)
|
||||
# if fastvideo_args.use_sf_wan:
|
||||
# flow_pred, x0_pred = pred_noise_btchw
|
||||
# logger.info(f"flow_pred.shape: {flow_pred.shape}")
|
||||
# logger.info(f"x0_pred.shape: {x0_pred.shape}")
|
||||
# # SF wan wrapper requires BTCHW output
|
||||
# pred_noise_btchw = flow_pred
|
||||
# else:
|
||||
# pred_noise_btchw = pred_noise_btchw.permute(0, 2, 1, 3, 4)
|
||||
|
||||
# Convert pred noise to pred video with FM Euler scheduler utilities
|
||||
pred_video_btchw = pred_noise_to_pred_video(
|
||||
@@ -338,6 +351,9 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
attn_metadata=attn_metadata,
|
||||
forward_batch=batch):
|
||||
t_expanded_context = t_context.unsqueeze(1)
|
||||
# if fastvideo_args.use_sf_wan:
|
||||
# SF wan wrapper requires BTCHW input
|
||||
# context_bcthw = context_bcthw.permute(0, 2, 1, 3, 4)
|
||||
_ = self.transformer(
|
||||
context_bcthw,
|
||||
prompt_embeds,
|
||||
@@ -349,7 +365,10 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
start_frame=start_index,
|
||||
**image_kwargs,
|
||||
**pos_cond_kwargs,
|
||||
)
|
||||
).permute(0, 2, 1, 3, 4)
|
||||
# if fastvideo_args.use_sf_wan:
|
||||
# SF wan wrapper requires BTCHW output
|
||||
# context_bcthw = context_bcthw.permute(0, 2, 1, 3, 4)
|
||||
start_index += current_num_frames
|
||||
|
||||
batch.latents = latents
|
||||
|
||||
@@ -0,0 +1,220 @@
|
||||
from utils.lmdb import get_array_shape_from_lmdb, retrieve_row_from_lmdb
|
||||
from torch.utils.data import Dataset
|
||||
import numpy as np
|
||||
import torch
|
||||
import lmdb
|
||||
import json
|
||||
from pathlib import Path
|
||||
from PIL import Image
|
||||
import os
|
||||
|
||||
|
||||
class TextDataset(Dataset):
|
||||
def __init__(self, prompt_path, extended_prompt_path=None):
|
||||
with open(prompt_path, encoding="utf-8") as f:
|
||||
self.prompt_list = [line.rstrip() for line in f]
|
||||
|
||||
if extended_prompt_path is not None:
|
||||
with open(extended_prompt_path, encoding="utf-8") as f:
|
||||
self.extended_prompt_list = [line.rstrip() for line in f]
|
||||
assert len(self.extended_prompt_list) == len(self.prompt_list)
|
||||
else:
|
||||
self.extended_prompt_list = None
|
||||
|
||||
def __len__(self):
|
||||
return len(self.prompt_list)
|
||||
|
||||
def __getitem__(self, idx):
|
||||
batch = {
|
||||
"prompts": self.prompt_list[idx],
|
||||
"idx": idx,
|
||||
}
|
||||
if self.extended_prompt_list is not None:
|
||||
batch["extended_prompts"] = self.extended_prompt_list[idx]
|
||||
return batch
|
||||
|
||||
|
||||
class ODERegressionLMDBDataset(Dataset):
|
||||
def __init__(self, data_path: str, max_pair: int = int(1e8)):
|
||||
self.env = lmdb.open(data_path, readonly=True,
|
||||
lock=False, readahead=False, meminit=False)
|
||||
|
||||
self.latents_shape = get_array_shape_from_lmdb(self.env, 'latents')
|
||||
self.max_pair = max_pair
|
||||
|
||||
def __len__(self):
|
||||
return min(self.latents_shape[0], self.max_pair)
|
||||
|
||||
def __getitem__(self, idx):
|
||||
"""
|
||||
Outputs:
|
||||
- prompts: List of Strings
|
||||
- latents: Tensor of shape (num_denoising_steps, num_frames, num_channels, height, width). It is ordered from pure noise to clean image.
|
||||
"""
|
||||
latents = retrieve_row_from_lmdb(
|
||||
self.env,
|
||||
"latents", np.float16, idx, shape=self.latents_shape[1:]
|
||||
)
|
||||
|
||||
if len(latents.shape) == 4:
|
||||
latents = latents[None, ...]
|
||||
|
||||
prompts = retrieve_row_from_lmdb(
|
||||
self.env,
|
||||
"prompts", str, idx
|
||||
)
|
||||
return {
|
||||
"prompts": prompts,
|
||||
"ode_latent": torch.tensor(latents, dtype=torch.float32)
|
||||
}
|
||||
|
||||
|
||||
class ShardingLMDBDataset(Dataset):
|
||||
def __init__(self, data_path: str, max_pair: int = int(1e8)):
|
||||
self.envs = []
|
||||
self.index = []
|
||||
|
||||
for fname in sorted(os.listdir(data_path)):
|
||||
path = os.path.join(data_path, fname)
|
||||
env = lmdb.open(path,
|
||||
readonly=True,
|
||||
lock=False,
|
||||
readahead=False,
|
||||
meminit=False)
|
||||
self.envs.append(env)
|
||||
|
||||
self.latents_shape = [None] * len(self.envs)
|
||||
for shard_id, env in enumerate(self.envs):
|
||||
self.latents_shape[shard_id] = get_array_shape_from_lmdb(env, 'latents')
|
||||
for local_i in range(self.latents_shape[shard_id][0]):
|
||||
self.index.append((shard_id, local_i))
|
||||
|
||||
# print("shard_id ", shard_id, " local_i ", local_i)
|
||||
|
||||
self.max_pair = max_pair
|
||||
|
||||
def __len__(self):
|
||||
return len(self.index)
|
||||
|
||||
def __getitem__(self, idx):
|
||||
"""
|
||||
Outputs:
|
||||
- prompts: List of Strings
|
||||
- latents: Tensor of shape (num_denoising_steps, num_frames, num_channels, height, width). It is ordered from pure noise to clean image.
|
||||
"""
|
||||
shard_id, local_idx = self.index[idx]
|
||||
|
||||
latents = retrieve_row_from_lmdb(
|
||||
self.envs[shard_id],
|
||||
"latents", np.float16, local_idx,
|
||||
shape=self.latents_shape[shard_id][1:]
|
||||
)
|
||||
|
||||
if len(latents.shape) == 4:
|
||||
latents = latents[None, ...]
|
||||
|
||||
prompts = retrieve_row_from_lmdb(
|
||||
self.envs[shard_id],
|
||||
"prompts", str, local_idx
|
||||
)
|
||||
|
||||
return {
|
||||
"prompts": prompts,
|
||||
"ode_latent": torch.tensor(latents, dtype=torch.float32)
|
||||
}
|
||||
|
||||
|
||||
class TextImagePairDataset(Dataset):
|
||||
def __init__(
|
||||
self,
|
||||
data_dir,
|
||||
transform=None,
|
||||
eval_first_n=-1,
|
||||
pad_to_multiple_of=None
|
||||
):
|
||||
"""
|
||||
Args:
|
||||
data_dir (str): Path to the directory containing:
|
||||
- target_crop_info_*.json (metadata file)
|
||||
- */ (subdirectory containing images with matching aspect ratio)
|
||||
transform (callable, optional): Optional transform to be applied on the image
|
||||
"""
|
||||
self.transform = transform
|
||||
data_dir = Path(data_dir)
|
||||
|
||||
# Find the metadata JSON file
|
||||
metadata_files = list(data_dir.glob('target_crop_info_*.json'))
|
||||
if not metadata_files:
|
||||
raise FileNotFoundError(f"No metadata file found in {data_dir}")
|
||||
if len(metadata_files) > 1:
|
||||
raise ValueError(f"Multiple metadata files found in {data_dir}")
|
||||
|
||||
metadata_path = metadata_files[0]
|
||||
# Extract aspect ratio from metadata filename (e.g. target_crop_info_26-15.json -> 26-15)
|
||||
aspect_ratio = metadata_path.stem.split('_')[-1]
|
||||
|
||||
# Use aspect ratio subfolder for images
|
||||
self.image_dir = data_dir / aspect_ratio
|
||||
if not self.image_dir.exists():
|
||||
raise FileNotFoundError(f"Image directory not found: {self.image_dir}")
|
||||
|
||||
# Load metadata
|
||||
with open(metadata_path, 'r') as f:
|
||||
self.metadata = json.load(f)
|
||||
|
||||
eval_first_n = eval_first_n if eval_first_n != -1 else len(self.metadata)
|
||||
self.metadata = self.metadata[:eval_first_n]
|
||||
|
||||
# Verify all images exist
|
||||
for item in self.metadata:
|
||||
image_path = self.image_dir / item['file_name']
|
||||
if not image_path.exists():
|
||||
raise FileNotFoundError(f"Image not found: {image_path}")
|
||||
|
||||
self.dummy_prompt = "DUMMY PROMPT"
|
||||
self.pre_pad_len = len(self.metadata)
|
||||
if pad_to_multiple_of is not None and len(self.metadata) % pad_to_multiple_of != 0:
|
||||
# Duplicate the last entry
|
||||
self.metadata += [self.metadata[-1]] * (
|
||||
pad_to_multiple_of - len(self.metadata) % pad_to_multiple_of
|
||||
)
|
||||
|
||||
def __len__(self):
|
||||
return len(self.metadata)
|
||||
|
||||
def __getitem__(self, idx):
|
||||
"""
|
||||
Returns:
|
||||
dict: A dictionary containing:
|
||||
- image: PIL Image
|
||||
- caption: str
|
||||
- target_bbox: list of int [x1, y1, x2, y2]
|
||||
- target_ratio: str
|
||||
- type: str
|
||||
- origin_size: tuple of int (width, height)
|
||||
"""
|
||||
item = self.metadata[idx]
|
||||
|
||||
# Load image
|
||||
image_path = self.image_dir / item['file_name']
|
||||
image = Image.open(image_path).convert('RGB')
|
||||
|
||||
# Apply transform if specified
|
||||
if self.transform:
|
||||
image = self.transform(image)
|
||||
|
||||
return {
|
||||
'image': image,
|
||||
'prompts': item['caption'],
|
||||
'target_bbox': item['target_crop']['target_bbox'],
|
||||
'target_ratio': item['target_crop']['target_ratio'],
|
||||
'type': item['type'],
|
||||
'origin_size': (item['origin_width'], item['origin_height']),
|
||||
'idx': idx
|
||||
}
|
||||
|
||||
|
||||
def cycle(dl):
|
||||
while True:
|
||||
for data in dl:
|
||||
yield data
|
||||
@@ -0,0 +1,125 @@
|
||||
from datetime import timedelta
|
||||
from functools import partial
|
||||
import os
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from torch.distributed.fsdp import FullStateDictConfig, FullyShardedDataParallel as FSDP, MixedPrecision, ShardingStrategy, StateDictType
|
||||
from torch.distributed.fsdp.api import CPUOffload
|
||||
from torch.distributed.fsdp.wrap import size_based_auto_wrap_policy, transformer_auto_wrap_policy
|
||||
|
||||
|
||||
def fsdp_state_dict(model):
|
||||
fsdp_fullstate_save_policy = FullStateDictConfig(
|
||||
offload_to_cpu=True, rank0_only=True
|
||||
)
|
||||
with FSDP.state_dict_type(
|
||||
model, StateDictType.FULL_STATE_DICT, fsdp_fullstate_save_policy
|
||||
):
|
||||
checkpoint = model.state_dict()
|
||||
|
||||
return checkpoint
|
||||
|
||||
|
||||
def fsdp_wrap(module, sharding_strategy="full", mixed_precision=False, wrap_strategy="size", min_num_params=int(5e7), transformer_module=None, cpu_offload=False):
|
||||
if mixed_precision:
|
||||
mixed_precision_policy = MixedPrecision(
|
||||
param_dtype=torch.bfloat16,
|
||||
reduce_dtype=torch.float32,
|
||||
buffer_dtype=torch.float32,
|
||||
cast_forward_inputs=False
|
||||
)
|
||||
else:
|
||||
mixed_precision_policy = None
|
||||
|
||||
if wrap_strategy == "transformer":
|
||||
auto_wrap_policy = partial(
|
||||
transformer_auto_wrap_policy,
|
||||
transformer_layer_cls=transformer_module
|
||||
)
|
||||
elif wrap_strategy == "size":
|
||||
auto_wrap_policy = partial(
|
||||
size_based_auto_wrap_policy,
|
||||
min_num_params=min_num_params
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"Invalid wrap strategy: {wrap_strategy}")
|
||||
|
||||
os.environ["NCCL_CROSS_NIC"] = "1"
|
||||
|
||||
sharding_strategy = {
|
||||
"full": ShardingStrategy.FULL_SHARD,
|
||||
"hybrid_full": ShardingStrategy.HYBRID_SHARD,
|
||||
"hybrid_zero2": ShardingStrategy._HYBRID_SHARD_ZERO2,
|
||||
"no_shard": ShardingStrategy.NO_SHARD,
|
||||
}[sharding_strategy]
|
||||
|
||||
module = FSDP(
|
||||
module,
|
||||
auto_wrap_policy=auto_wrap_policy,
|
||||
sharding_strategy=sharding_strategy,
|
||||
mixed_precision=mixed_precision_policy,
|
||||
device_id=torch.cuda.current_device(),
|
||||
limit_all_gathers=True,
|
||||
use_orig_params=True,
|
||||
cpu_offload=CPUOffload(offload_params=cpu_offload),
|
||||
sync_module_states=False # Load ckpt on rank 0 and sync to other ranks
|
||||
)
|
||||
return module
|
||||
|
||||
|
||||
def barrier():
|
||||
if dist.is_initialized():
|
||||
dist.barrier()
|
||||
|
||||
|
||||
def launch_distributed_job(backend: str = "nccl"):
|
||||
rank = int(os.environ["RANK"])
|
||||
local_rank = int(os.environ["LOCAL_RANK"])
|
||||
world_size = int(os.environ["WORLD_SIZE"])
|
||||
host = os.environ["MASTER_ADDR"]
|
||||
port = int(os.environ["MASTER_PORT"])
|
||||
|
||||
if ":" in host: # IPv6
|
||||
init_method = f"tcp://[{host}]:{port}"
|
||||
else: # IPv4
|
||||
init_method = f"tcp://{host}:{port}"
|
||||
dist.init_process_group(rank=rank, world_size=world_size, backend=backend,
|
||||
init_method=init_method, timeout=timedelta(minutes=30))
|
||||
torch.cuda.set_device(local_rank)
|
||||
|
||||
|
||||
class EMA_FSDP:
|
||||
def __init__(self, fsdp_module: torch.nn.Module, decay: float = 0.999):
|
||||
self.decay = decay
|
||||
self.shadow = {}
|
||||
self._init_shadow(fsdp_module)
|
||||
|
||||
@torch.no_grad()
|
||||
def _init_shadow(self, fsdp_module):
|
||||
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
|
||||
with FSDP.summon_full_params(fsdp_module, writeback=False):
|
||||
for n, p in fsdp_module.module.named_parameters():
|
||||
self.shadow[n] = p.detach().clone().float().cpu()
|
||||
|
||||
@torch.no_grad()
|
||||
def update(self, fsdp_module):
|
||||
d = self.decay
|
||||
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
|
||||
with FSDP.summon_full_params(fsdp_module, writeback=False):
|
||||
for n, p in fsdp_module.module.named_parameters():
|
||||
self.shadow[n].mul_(d).add_(p.detach().float().cpu(), alpha=1. - d)
|
||||
|
||||
# Optional helpers ---------------------------------------------------
|
||||
def state_dict(self):
|
||||
return self.shadow # picklable
|
||||
|
||||
def load_state_dict(self, sd):
|
||||
self.shadow = {k: v.clone() for k, v in sd.items()}
|
||||
|
||||
def copy_to(self, fsdp_module):
|
||||
# load EMA weights into an (unwrapped) copy of the generator
|
||||
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
|
||||
with FSDP.summon_full_params(fsdp_module, writeback=True):
|
||||
for n, p in fsdp_module.module.named_parameters():
|
||||
if n in self.shadow:
|
||||
p.data.copy_(self.shadow[n].to(p.dtype, device=p.device))
|
||||
@@ -0,0 +1,72 @@
|
||||
import numpy as np
|
||||
|
||||
|
||||
def get_array_shape_from_lmdb(env, array_name):
|
||||
with env.begin() as txn:
|
||||
image_shape = txn.get(f"{array_name}_shape".encode()).decode()
|
||||
image_shape = tuple(map(int, image_shape.split()))
|
||||
return image_shape
|
||||
|
||||
|
||||
def store_arrays_to_lmdb(env, arrays_dict, start_index=0):
|
||||
"""
|
||||
Store rows of multiple numpy arrays in a single LMDB.
|
||||
Each row is stored separately with a naming convention.
|
||||
"""
|
||||
with env.begin(write=True) as txn:
|
||||
for array_name, array in arrays_dict.items():
|
||||
for i, row in enumerate(array):
|
||||
# Convert row to bytes
|
||||
if isinstance(row, str):
|
||||
row_bytes = row.encode()
|
||||
else:
|
||||
row_bytes = row.tobytes()
|
||||
|
||||
data_key = f'{array_name}_{start_index + i}_data'.encode()
|
||||
|
||||
txn.put(data_key, row_bytes)
|
||||
|
||||
|
||||
def process_data_dict(data_dict, seen_prompts):
|
||||
output_dict = {}
|
||||
|
||||
all_videos = []
|
||||
all_prompts = []
|
||||
for prompt, video in data_dict.items():
|
||||
if prompt in seen_prompts:
|
||||
continue
|
||||
else:
|
||||
seen_prompts.add(prompt)
|
||||
|
||||
video = video.half().numpy()
|
||||
all_videos.append(video)
|
||||
all_prompts.append(prompt)
|
||||
|
||||
if len(all_videos) == 0:
|
||||
return {"latents": np.array([]), "prompts": np.array([])}
|
||||
|
||||
all_videos = np.concatenate(all_videos, axis=0)
|
||||
|
||||
output_dict['latents'] = all_videos
|
||||
output_dict['prompts'] = np.array(all_prompts)
|
||||
|
||||
return output_dict
|
||||
|
||||
|
||||
def retrieve_row_from_lmdb(lmdb_env, array_name, dtype, row_index, shape=None):
|
||||
"""
|
||||
Retrieve a specific row from a specific array in the LMDB.
|
||||
"""
|
||||
data_key = f'{array_name}_{row_index}_data'.encode()
|
||||
|
||||
with lmdb_env.begin() as txn:
|
||||
row_bytes = txn.get(data_key)
|
||||
|
||||
if dtype == str:
|
||||
array = row_bytes.decode()
|
||||
else:
|
||||
array = np.frombuffer(row_bytes, dtype=dtype)
|
||||
|
||||
if shape is not None and len(shape) > 0:
|
||||
array = array.reshape(shape)
|
||||
return array
|
||||
@@ -0,0 +1,81 @@
|
||||
from abc import ABC, abstractmethod
|
||||
import torch
|
||||
|
||||
|
||||
class DenoisingLoss(ABC):
|
||||
@abstractmethod
|
||||
def __call__(
|
||||
self, x: torch.Tensor, x_pred: torch.Tensor,
|
||||
noise: torch.Tensor, noise_pred: torch.Tensor,
|
||||
alphas_cumprod: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
**kwargs
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Base class for denoising loss.
|
||||
Input:
|
||||
- x: the clean data with shape [B, F, C, H, W]
|
||||
- x_pred: the predicted clean data with shape [B, F, C, H, W]
|
||||
- noise: the noise with shape [B, F, C, H, W]
|
||||
- noise_pred: the predicted noise with shape [B, F, C, H, W]
|
||||
- alphas_cumprod: the cumulative product of alphas (defining the noise schedule) with shape [T]
|
||||
- timestep: the current timestep with shape [B, F]
|
||||
"""
|
||||
pass
|
||||
|
||||
|
||||
class X0PredLoss(DenoisingLoss):
|
||||
def __call__(
|
||||
self, x: torch.Tensor, x_pred: torch.Tensor,
|
||||
noise: torch.Tensor, noise_pred: torch.Tensor,
|
||||
alphas_cumprod: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
**kwargs
|
||||
) -> torch.Tensor:
|
||||
return torch.mean((x - x_pred) ** 2)
|
||||
|
||||
|
||||
class VPredLoss(DenoisingLoss):
|
||||
def __call__(
|
||||
self, x: torch.Tensor, x_pred: torch.Tensor,
|
||||
noise: torch.Tensor, noise_pred: torch.Tensor,
|
||||
alphas_cumprod: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
**kwargs
|
||||
) -> torch.Tensor:
|
||||
weights = 1 / (1 - alphas_cumprod[timestep].reshape(*timestep.shape, 1, 1, 1))
|
||||
return torch.mean(weights * (x - x_pred) ** 2)
|
||||
|
||||
|
||||
class NoisePredLoss(DenoisingLoss):
|
||||
def __call__(
|
||||
self, x: torch.Tensor, x_pred: torch.Tensor,
|
||||
noise: torch.Tensor, noise_pred: torch.Tensor,
|
||||
alphas_cumprod: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
**kwargs
|
||||
) -> torch.Tensor:
|
||||
return torch.mean((noise - noise_pred) ** 2)
|
||||
|
||||
|
||||
class FlowPredLoss(DenoisingLoss):
|
||||
def __call__(
|
||||
self, x: torch.Tensor, x_pred: torch.Tensor,
|
||||
noise: torch.Tensor, noise_pred: torch.Tensor,
|
||||
alphas_cumprod: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
**kwargs
|
||||
) -> torch.Tensor:
|
||||
return torch.mean((kwargs["flow_pred"] - (noise - x)) ** 2)
|
||||
|
||||
|
||||
NAME_TO_CLASS = {
|
||||
"x0": X0PredLoss,
|
||||
"v": VPredLoss,
|
||||
"noise": NoisePredLoss,
|
||||
"flow": FlowPredLoss
|
||||
}
|
||||
|
||||
|
||||
def get_denoising_loss(loss_type: str) -> DenoisingLoss:
|
||||
return NAME_TO_CLASS[loss_type]
|
||||
@@ -0,0 +1,39 @@
|
||||
import numpy as np
|
||||
import random
|
||||
import torch
|
||||
|
||||
|
||||
def set_seed(seed: int, deterministic: bool = False):
|
||||
"""
|
||||
Helper function for reproducible behavior to set the seed in `random`, `numpy`, `torch`.
|
||||
|
||||
Args:
|
||||
seed (`int`):
|
||||
The seed to set.
|
||||
deterministic (`bool`, *optional*, defaults to `False`):
|
||||
Whether to use deterministic algorithms where available. Can slow down training.
|
||||
"""
|
||||
random.seed(seed)
|
||||
np.random.seed(seed)
|
||||
torch.manual_seed(seed)
|
||||
torch.cuda.manual_seed_all(seed)
|
||||
|
||||
if deterministic:
|
||||
torch.use_deterministic_algorithms(True)
|
||||
|
||||
|
||||
def merge_dict_list(dict_list):
|
||||
if len(dict_list) == 1:
|
||||
return dict_list[0]
|
||||
|
||||
merged_dict = {}
|
||||
for k, v in dict_list[0].items():
|
||||
if isinstance(v, torch.Tensor):
|
||||
if v.ndim == 0:
|
||||
merged_dict[k] = torch.stack([d[k] for d in dict_list], dim=0)
|
||||
else:
|
||||
merged_dict[k] = torch.cat([d[k] for d in dict_list], dim=0)
|
||||
else:
|
||||
# for non-tensor values, we just copy the value from the first item
|
||||
merged_dict[k] = v
|
||||
return merged_dict
|
||||
@@ -0,0 +1,194 @@
|
||||
from abc import abstractmethod, ABC
|
||||
import torch
|
||||
|
||||
|
||||
class SchedulerInterface(ABC):
|
||||
"""
|
||||
Base class for diffusion noise schedule.
|
||||
"""
|
||||
alphas_cumprod: torch.Tensor # [T], alphas for defining the noise schedule
|
||||
|
||||
@abstractmethod
|
||||
def add_noise(
|
||||
self, clean_latent: torch.Tensor,
|
||||
noise: torch.Tensor, timestep: torch.Tensor
|
||||
):
|
||||
"""
|
||||
Diffusion forward corruption process.
|
||||
Input:
|
||||
- clean_latent: the clean latent with shape [B, C, H, W]
|
||||
- noise: the noise with shape [B, C, H, W]
|
||||
- timestep: the timestep with shape [B]
|
||||
Output: the corrupted latent with shape [B, C, H, W]
|
||||
"""
|
||||
pass
|
||||
|
||||
def convert_x0_to_noise(
|
||||
self, x0: torch.Tensor, xt: torch.Tensor,
|
||||
timestep: torch.Tensor
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Convert the diffusion network's x0 prediction to noise predidction.
|
||||
x0: the predicted clean data with shape [B, C, H, W]
|
||||
xt: the input noisy data with shape [B, C, H, W]
|
||||
timestep: the timestep with shape [B]
|
||||
|
||||
noise = (xt-sqrt(alpha_t)*x0) / sqrt(beta_t) (eq 11 in https://arxiv.org/abs/2311.18828)
|
||||
"""
|
||||
# use higher precision for calculations
|
||||
original_dtype = x0.dtype
|
||||
x0, xt, alphas_cumprod = map(
|
||||
lambda x: x.double().to(x0.device), [x0, xt,
|
||||
self.alphas_cumprod]
|
||||
)
|
||||
|
||||
alpha_prod_t = alphas_cumprod[timestep].reshape(-1, 1, 1, 1)
|
||||
beta_prod_t = 1 - alpha_prod_t
|
||||
|
||||
noise_pred = (xt - alpha_prod_t **
|
||||
(0.5) * x0) / beta_prod_t ** (0.5)
|
||||
return noise_pred.to(original_dtype)
|
||||
|
||||
def convert_noise_to_x0(
|
||||
self, noise: torch.Tensor, xt: torch.Tensor,
|
||||
timestep: torch.Tensor
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Convert the diffusion network's noise prediction to x0 predidction.
|
||||
noise: the predicted noise with shape [B, C, H, W]
|
||||
xt: the input noisy data with shape [B, C, H, W]
|
||||
timestep: the timestep with shape [B]
|
||||
|
||||
x0 = (x_t - sqrt(beta_t) * noise) / sqrt(alpha_t) (eq 11 in https://arxiv.org/abs/2311.18828)
|
||||
"""
|
||||
# use higher precision for calculations
|
||||
original_dtype = noise.dtype
|
||||
noise, xt, alphas_cumprod = map(
|
||||
lambda x: x.double().to(noise.device), [noise, xt,
|
||||
self.alphas_cumprod]
|
||||
)
|
||||
alpha_prod_t = alphas_cumprod[timestep].reshape(-1, 1, 1, 1)
|
||||
beta_prod_t = 1 - alpha_prod_t
|
||||
|
||||
x0_pred = (xt - beta_prod_t **
|
||||
(0.5) * noise) / alpha_prod_t ** (0.5)
|
||||
return x0_pred.to(original_dtype)
|
||||
|
||||
def convert_velocity_to_x0(
|
||||
self, velocity: torch.Tensor, xt: torch.Tensor,
|
||||
timestep: torch.Tensor
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Convert the diffusion network's velocity prediction to x0 predidction.
|
||||
velocity: the predicted noise with shape [B, C, H, W]
|
||||
xt: the input noisy data with shape [B, C, H, W]
|
||||
timestep: the timestep with shape [B]
|
||||
|
||||
v = sqrt(alpha_t) * noise - sqrt(beta_t) x0
|
||||
noise = (xt-sqrt(alpha_t)*x0) / sqrt(beta_t)
|
||||
given v, x_t, we have
|
||||
x0 = sqrt(alpha_t) * x_t - sqrt(beta_t) * v
|
||||
see derivations https://chatgpt.com/share/679fb6c8-3a30-8008-9b0e-d1ae892dac56
|
||||
"""
|
||||
# use higher precision for calculations
|
||||
original_dtype = velocity.dtype
|
||||
velocity, xt, alphas_cumprod = map(
|
||||
lambda x: x.double().to(velocity.device), [velocity, xt,
|
||||
self.alphas_cumprod]
|
||||
)
|
||||
alpha_prod_t = alphas_cumprod[timestep].reshape(-1, 1, 1, 1)
|
||||
beta_prod_t = 1 - alpha_prod_t
|
||||
|
||||
x0_pred = (alpha_prod_t ** 0.5) * xt - (beta_prod_t ** 0.5) * velocity
|
||||
return x0_pred.to(original_dtype)
|
||||
|
||||
|
||||
class FlowMatchScheduler():
|
||||
|
||||
def __init__(self, num_inference_steps=100, num_train_timesteps=1000, shift=3.0, sigma_max=1.0, sigma_min=0.003 / 1.002, inverse_timesteps=False, extra_one_step=False, reverse_sigmas=False):
|
||||
self.num_train_timesteps = num_train_timesteps
|
||||
self.shift = shift
|
||||
self.sigma_max = sigma_max
|
||||
self.sigma_min = sigma_min
|
||||
self.inverse_timesteps = inverse_timesteps
|
||||
self.extra_one_step = extra_one_step
|
||||
self.reverse_sigmas = reverse_sigmas
|
||||
self.set_timesteps(num_inference_steps)
|
||||
|
||||
def set_timesteps(self, num_inference_steps=100, denoising_strength=1.0, training=False):
|
||||
sigma_start = self.sigma_min + \
|
||||
(self.sigma_max - self.sigma_min) * denoising_strength
|
||||
if self.extra_one_step:
|
||||
self.sigmas = torch.linspace(
|
||||
sigma_start, self.sigma_min, num_inference_steps + 1)[:-1]
|
||||
else:
|
||||
self.sigmas = torch.linspace(
|
||||
sigma_start, self.sigma_min, num_inference_steps)
|
||||
if self.inverse_timesteps:
|
||||
self.sigmas = torch.flip(self.sigmas, dims=[0])
|
||||
self.sigmas = self.shift * self.sigmas / \
|
||||
(1 + (self.shift - 1) * self.sigmas)
|
||||
if self.reverse_sigmas:
|
||||
self.sigmas = 1 - self.sigmas
|
||||
self.timesteps = self.sigmas * self.num_train_timesteps
|
||||
if training:
|
||||
x = self.timesteps
|
||||
y = torch.exp(-2 * ((x - num_inference_steps / 2) /
|
||||
num_inference_steps) ** 2)
|
||||
y_shifted = y - y.min()
|
||||
bsmntw_weighing = y_shifted * \
|
||||
(num_inference_steps / y_shifted.sum())
|
||||
self.linear_timesteps_weights = bsmntw_weighing
|
||||
|
||||
def step(self, model_output, timestep, sample, to_final=False):
|
||||
if timestep.ndim == 2:
|
||||
timestep = timestep.flatten(0, 1)
|
||||
self.sigmas = self.sigmas.to(model_output.device)
|
||||
self.timesteps = self.timesteps.to(model_output.device)
|
||||
timestep_id = torch.argmin(
|
||||
(self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
|
||||
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
|
||||
if to_final or (timestep_id + 1 >= len(self.timesteps)).any():
|
||||
sigma_ = 1 if (
|
||||
self.inverse_timesteps or self.reverse_sigmas) else 0
|
||||
else:
|
||||
sigma_ = self.sigmas[timestep_id + 1].reshape(-1, 1, 1, 1)
|
||||
prev_sample = sample + model_output * (sigma_ - sigma)
|
||||
return prev_sample
|
||||
|
||||
def add_noise(self, original_samples, noise, timestep):
|
||||
"""
|
||||
Diffusion forward corruption process.
|
||||
Input:
|
||||
- clean_latent: the clean latent with shape [B*T, C, H, W]
|
||||
- noise: the noise with shape [B*T, C, H, W]
|
||||
- timestep: the timestep with shape [B*T]
|
||||
Output: the corrupted latent with shape [B*T, C, H, W]
|
||||
"""
|
||||
if timestep.ndim == 2:
|
||||
timestep = timestep.flatten(0, 1)
|
||||
self.sigmas = self.sigmas.to(noise.device)
|
||||
self.timesteps = self.timesteps.to(noise.device)
|
||||
timestep_id = torch.argmin(
|
||||
(self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
|
||||
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
|
||||
sample = (1 - sigma) * original_samples + sigma * noise
|
||||
return sample.type_as(noise)
|
||||
|
||||
def training_target(self, sample, noise, timestep):
|
||||
target = noise - sample
|
||||
return target
|
||||
|
||||
def training_weight(self, timestep):
|
||||
"""
|
||||
Input:
|
||||
- timestep: the timestep with shape [B*T]
|
||||
Output: the corresponding weighting [B*T]
|
||||
"""
|
||||
if timestep.ndim == 2:
|
||||
timestep = timestep.flatten(0, 1)
|
||||
self.linear_timesteps_weights = self.linear_timesteps_weights.to(timestep.device)
|
||||
timestep_id = torch.argmin(
|
||||
(self.timesteps.unsqueeze(1) - timestep.unsqueeze(0)).abs(), dim=0)
|
||||
weights = self.linear_timesteps_weights[timestep_id]
|
||||
return weights
|
||||
@@ -0,0 +1,381 @@
|
||||
import types
|
||||
from typing import List, Optional
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from fastvideo.sf_utils.scheduler import SchedulerInterface, FlowMatchScheduler
|
||||
from fastvideo.wan.modules.tokenizers import HuggingfaceTokenizer
|
||||
from fastvideo.wan.modules.model import WanModel, RegisterTokens, GanAttentionBlock
|
||||
from fastvideo.wan.modules.vae import _video_vae
|
||||
from fastvideo.wan.modules.t5 import umt5_xxl
|
||||
from fastvideo.wan.modules.causal_model import CausalWanModel
|
||||
|
||||
|
||||
class WanTextEncoder(torch.nn.Module):
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
|
||||
self.text_encoder = umt5_xxl(
|
||||
encoder_only=True,
|
||||
return_tokenizer=False,
|
||||
dtype=torch.float32,
|
||||
device=torch.device('cpu')
|
||||
).eval().requires_grad_(False)
|
||||
self.text_encoder.load_state_dict(
|
||||
torch.load("wan_models/Wan2.1-T2V-1.3B/models_t5_umt5-xxl-enc-bf16.pth",
|
||||
map_location='cpu', weights_only=False)
|
||||
)
|
||||
|
||||
self.tokenizer = HuggingfaceTokenizer(
|
||||
name="wan_models/Wan2.1-T2V-1.3B/google/umt5-xxl/", seq_len=512, clean='whitespace')
|
||||
|
||||
@property
|
||||
def device(self):
|
||||
# Assume we are always on GPU
|
||||
return torch.cuda.current_device()
|
||||
|
||||
def forward(self, text_prompts: List[str]) -> dict:
|
||||
ids, mask = self.tokenizer(
|
||||
text_prompts, return_mask=True, add_special_tokens=True)
|
||||
ids = ids.to(self.device)
|
||||
mask = mask.to(self.device)
|
||||
seq_lens = mask.gt(0).sum(dim=1).long()
|
||||
context = self.text_encoder(ids, mask)
|
||||
|
||||
for u, v in zip(context, seq_lens):
|
||||
u[v:] = 0.0 # set padding to 0.0
|
||||
|
||||
return {
|
||||
"prompt_embeds": context
|
||||
}
|
||||
|
||||
|
||||
class WanVAEWrapper(torch.nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
mean = [
|
||||
-0.7571, -0.7089, -0.9113, 0.1075, -0.1745, 0.9653, -0.1517, 1.5508,
|
||||
0.4134, -0.0715, 0.5517, -0.3632, -0.1922, -0.9497, 0.2503, -0.2921
|
||||
]
|
||||
std = [
|
||||
2.8184, 1.4541, 2.3275, 2.6558, 1.2196, 1.7708, 2.6052, 2.0743,
|
||||
3.2687, 2.1526, 2.8652, 1.5579, 1.6382, 1.1253, 2.8251, 1.9160
|
||||
]
|
||||
self.mean = torch.tensor(mean, dtype=torch.float32)
|
||||
self.std = torch.tensor(std, dtype=torch.float32)
|
||||
|
||||
# init model
|
||||
self.model = _video_vae(
|
||||
pretrained_path="wan_models/Wan2.1-T2V-1.3B/Wan2.1_VAE.pth",
|
||||
z_dim=16,
|
||||
).eval().requires_grad_(False)
|
||||
|
||||
def encode_to_latent(self, pixel: torch.Tensor) -> torch.Tensor:
|
||||
# pixel: [batch_size, num_channels, num_frames, height, width]
|
||||
device, dtype = pixel.device, pixel.dtype
|
||||
scale = [self.mean.to(device=device, dtype=dtype),
|
||||
1.0 / self.std.to(device=device, dtype=dtype)]
|
||||
|
||||
output = [
|
||||
self.model.encode(u.unsqueeze(0), scale).float().squeeze(0)
|
||||
for u in pixel
|
||||
]
|
||||
output = torch.stack(output, dim=0)
|
||||
# from [batch_size, num_channels, num_frames, height, width]
|
||||
# to [batch_size, num_frames, num_channels, height, width]
|
||||
output = output.permute(0, 2, 1, 3, 4)
|
||||
return output
|
||||
|
||||
def decode_to_pixel(self, latent: torch.Tensor, use_cache: bool = False) -> torch.Tensor:
|
||||
# from [batch_size, num_frames, num_channels, height, width]
|
||||
# to [batch_size, num_channels, num_frames, height, width]
|
||||
zs = latent.permute(0, 2, 1, 3, 4)
|
||||
if use_cache:
|
||||
assert latent.shape[0] == 1, "Batch size must be 1 when using cache"
|
||||
|
||||
device, dtype = latent.device, latent.dtype
|
||||
scale = [self.mean.to(device=device, dtype=dtype),
|
||||
1.0 / self.std.to(device=device, dtype=dtype)]
|
||||
|
||||
if use_cache:
|
||||
decode_function = self.model.cached_decode
|
||||
else:
|
||||
decode_function = self.model.decode
|
||||
|
||||
output = []
|
||||
for u in zs:
|
||||
output.append(decode_function(u.unsqueeze(0), scale).float().clamp_(-1, 1).squeeze(0))
|
||||
output = torch.stack(output, dim=0)
|
||||
# from [batch_size, num_channels, num_frames, height, width]
|
||||
# to [batch_size, num_frames, num_channels, height, width]
|
||||
output = output.permute(0, 2, 1, 3, 4)
|
||||
return output
|
||||
|
||||
|
||||
class WanDiffusionWrapper(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
model_name="Wan2.1-T2V-1.3B",
|
||||
timestep_shift=8.0,
|
||||
is_causal=False,
|
||||
local_attn_size=-1,
|
||||
sink_size=0,
|
||||
config=None,
|
||||
**kwargs
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
assert config is not None, "config is required"
|
||||
|
||||
self.config = config
|
||||
self.hidden_size = config.hidden_size
|
||||
self.num_attention_heads = config.num_attention_heads
|
||||
# self.num_single_layers = config.num_single_layers
|
||||
self.num_layers = config.num_layers
|
||||
self.hidden_size = config.hidden_size
|
||||
self.num_attention_heads = config.num_attention_heads
|
||||
self.attention_head_dim = config.attention_head_dim
|
||||
self.in_channels = config.in_channels
|
||||
self.out_channels = config.out_channels
|
||||
self.num_channels_latents = config.num_channels_latents
|
||||
self.patch_size = config.patch_size
|
||||
self.text_len = config.text_len
|
||||
self.local_attn_size = config.local_attn_size
|
||||
self.independent_first_frame = False
|
||||
# self.num_single_layers = config.num_single_layers
|
||||
# self.num_registers = config.num_registers
|
||||
# self.num_frame_per_block = config.num_frame_per_block
|
||||
# self.independent_first_frame = config.independent_first_frame
|
||||
# self.local_attn_size = config.local_attn_size
|
||||
|
||||
if is_causal:
|
||||
self.model = CausalWanModel.from_pretrained(
|
||||
f"wan_models/{model_name}/", local_attn_size=local_attn_size, sink_size=sink_size)
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
self.model = self.model.to(torch.bfloat16).to(get_local_torch_device())
|
||||
else:
|
||||
self.model = WanModel.from_pretrained(f"wan_models/{model_name}/")
|
||||
self.model.eval()
|
||||
|
||||
# For non-causal diffusion, all frames share the same timestep
|
||||
self.uniform_timestep = not is_causal
|
||||
|
||||
self.scheduler = FlowMatchScheduler(
|
||||
shift=timestep_shift, sigma_min=0.0, extra_one_step=True
|
||||
)
|
||||
self.scheduler.set_timesteps(1000, training=True)
|
||||
|
||||
self.seq_len = 32760 # [1, 21, 16, 60, 104]
|
||||
self.post_init()
|
||||
|
||||
@property
|
||||
def blocks(self):
|
||||
return self.model.blocks
|
||||
|
||||
def enable_gradient_checkpointing(self) -> None:
|
||||
self.model.enable_gradient_checkpointing()
|
||||
|
||||
def adding_cls_branch(self, atten_dim=1536, num_class=4, time_embed_dim=0) -> None:
|
||||
# NOTE: This is hard coded for WAN2.1-T2V-1.3B for now!!!!!!!!!!!!!!!!!!!!
|
||||
self._cls_pred_branch = nn.Sequential(
|
||||
# Input: [B, 384, 21, 60, 104]
|
||||
nn.LayerNorm(atten_dim * 3 + time_embed_dim),
|
||||
nn.Linear(atten_dim * 3 + time_embed_dim, 1536),
|
||||
nn.SiLU(),
|
||||
nn.Linear(atten_dim, num_class)
|
||||
)
|
||||
self._cls_pred_branch.requires_grad_(True)
|
||||
num_registers = 3
|
||||
self._register_tokens = RegisterTokens(num_registers=num_registers, dim=atten_dim)
|
||||
self._register_tokens.requires_grad_(True)
|
||||
|
||||
gan_ca_blocks = []
|
||||
for _ in range(num_registers):
|
||||
block = GanAttentionBlock()
|
||||
gan_ca_blocks.append(block)
|
||||
self._gan_ca_blocks = nn.ModuleList(gan_ca_blocks)
|
||||
self._gan_ca_blocks.requires_grad_(True)
|
||||
# self.has_cls_branch = True
|
||||
|
||||
def _convert_flow_pred_to_x0(self, flow_pred: torch.Tensor, xt: torch.Tensor, timestep: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Convert flow matching's prediction to x0 prediction.
|
||||
flow_pred: the prediction with shape [B, C, H, W]
|
||||
xt: the input noisy data with shape [B, C, H, W]
|
||||
timestep: the timestep with shape [B]
|
||||
|
||||
pred = noise - x0
|
||||
x_t = (1-sigma_t) * x0 + sigma_t * noise
|
||||
we have x0 = x_t - sigma_t * pred
|
||||
see derivations https://chatgpt.com/share/67bf8589-3d04-8008-bc6e-4cf1a24e2d0e
|
||||
"""
|
||||
# use higher precision for calculations
|
||||
original_dtype = flow_pred.dtype
|
||||
flow_pred, xt, sigmas, timesteps = map(
|
||||
lambda x: x.double().to(flow_pred.device), [flow_pred, xt,
|
||||
self.scheduler.sigmas,
|
||||
self.scheduler.timesteps]
|
||||
)
|
||||
|
||||
timestep_id = torch.argmin(
|
||||
(timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
|
||||
sigma_t = sigmas[timestep_id].reshape(-1, 1, 1, 1)
|
||||
x0_pred = xt - sigma_t * flow_pred
|
||||
return x0_pred.to(original_dtype)
|
||||
|
||||
@staticmethod
|
||||
def _convert_x0_to_flow_pred(scheduler, x0_pred: torch.Tensor, xt: torch.Tensor, timestep: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Convert x0 prediction to flow matching's prediction.
|
||||
x0_pred: the x0 prediction with shape [B, C, H, W]
|
||||
xt: the input noisy data with shape [B, C, H, W]
|
||||
timestep: the timestep with shape [B]
|
||||
|
||||
pred = (x_t - x_0) / sigma_t
|
||||
"""
|
||||
# use higher precision for calculations
|
||||
original_dtype = x0_pred.dtype
|
||||
x0_pred, xt, sigmas, timesteps = map(
|
||||
lambda x: x.double().to(x0_pred.device), [x0_pred, xt,
|
||||
scheduler.sigmas,
|
||||
scheduler.timesteps]
|
||||
)
|
||||
timestep_id = torch.argmin(
|
||||
(timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
|
||||
sigma_t = sigmas[timestep_id].reshape(-1, 1, 1, 1)
|
||||
flow_pred = (xt - x0_pred) / sigma_t
|
||||
return flow_pred.to(original_dtype)
|
||||
|
||||
|
||||
def forward(self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
|
||||
timestep: torch.LongTensor,
|
||||
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor]
|
||||
| None = None,
|
||||
kv_cache = None,
|
||||
crossattn_cache = None,
|
||||
current_start = None,
|
||||
**kwargs) -> torch.Tensor:
|
||||
assert encoder_hidden_states_image is None, "encoder_hidden_states_image is not supported"
|
||||
|
||||
return self._forward(
|
||||
noisy_image_or_video=hidden_states,
|
||||
conditional_dict={'prompt_embeds': encoder_hidden_states},
|
||||
timestep=timestep,
|
||||
kv_cache=kv_cache,
|
||||
crossattn_cache=crossattn_cache,
|
||||
current_start=current_start,
|
||||
)
|
||||
|
||||
|
||||
def _forward(
|
||||
self,
|
||||
noisy_image_or_video: torch.Tensor, conditional_dict: dict,
|
||||
timestep: torch.Tensor, kv_cache: Optional[List[dict]] = None,
|
||||
crossattn_cache: Optional[List[dict]] = None,
|
||||
current_start: Optional[int] = None,
|
||||
classify_mode: Optional[bool] = False,
|
||||
concat_time_embeddings: Optional[bool] = False,
|
||||
clean_x: Optional[torch.Tensor] = None,
|
||||
aug_t: Optional[torch.Tensor] = None,
|
||||
cache_start: Optional[int] = None
|
||||
) -> torch.Tensor:
|
||||
noisy_image_or_video = noisy_image_or_video.permute(0, 2, 1, 3, 4)
|
||||
prompt_embeds = conditional_dict["prompt_embeds"]
|
||||
|
||||
# [B, F] -> [B]
|
||||
print(f"timestep: {timestep}")
|
||||
print(f"self.uniform_timestep: {self.uniform_timestep}")
|
||||
print(f"timestep.ndim: {timestep.ndim}")
|
||||
print(f"timestep.shape: {timestep.shape}")
|
||||
print(f"noisy_image_or_video.shape: {noisy_image_or_video.shape}")
|
||||
if self.uniform_timestep:
|
||||
# input_timestep = timestep[:, 0]
|
||||
# if timestep.ndim == 1:
|
||||
# input_timestep = timestep.unsqueeze(0)
|
||||
# else:
|
||||
input_timestep = timestep
|
||||
pass
|
||||
else:
|
||||
if timestep.ndim == 1:
|
||||
print(f"not uniform timestep, timestep.ndim == 1")
|
||||
input_timestep = timestep.unsqueeze(0)
|
||||
else:
|
||||
input_timestep = timestep
|
||||
|
||||
logits = None
|
||||
# X0 prediction
|
||||
if kv_cache is not None:
|
||||
flow_pred = self.model(
|
||||
noisy_image_or_video.permute(0, 2, 1, 3, 4),
|
||||
t=input_timestep, context=prompt_embeds,
|
||||
seq_len=self.seq_len,
|
||||
kv_cache=kv_cache,
|
||||
crossattn_cache=crossattn_cache,
|
||||
current_start=current_start,
|
||||
cache_start=cache_start
|
||||
).permute(0, 2, 1, 3, 4)
|
||||
else:
|
||||
if clean_x is not None:
|
||||
# teacher forcing
|
||||
flow_pred = self.model(
|
||||
noisy_image_or_video.permute(0, 2, 1, 3, 4),
|
||||
t=input_timestep, context=prompt_embeds,
|
||||
seq_len=self.seq_len,
|
||||
clean_x=clean_x.permute(0, 2, 1, 3, 4),
|
||||
aug_t=aug_t,
|
||||
).permute(0, 2, 1, 3, 4)
|
||||
else:
|
||||
if classify_mode:
|
||||
flow_pred, logits = self.model(
|
||||
noisy_image_or_video.permute(0, 2, 1, 3, 4),
|
||||
t=input_timestep, context=prompt_embeds,
|
||||
seq_len=self.seq_len,
|
||||
classify_mode=True,
|
||||
register_tokens=self._register_tokens,
|
||||
cls_pred_branch=self._cls_pred_branch,
|
||||
gan_ca_blocks=self._gan_ca_blocks,
|
||||
concat_time_embeddings=concat_time_embeddings
|
||||
)
|
||||
flow_pred = flow_pred.permute(0, 2, 1, 3, 4)
|
||||
else:
|
||||
flow_pred = self.model(
|
||||
noisy_image_or_video.permute(0, 2, 1, 3, 4),
|
||||
t=input_timestep, context=prompt_embeds,
|
||||
seq_len=self.seq_len
|
||||
).permute(0, 2, 1, 3, 4)
|
||||
|
||||
# pred_x0 = self._convert_flow_pred_to_x0(
|
||||
# flow_pred=flow_pred.flatten(0, 1),
|
||||
# xt=noisy_image_or_video.flatten(0, 1),
|
||||
# timestep=timestep.flatten(0, 1)
|
||||
# ).unflatten(0, flow_pred.shape[:2])
|
||||
|
||||
if logits is not None:
|
||||
return flow_pred.permute(0, 2, 1, 3, 4), pred_x0.permute(0, 2, 1, 3, 4), logits
|
||||
|
||||
return flow_pred.permute(0, 2, 1, 3, 4)
|
||||
# return flow_pred.permute(0, 2, 1, 3, 4), pred_x0.permute(0, 2, 1, 3, 4)
|
||||
|
||||
def get_scheduler(self) -> SchedulerInterface:
|
||||
"""
|
||||
Update the current scheduler with the interface's static method
|
||||
"""
|
||||
scheduler = self.scheduler
|
||||
scheduler.convert_x0_to_noise = types.MethodType(
|
||||
SchedulerInterface.convert_x0_to_noise, scheduler)
|
||||
scheduler.convert_noise_to_x0 = types.MethodType(
|
||||
SchedulerInterface.convert_noise_to_x0, scheduler)
|
||||
scheduler.convert_velocity_to_x0 = types.MethodType(
|
||||
SchedulerInterface.convert_velocity_to_x0, scheduler)
|
||||
self.scheduler = scheduler
|
||||
return scheduler
|
||||
|
||||
def post_init(self):
|
||||
"""
|
||||
A few custom initialization steps that should be called after the object is created.
|
||||
Currently, the only one we have is to bind a few methods to scheduler.
|
||||
We can gradually add more methods here if needed.
|
||||
"""
|
||||
self.get_scheduler()
|
||||
@@ -1,292 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import os
|
||||
import math
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
from fastvideo.wan.modules.causal_model import CausalWanModel
|
||||
|
||||
from fastvideo.configs.pipelines import PipelineConfig
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.loader.component_loader import TransformerLoader
|
||||
from fastvideo.models.dits.causal_wanvideo import CausalWanTransformer3DModel
|
||||
from fastvideo.utils import maybe_download_model
|
||||
from fastvideo.configs.models.dits import WanVideoConfig
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
os.environ["MASTER_ADDR"] = "localhost"
|
||||
os.environ["MASTER_PORT"] = "29503"
|
||||
|
||||
BASE_MODEL_PATH = "wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
|
||||
local_dir=os.path.join(
|
||||
'data', BASE_MODEL_PATH))
|
||||
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("distributed_setup")
|
||||
def test_ori_causal_wan_transformer():
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
precision = torch.bfloat16
|
||||
precision_str = "bf16"
|
||||
args = FastVideoArgs(model_path=TRANSFORMER_PATH,
|
||||
dit_cpu_offload=True,
|
||||
pipeline_config=PipelineConfig(dit_config=WanVideoConfig(), dit_precision=precision_str))
|
||||
args.device = device
|
||||
|
||||
loader = TransformerLoader()
|
||||
model2 = loader.load(TRANSFORMER_PATH, args).to(dtype=precision)
|
||||
|
||||
model1 = CausalWanModel.from_pretrained(
|
||||
"/mnt/weka/home/hao.zhang/wei/Self-Forcing/wan_models/Wan2.1-T2V-1.3B", device=device,
|
||||
torch_dtype=precision).to(device, dtype=precision).requires_grad_(False)
|
||||
causal_state_dict = torch.load("/mnt/weka/home/hao.zhang/wei/Self-Forcing/checkpoints/self_forcing_dmd.pt")["generator_ema"]
|
||||
new_state_dict = {}
|
||||
for k, v in causal_state_dict.items():
|
||||
if k.startswith("model."):
|
||||
new_state_dict[k.replace("model.", "")] = v
|
||||
causal_state_dict = new_state_dict
|
||||
model1.load_state_dict(causal_state_dict)
|
||||
|
||||
total_params = sum(p.numel() for p in model1.parameters())
|
||||
# Calculate weight sum for model1 (converting to float64 to avoid overflow)
|
||||
weight_sum_model1 = sum(
|
||||
p.to(torch.float64).sum().item() for p in model1.parameters())
|
||||
# Also calculate mean for more stable comparison
|
||||
weight_mean_model1 = weight_sum_model1 / total_params
|
||||
logger.info("Model 1 weight sum: %s", weight_sum_model1)
|
||||
logger.info("Model 1 weight mean: %s", weight_mean_model1)
|
||||
|
||||
# Calculate weight sum for model2 (converting to float64 to avoid overflow)
|
||||
total_params_model2 = sum(p.numel() for p in model2.parameters())
|
||||
weight_sum_model2 = sum(
|
||||
p.to(torch.float64).sum().item() for p in model2.parameters())
|
||||
# Also calculate mean for more stable comparison
|
||||
weight_mean_model2 = weight_sum_model2 / total_params_model2
|
||||
logger.info("Model 2 weight sum: %s", weight_sum_model2)
|
||||
logger.info("Model 2 weight mean: %s", weight_mean_model2)
|
||||
|
||||
weight_sum_diff = abs(weight_sum_model1 - weight_sum_model2)
|
||||
logger.info("Weight sum difference: %s", weight_sum_diff)
|
||||
weight_mean_diff = abs(weight_mean_model1 - weight_mean_model2)
|
||||
logger.info("Weight mean difference: %s", weight_mean_diff)
|
||||
|
||||
# Set both models to eval mode
|
||||
model1 = model1.eval()
|
||||
model2 = model2.eval()
|
||||
|
||||
# Create identical inputs for both models
|
||||
batch_size = 1
|
||||
text_seq_len = 30
|
||||
|
||||
# Video latents [B, C, T, H, W]
|
||||
hidden_states = torch.randn(batch_size,
|
||||
16,
|
||||
12,
|
||||
160,
|
||||
90,
|
||||
device=device,
|
||||
dtype=precision)
|
||||
block_sizes = [3 for _ in range(4)]
|
||||
timesteps = [1000, 750, 500, 250]
|
||||
|
||||
# Text embeddings [B, L, D] (including global token)
|
||||
encoder_hidden_states = torch.randn(batch_size,
|
||||
text_seq_len + 1,
|
||||
4096,
|
||||
device=device,
|
||||
dtype=precision)
|
||||
|
||||
output1 = _causal_inference(model1, hidden_states.clone(), encoder_hidden_states.clone(), block_sizes, timesteps, precision)
|
||||
logger.info("Finish inference for model1")
|
||||
output2 = _causal_inference(model2, hidden_states.clone(), encoder_hidden_states.clone(), block_sizes, timesteps, precision)
|
||||
|
||||
# Check if outputs have the same shape
|
||||
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
|
||||
assert output1.dtype == output2.dtype, f"Output dtype don't match: {output1.dtype} vs {output2.dtype}"
|
||||
|
||||
logger.info("Output 1 Sum: %s", output1.float().sum().item())
|
||||
logger.info("Output 2 Sum: %s", output2.float().sum().item())
|
||||
|
||||
# Check if outputs are similar (allowing for small numerical differences)
|
||||
max_diff = torch.max(torch.abs(output1 - output2))
|
||||
mean_diff = torch.mean(torch.abs(output1 - output2))
|
||||
logger.info("Max Diff: %s", max_diff.item())
|
||||
logger.info("Mean Diff: %s", mean_diff.item())
|
||||
assert max_diff < 1e-4, f"Maximum difference between outputs: {max_diff.item()}"
|
||||
# mean diff
|
||||
assert mean_diff < 1e-4, f"Mean difference between outputs: {mean_diff.item()}"
|
||||
|
||||
def _causal_inference(transformer, latents, prompt_embeds, block_sizes, timesteps, target_dtype):
|
||||
forward_batch = ForwardBatch(
|
||||
data_type="dummy",
|
||||
)
|
||||
start_index = 0
|
||||
pos_start_base = 0
|
||||
frame_seq_length = latents.shape[-1] * latents.shape[-2] // (WanVideoConfig().arch_config.patch_size[-1] * WanVideoConfig().arch_config.patch_size[-2])
|
||||
seq_len = frame_seq_length * latents.shape[2]
|
||||
kv_cache1 = _initialize_kv_cache(transformer, batch_size=latents.shape[0],
|
||||
kv_cache_size=frame_seq_length * latents.shape[2],
|
||||
dtype=target_dtype,
|
||||
device=latents.device)
|
||||
crossattn_cache = _initialize_crossattn_cache(
|
||||
transformer,
|
||||
batch_size=latents.shape[0],
|
||||
max_text_len=WanVideoConfig().arch_config.text_len,
|
||||
dtype=target_dtype,
|
||||
device=latents.device)
|
||||
for current_num_frames, t_cur in zip(block_sizes, timesteps):
|
||||
# logger.info(f"Current frame idx: {start_index}, Current timestep: {t_cur}")
|
||||
# logger.info(f"k cache sum: {sum(kv_cache['k'].float().sum().item() for kv_cache in kv_cache1)}, v cache sum: {sum(kv_cache['v'].float().sum().item() for kv_cache in kv_cache1)}")
|
||||
# logger.info(f"latents sum: {latents.float().sum().item()}, encoder_hidden_states sum: {prompt_embeds.float().sum().item()}")
|
||||
current_latents = latents[:, :, start_index:start_index +
|
||||
current_num_frames, :, :]
|
||||
|
||||
attn_metadata = None
|
||||
|
||||
with set_forward_context(current_timestep=0,
|
||||
attn_metadata=attn_metadata,
|
||||
forward_batch=forward_batch):
|
||||
# Run transformer; follow DMD stage pattern
|
||||
t_expanded_noise = t_cur * torch.ones(
|
||||
(current_latents.shape[0], 1),
|
||||
device=current_latents.device,
|
||||
dtype=torch.long)
|
||||
if isinstance(transformer, CausalWanModel):
|
||||
pred_noise_btchw = transformer(
|
||||
x=current_latents,
|
||||
context=prompt_embeds,
|
||||
t=t_expanded_noise,
|
||||
seq_len=seq_len,
|
||||
kv_cache=kv_cache1,
|
||||
crossattn_cache=crossattn_cache,
|
||||
current_start=(pos_start_base + start_index) *
|
||||
frame_seq_length
|
||||
)
|
||||
elif isinstance(transformer, CausalWanTransformer3DModel):
|
||||
pred_noise_btchw = transformer(
|
||||
current_latents,
|
||||
prompt_embeds,
|
||||
t_expanded_noise,
|
||||
kv_cache=kv_cache1,
|
||||
crossattn_cache=crossattn_cache,
|
||||
current_start=(pos_start_base + start_index) *
|
||||
frame_seq_length,
|
||||
start_frame=start_index
|
||||
)
|
||||
|
||||
# Write back and advance
|
||||
latents[:, :, start_index:start_index +
|
||||
current_num_frames, :, :] = pred_noise_btchw.clone()
|
||||
|
||||
# Re-run with context timestep to update KV cache using clean context
|
||||
context_noise = 0
|
||||
t_context = torch.ones([latents.shape[0]],
|
||||
device=latents.device,
|
||||
dtype=torch.long) * int(context_noise)
|
||||
context_bcthw = pred_noise_btchw.to(target_dtype)
|
||||
with set_forward_context(current_timestep=0,
|
||||
attn_metadata=attn_metadata,
|
||||
forward_batch=forward_batch):
|
||||
t_expanded_context = t_context.unsqueeze(1)
|
||||
if isinstance(transformer, CausalWanModel):
|
||||
_ = transformer(
|
||||
x=context_bcthw,
|
||||
context=prompt_embeds,
|
||||
t=t_expanded_context,
|
||||
seq_len=seq_len,
|
||||
kv_cache=kv_cache1,
|
||||
crossattn_cache=crossattn_cache,
|
||||
current_start=(pos_start_base + start_index) *
|
||||
frame_seq_length
|
||||
)
|
||||
elif isinstance(transformer, CausalWanTransformer3DModel):
|
||||
_ = transformer(
|
||||
context_bcthw,
|
||||
prompt_embeds,
|
||||
t_expanded_context,
|
||||
kv_cache=kv_cache1,
|
||||
crossattn_cache=crossattn_cache,
|
||||
current_start=(pos_start_base + start_index) *
|
||||
frame_seq_length,
|
||||
start_frame=start_index
|
||||
)
|
||||
start_index += current_num_frames
|
||||
|
||||
return latents
|
||||
|
||||
def _initialize_kv_cache(transformer, batch_size, kv_cache_size, dtype, device) -> None:
|
||||
"""
|
||||
Initialize a Per-GPU KV cache aligned with the Wan model assumptions.
|
||||
"""
|
||||
kv_cache1 = []
|
||||
if isinstance(transformer, CausalWanModel):
|
||||
num_attention_heads = transformer.num_heads
|
||||
attention_head_dim = transformer.dim // transformer.num_heads
|
||||
elif isinstance(transformer, CausalWanTransformer3DModel):
|
||||
num_attention_heads = transformer.num_attention_heads
|
||||
attention_head_dim = transformer.attention_head_dim
|
||||
|
||||
for _ in range(len(transformer.blocks)):
|
||||
kv_cache1.append({
|
||||
"k":
|
||||
torch.zeros([
|
||||
batch_size, kv_cache_size, num_attention_heads,
|
||||
attention_head_dim
|
||||
],
|
||||
dtype=dtype,
|
||||
device=device),
|
||||
"v":
|
||||
torch.zeros([
|
||||
batch_size, kv_cache_size, num_attention_heads,
|
||||
attention_head_dim
|
||||
],
|
||||
dtype=dtype,
|
||||
device=device),
|
||||
"global_end_index":
|
||||
torch.tensor([0], dtype=torch.long, device=device),
|
||||
"local_end_index":
|
||||
torch.tensor([0], dtype=torch.long, device=device),
|
||||
})
|
||||
|
||||
return kv_cache1
|
||||
|
||||
def _initialize_crossattn_cache(transformer, batch_size, max_text_len, dtype,
|
||||
device) -> None:
|
||||
"""
|
||||
Initialize a Per-GPU cross-attention cache aligned with the Wan model assumptions.
|
||||
"""
|
||||
crossattn_cache = []
|
||||
if isinstance(transformer, CausalWanModel):
|
||||
num_attention_heads = transformer.num_heads
|
||||
attention_head_dim = transformer.dim // transformer.num_heads
|
||||
elif isinstance(transformer, CausalWanTransformer3DModel):
|
||||
num_attention_heads = transformer.num_attention_heads
|
||||
attention_head_dim = transformer.attention_head_dim
|
||||
|
||||
for _ in range(len(transformer.blocks)):
|
||||
crossattn_cache.append({
|
||||
"k":
|
||||
torch.zeros([
|
||||
batch_size, max_text_len, num_attention_heads,
|
||||
attention_head_dim
|
||||
],
|
||||
dtype=dtype,
|
||||
device=device),
|
||||
"v":
|
||||
torch.zeros([
|
||||
batch_size, max_text_len, num_attention_heads,
|
||||
attention_head_dim
|
||||
],
|
||||
dtype=dtype,
|
||||
device=device),
|
||||
"is_init":
|
||||
False,
|
||||
})
|
||||
return crossattn_cache
|
||||
@@ -1,133 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import os
|
||||
import math
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
from fastvideo.wan.modules.model import WanModel
|
||||
|
||||
from fastvideo.configs.pipelines import PipelineConfig
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.loader.component_loader import TransformerLoader
|
||||
from fastvideo.utils import maybe_download_model
|
||||
from fastvideo.configs.models.dits import WanVideoConfig
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
os.environ["MASTER_ADDR"] = "localhost"
|
||||
os.environ["MASTER_PORT"] = "29503"
|
||||
|
||||
BASE_MODEL_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
|
||||
local_dir=os.path.join(
|
||||
'data', BASE_MODEL_PATH))
|
||||
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("distributed_setup")
|
||||
def test_ori_wan_transformer():
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
precision = torch.bfloat16
|
||||
precision_str = "bf16"
|
||||
args = FastVideoArgs(model_path=TRANSFORMER_PATH,
|
||||
dit_cpu_offload=True,
|
||||
pipeline_config=PipelineConfig(dit_config=WanVideoConfig(), dit_precision=precision_str))
|
||||
args.device = device
|
||||
|
||||
loader = TransformerLoader()
|
||||
model2 = loader.load(TRANSFORMER_PATH, args).to(dtype=precision)
|
||||
|
||||
model1 = WanModel.from_pretrained(
|
||||
"/mnt/weka/home/hao.zhang/wei/Self-Forcing/wan_models/Wan2.1-T2V-1.3B", device=device,
|
||||
torch_dtype=precision).to(device, dtype=precision).requires_grad_(False)
|
||||
|
||||
total_params = sum(p.numel() for p in model1.parameters())
|
||||
# Calculate weight sum for model1 (converting to float64 to avoid overflow)
|
||||
weight_sum_model1 = sum(
|
||||
p.to(torch.float64).sum().item() for p in model1.parameters())
|
||||
# Also calculate mean for more stable comparison
|
||||
weight_mean_model1 = weight_sum_model1 / total_params
|
||||
logger.info("Model 1 weight sum: %s", weight_sum_model1)
|
||||
logger.info("Model 1 weight mean: %s", weight_mean_model1)
|
||||
|
||||
# Calculate weight sum for model2 (converting to float64 to avoid overflow)
|
||||
total_params_model2 = sum(p.numel() for p in model2.parameters())
|
||||
weight_sum_model2 = sum(
|
||||
p.to(torch.float64).sum().item() for p in model2.parameters())
|
||||
# Also calculate mean for more stable comparison
|
||||
weight_mean_model2 = weight_sum_model2 / total_params_model2
|
||||
logger.info("Model 2 weight sum: %s", weight_sum_model2)
|
||||
logger.info("Model 2 weight mean: %s", weight_mean_model2)
|
||||
|
||||
weight_sum_diff = abs(weight_sum_model1 - weight_sum_model2)
|
||||
logger.info("Weight sum difference: %s", weight_sum_diff)
|
||||
weight_mean_diff = abs(weight_mean_model1 - weight_mean_model2)
|
||||
logger.info("Weight mean difference: %s", weight_mean_diff)
|
||||
|
||||
# Set both models to eval mode
|
||||
model1 = model1.eval()
|
||||
model2 = model2.eval()
|
||||
|
||||
# Create identical inputs for both models
|
||||
batch_size = 1
|
||||
text_seq_len = 30
|
||||
seq_len = math.ceil((160 * 90) /
|
||||
(2 * 2) *
|
||||
21)
|
||||
|
||||
# Video latents [B, C, T, H, W]
|
||||
hidden_states = torch.randn(batch_size,
|
||||
16,
|
||||
21,
|
||||
160,
|
||||
90,
|
||||
device=device,
|
||||
dtype=precision)
|
||||
|
||||
# Text embeddings [B, L, D] (including global token)
|
||||
encoder_hidden_states = torch.randn(batch_size,
|
||||
text_seq_len + 1,
|
||||
4096,
|
||||
device=device,
|
||||
dtype=precision)
|
||||
|
||||
# Timestep
|
||||
timestep = torch.tensor([500], device=device, dtype=precision)
|
||||
|
||||
forward_batch = ForwardBatch(
|
||||
data_type="dummy",
|
||||
)
|
||||
|
||||
# with torch.amp.autocast('cuda', dtype=precision):
|
||||
output1 = model1(
|
||||
x=hidden_states,
|
||||
context=encoder_hidden_states,
|
||||
t=timestep,
|
||||
seq_len=seq_len,
|
||||
)
|
||||
with set_forward_context(
|
||||
current_timestep=0,
|
||||
attn_metadata=None,
|
||||
forward_batch=forward_batch,
|
||||
):
|
||||
output2 = model2(hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
timestep=timestep)
|
||||
|
||||
# Check if outputs have the same shape
|
||||
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
|
||||
assert output1.dtype == output2.dtype, f"Output dtype don't match: {output1.dtype} vs {output2.dtype}"
|
||||
|
||||
# Check if outputs are similar (allowing for small numerical differences)
|
||||
max_diff = torch.max(torch.abs(output1 - output2))
|
||||
mean_diff = torch.mean(torch.abs(output1 - output2))
|
||||
logger.info("Max Diff: %s", max_diff.item())
|
||||
logger.info("Mean Diff: %s", mean_diff.item())
|
||||
assert max_diff < 1e-4, f"Maximum difference between outputs: {max_diff.item()}"
|
||||
# mean diff
|
||||
assert mean_diff < 1e-4, f"Mean difference between outputs: {mean_diff.item()}"
|
||||
@@ -1,144 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import os
|
||||
import math
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
from fastvideo.wan.modules.causal_model import CausalWanModel
|
||||
|
||||
from fastvideo.configs.pipelines import PipelineConfig
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.loader.component_loader import TransformerLoader
|
||||
from fastvideo.utils import maybe_download_model
|
||||
from fastvideo.configs.models.dits import WanVideoConfig
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
os.environ["MASTER_ADDR"] = "localhost"
|
||||
os.environ["MASTER_PORT"] = "29503"
|
||||
|
||||
BASE_MODEL_PATH = "wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
|
||||
local_dir=os.path.join(
|
||||
'data', BASE_MODEL_PATH))
|
||||
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("distributed_setup")
|
||||
def test_train_ori_causal_wan_transformer():
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
precision = torch.bfloat16
|
||||
precision_str = "bf16"
|
||||
args = FastVideoArgs(model_path=TRANSFORMER_PATH,
|
||||
dit_cpu_offload=True,
|
||||
pipeline_config=PipelineConfig(dit_config=WanVideoConfig(), dit_precision=precision_str))
|
||||
args.device = device
|
||||
|
||||
loader = TransformerLoader()
|
||||
model2 = loader.load(TRANSFORMER_PATH, args).to(dtype=precision)
|
||||
|
||||
model1 = CausalWanModel.from_pretrained(
|
||||
"/mnt/weka/home/hao.zhang/wei/Self-Forcing/wan_models/Wan2.1-T2V-1.3B", device=device,
|
||||
torch_dtype=precision).to(device, dtype=precision).requires_grad_(False)
|
||||
causal_state_dict = torch.load("/mnt/weka/home/hao.zhang/wei/Self-Forcing/checkpoints/self_forcing_dmd.pt")["generator_ema"]
|
||||
new_state_dict = {}
|
||||
for k, v in causal_state_dict.items():
|
||||
if k.startswith("model."):
|
||||
new_state_dict[k.replace("model.", "")] = v
|
||||
causal_state_dict = new_state_dict
|
||||
model1.load_state_dict(causal_state_dict)
|
||||
|
||||
model1.num_frame_per_block = 3
|
||||
model2.num_frame_per_block = 3
|
||||
|
||||
total_params = sum(p.numel() for p in model1.parameters())
|
||||
# Calculate weight sum for model1 (converting to float64 to avoid overflow)
|
||||
weight_sum_model1 = sum(
|
||||
p.to(torch.float64).sum().item() for p in model1.parameters())
|
||||
# Also calculate mean for more stable comparison
|
||||
weight_mean_model1 = weight_sum_model1 / total_params
|
||||
logger.info("Model 1 weight sum: %s", weight_sum_model1)
|
||||
logger.info("Model 1 weight mean: %s", weight_mean_model1)
|
||||
|
||||
# Calculate weight sum for model2 (converting to float64 to avoid overflow)
|
||||
total_params_model2 = sum(p.numel() for p in model2.parameters())
|
||||
weight_sum_model2 = sum(
|
||||
p.to(torch.float64).sum().item() for p in model2.parameters())
|
||||
# Also calculate mean for more stable comparison
|
||||
weight_mean_model2 = weight_sum_model2 / total_params_model2
|
||||
logger.info("Model 2 weight sum: %s", weight_sum_model2)
|
||||
logger.info("Model 2 weight mean: %s", weight_mean_model2)
|
||||
|
||||
weight_sum_diff = abs(weight_sum_model1 - weight_sum_model2)
|
||||
logger.info("Weight sum difference: %s", weight_sum_diff)
|
||||
weight_mean_diff = abs(weight_mean_model1 - weight_mean_model2)
|
||||
logger.info("Weight mean difference: %s", weight_mean_diff)
|
||||
|
||||
# Set both models to eval mode
|
||||
model1 = model1.eval()
|
||||
model2 = model2.eval()
|
||||
|
||||
# Create identical inputs for both models
|
||||
batch_size = 1
|
||||
text_seq_len = 30
|
||||
seq_len = math.ceil((160 * 90) /
|
||||
(2 * 2) *
|
||||
21)
|
||||
|
||||
# Video latents [B, C, T, H, W]
|
||||
hidden_states = torch.randn(batch_size,
|
||||
16,
|
||||
21,
|
||||
160,
|
||||
90,
|
||||
device=device,
|
||||
dtype=precision)
|
||||
|
||||
# Text embeddings [B, L, D] (including global token)
|
||||
encoder_hidden_states = torch.randn(batch_size,
|
||||
text_seq_len + 1,
|
||||
4096,
|
||||
device=device,
|
||||
dtype=precision)
|
||||
|
||||
# Timestep
|
||||
timestep = torch.randint(0, 1000, (batch_size, 21), device=device, dtype=torch.long)
|
||||
logger.info("timestep: %s", timestep)
|
||||
|
||||
forward_batch = ForwardBatch(
|
||||
data_type="dummy",
|
||||
)
|
||||
|
||||
# with torch.amp.autocast('cuda', dtype=precision):
|
||||
output1 = model1(
|
||||
x=hidden_states,
|
||||
context=encoder_hidden_states,
|
||||
t=timestep,
|
||||
seq_len=seq_len,
|
||||
)
|
||||
with set_forward_context(
|
||||
current_timestep=0,
|
||||
attn_metadata=None,
|
||||
forward_batch=forward_batch,
|
||||
):
|
||||
output2 = model2(hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
timestep=timestep)
|
||||
|
||||
# Check if outputs have the same shape
|
||||
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
|
||||
assert output1.dtype == output2.dtype, f"Output dtype don't match: {output1.dtype} vs {output2.dtype}"
|
||||
|
||||
# Check if outputs are similar (allowing for small numerical differences)
|
||||
max_diff = torch.max(torch.abs(output1 - output2))
|
||||
mean_diff = torch.mean(torch.abs(output1 - output2))
|
||||
logger.info("Max Diff: %s", max_diff.item())
|
||||
logger.info("Mean Diff: %s", mean_diff.item())
|
||||
assert max_diff < 1e-4, f"Maximum difference between outputs: {max_diff.item()}"
|
||||
# mean diff
|
||||
assert mean_diff < 1e-4, f"Mean difference between outputs: {mean_diff.item()}"
|
||||
@@ -1,4 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from torch.distributed.fsdp import (CPUOffloadPolicy, FSDPModule,
|
||||
MixedPrecisionPolicy, fully_shard)
|
||||
from torch.distributed import DeviceMesh, init_device_mesh
|
||||
import copy
|
||||
import gc
|
||||
import json
|
||||
@@ -23,9 +26,11 @@ from tqdm.auto import tqdm
|
||||
import fastvideo.envs as envs
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.dataset.validation_dataset import ValidationDataset
|
||||
from fastvideo.models.hf_transformer_utils import get_diffusers_config
|
||||
from fastvideo.distributed import (cleanup_dist_env_and_memory,
|
||||
get_local_torch_device, get_sp_group,
|
||||
get_world_group)
|
||||
from fastvideo.models.loader.fsdp_load import shard_model
|
||||
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.logger import init_logger
|
||||
@@ -43,6 +48,8 @@ from fastvideo.training.training_utils import (
|
||||
shift_timestep)
|
||||
from fastvideo.utils import (is_vsa_available, maybe_download_model,
|
||||
set_random_seed, verify_model_config_and_directory)
|
||||
from fastvideo.sf_utils.wan_wrapper import WanDiffusionWrapper
|
||||
from fastvideo.sf_utils.distributed import fsdp_wrap
|
||||
|
||||
import wandb # isort: skip
|
||||
|
||||
@@ -209,6 +216,84 @@ class DistillationPipeline(TrainingPipeline):
|
||||
training_args: "TrainingArgs"):
|
||||
"""
|
||||
Load a module from a specific path using the same loading logic as the pipeline.
|
||||
"""
|
||||
if training_args.use_sf_wan:
|
||||
return self.load_sf_wan_module_from_path(model_path, module_type, training_args)
|
||||
else:
|
||||
return self.load_fastvideo_module_from_path(model_path, module_type, training_args)
|
||||
|
||||
def load_sf_wan_module_from_path(self, model_path: str, module_type: str,
|
||||
training_args: "TrainingArgs", use_ode_init: bool = False):
|
||||
"""
|
||||
Load a module from a specific path using the same loading logic as the pipeline.
|
||||
"""
|
||||
# get config.json from model_path
|
||||
local_model_path = maybe_download_model(model_path)
|
||||
model_path = os.path.join(local_model_path, 'transformer')
|
||||
config = get_diffusers_config(model=model_path)
|
||||
config.pop('_class_name')
|
||||
dit_config = training_args.pipeline_config.dit_config
|
||||
dit_config.update_model_arch(config)
|
||||
|
||||
# convert to WanDiffusionWrapper format
|
||||
if "1.3B" in model_path:
|
||||
model_name = "Wan2.1-T2V-1.3B"
|
||||
# assert not use_ode_init, "ODE init is not supported for 1.3B model"
|
||||
elif "14B" in model_path:
|
||||
model_name = "Wan2.1-T2V-14B"
|
||||
assert not use_ode_init, "ODE init is not supported for 14B model"
|
||||
else:
|
||||
raise ValueError(f"Unsupported SF Wan model name: {model_path}")
|
||||
is_causal = use_ode_init
|
||||
logger.info(f"Loading SF Wan model: {model_name}")
|
||||
module = WanDiffusionWrapper(
|
||||
model_name=model_name, timestep_shift=5.0, is_causal=is_causal, config=dit_config)
|
||||
if use_ode_init:
|
||||
logger.info(f"Loading ODE init weights from: {training_args.sf_ode_init_path}")
|
||||
state_dict = torch.load(training_args.sf_ode_init_path)
|
||||
if 'generator_ema' in state_dict:
|
||||
state_dict = state_dict['generator_ema']
|
||||
else:
|
||||
state_dict = state_dict['generator']
|
||||
module.load_state_dict(state_dict)
|
||||
|
||||
|
||||
mp_policy = MixedPrecisionPolicy(torch.bfloat16,
|
||||
torch.float32,
|
||||
None,
|
||||
cast_forward_inputs=False)
|
||||
device_mesh = init_device_mesh(
|
||||
"cuda",
|
||||
mesh_shape=(training_args.hsdp_replicate_dim, training_args.hsdp_shard_dim),
|
||||
mesh_dim_names=("replicate", "shard"),
|
||||
)
|
||||
from fastvideo.configs.models.dits.wanvideo import WanVideoArchConfig
|
||||
config = WanVideoArchConfig()
|
||||
shard_conditions = config._fsdp_shard_conditions
|
||||
shard_model(
|
||||
module,
|
||||
cpu_offload=True,
|
||||
reshard_after_forward=True,
|
||||
mp_policy=mp_policy,
|
||||
mesh=device_mesh,
|
||||
fsdp_shard_conditions=shard_conditions,
|
||||
pin_cpu_memory=True,
|
||||
)
|
||||
# module = fsdp_wrap(
|
||||
# module,
|
||||
# cpu_offload=True,
|
||||
# # sharding_strategy='hybrid_full',
|
||||
# mixed_precision=True,
|
||||
# wrap_strategy='size'
|
||||
# )
|
||||
module.to(get_local_torch_device())
|
||||
|
||||
return module
|
||||
|
||||
def load_fastvideo_module_from_path(self, model_path: str, module_type: str,
|
||||
training_args: "TrainingArgs"):
|
||||
"""
|
||||
Load a module from a specific path using the same loading logic as the pipeline.
|
||||
|
||||
Args:
|
||||
model_path: Path to the model
|
||||
@@ -730,11 +815,12 @@ class DistillationPipeline(TrainingPipeline):
|
||||
"encoder_hidden_states": training_batch.encoder_hidden_states,
|
||||
"encoder_attention_mask": training_batch.encoder_attention_mask,
|
||||
}
|
||||
unconditional_dict = {
|
||||
"encoder_hidden_states": self.negative_prompt_embeds,
|
||||
"encoder_attention_mask": self.negative_prompt_attention_mask,
|
||||
}
|
||||
training_batch.unconditional_dict = unconditional_dict
|
||||
if getattr(self, "negative_prompt_embeds", None) is not None:
|
||||
unconditional_dict = {
|
||||
"encoder_hidden_states": self.negative_prompt_embeds,
|
||||
"encoder_attention_mask": self.negative_prompt_attention_mask,
|
||||
}
|
||||
training_batch.unconditional_dict = unconditional_dict
|
||||
|
||||
training_batch.dmd_latent_vis_dict = {}
|
||||
training_batch.fake_score_latent_vis_dict = {}
|
||||
|
||||
@@ -29,6 +29,7 @@ import numpy as np
|
||||
from fastvideo.utils import set_random_seed, is_vsa_available
|
||||
import fastvideo.envs as envs
|
||||
from einops import rearrange
|
||||
from fastvideo.sf_utils.wan_wrapper import WanDiffusionWrapper
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -72,6 +73,15 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
|
||||
logger.info(f"Self-forcing generator update ratio: {self.dfake_gen_update_ratio}")
|
||||
|
||||
if training_args.use_sf_wan:
|
||||
sf_transformer = self.load_sf_wan_module_from_path(training_args.model_path, "transformer", training_args, use_ode_init=True)
|
||||
del self.transformer
|
||||
self.transformer = sf_transformer
|
||||
self.modules["transformer"] = sf_transformer
|
||||
logger.info("Using self-forcing Wan model for DMD training")
|
||||
else:
|
||||
logger.info("transformer is not a WanDiffusionWrapper")
|
||||
|
||||
def generate_and_sync_list(self, num_blocks, num_denoising_steps, device):
|
||||
"""Generate and synchronize random exit flags across distributed processes."""
|
||||
rank = dist.get_rank() if dist.is_initialized() else 0
|
||||
@@ -319,8 +329,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
encoder_hidden_states_image=training_batch_temp.input_kwargs.get('encoder_hidden_states_image'),
|
||||
kv_cache=self.kv_cache1,
|
||||
crossattn_cache=self.crossattn_cache,
|
||||
current_start=current_start_frame * self.frame_seq_length,
|
||||
start_frame=current_start_frame
|
||||
current_start=current_start_frame * self.frame_seq_length
|
||||
).permute(0, 2, 1, 3, 4)
|
||||
|
||||
denoised_pred = pred_noise_to_pred_video(
|
||||
@@ -349,8 +358,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
encoder_hidden_states_image=training_batch_temp.input_kwargs.get('encoder_hidden_states_image'),
|
||||
kv_cache=self.kv_cache1,
|
||||
crossattn_cache=self.crossattn_cache,
|
||||
current_start=current_start_frame * self.frame_seq_length,
|
||||
start_frame=current_start_frame
|
||||
current_start=current_start_frame * self.frame_seq_length
|
||||
).permute(0, 2, 1, 3, 4)
|
||||
else:
|
||||
training_batch_temp = self._build_distill_input_kwargs(
|
||||
@@ -363,8 +371,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
encoder_hidden_states_image=training_batch_temp.input_kwargs.get('encoder_hidden_states_image'),
|
||||
kv_cache=self.kv_cache1,
|
||||
crossattn_cache=self.crossattn_cache,
|
||||
current_start=current_start_frame * self.frame_seq_length,
|
||||
start_frame=current_start_frame
|
||||
current_start=current_start_frame * self.frame_seq_length
|
||||
).permute(0, 2, 1, 3, 4)
|
||||
|
||||
denoised_pred = pred_noise_to_pred_video(
|
||||
@@ -396,8 +403,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
encoder_hidden_states_image=training_batch_temp.input_kwargs.get('encoder_hidden_states_image'),
|
||||
kv_cache=self.kv_cache1,
|
||||
crossattn_cache=self.crossattn_cache,
|
||||
current_start=current_start_frame * self.frame_seq_length,
|
||||
start_frame=current_start_frame
|
||||
current_start=current_start_frame * self.frame_seq_length
|
||||
)
|
||||
|
||||
# Step 3.4: update the start and end frame indices
|
||||
|
||||
@@ -41,6 +41,7 @@ class WanDistillationPipeline(DistillationPipeline):
|
||||
args_copy = deepcopy(training_args)
|
||||
|
||||
args_copy.inference_mode = True
|
||||
assert self.get_module("transformer") is not None, "transformer is not initialized for validation"
|
||||
validation_pipeline = WanDMDPipeline.from_pretrained(
|
||||
training_args.model_path,
|
||||
args=args_copy, # type: ignore
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
from . import configs, distributed, modules
|
||||
from .image2video import WanI2V
|
||||
from .text2video import WanT2V
|
||||
from . import modules
|
||||
# from . import configs, distributed, modules
|
||||
# from .image2video import WanI2V
|
||||
# from .text2video import WanT2V
|
||||
|
||||
@@ -1,33 +0,0 @@
|
||||
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
|
||||
from functools import partial
|
||||
|
||||
import torch
|
||||
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
|
||||
from torch.distributed.fsdp import MixedPrecision, ShardingStrategy
|
||||
from torch.distributed.fsdp.wrap import lambda_auto_wrap_policy
|
||||
|
||||
|
||||
def shard_model(
|
||||
model,
|
||||
device_id,
|
||||
param_dtype=torch.bfloat16,
|
||||
reduce_dtype=torch.float32,
|
||||
buffer_dtype=torch.float32,
|
||||
process_group=None,
|
||||
sharding_strategy=ShardingStrategy.FULL_SHARD,
|
||||
sync_module_states=True,
|
||||
):
|
||||
model = FSDP(
|
||||
module=model,
|
||||
process_group=process_group,
|
||||
sharding_strategy=sharding_strategy,
|
||||
auto_wrap_policy=partial(
|
||||
lambda_auto_wrap_policy, lambda_fn=lambda m: m in model.blocks),
|
||||
mixed_precision=MixedPrecision(
|
||||
param_dtype=param_dtype,
|
||||
reduce_dtype=reduce_dtype,
|
||||
buffer_dtype=buffer_dtype),
|
||||
device_id=device_id,
|
||||
use_orig_params=True,
|
||||
sync_module_states=sync_module_states)
|
||||
return model
|
||||
@@ -1,192 +0,0 @@
|
||||
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
|
||||
import torch
|
||||
import torch.cuda.amp as amp
|
||||
from xfuser.core.distributed import (get_sequence_parallel_rank,
|
||||
get_sequence_parallel_world_size,
|
||||
get_sp_group)
|
||||
from xfuser.core.long_ctx_attention import xFuserLongContextAttention
|
||||
|
||||
from ..modules.model import sinusoidal_embedding_1d
|
||||
|
||||
|
||||
def pad_freqs(original_tensor, target_len):
|
||||
seq_len, s1, s2 = original_tensor.shape
|
||||
pad_size = target_len - seq_len
|
||||
padding_tensor = torch.ones(
|
||||
pad_size,
|
||||
s1,
|
||||
s2,
|
||||
dtype=original_tensor.dtype,
|
||||
device=original_tensor.device)
|
||||
padded_tensor = torch.cat([original_tensor, padding_tensor], dim=0)
|
||||
return padded_tensor
|
||||
|
||||
|
||||
@amp.autocast(enabled=False)
|
||||
def rope_apply(x, grid_sizes, freqs):
|
||||
"""
|
||||
x: [B, L, N, C].
|
||||
grid_sizes: [B, 3].
|
||||
freqs: [M, C // 2].
|
||||
"""
|
||||
s, n, c = x.size(1), x.size(2), x.size(3) // 2
|
||||
# split freqs
|
||||
freqs = freqs.split([c - 2 * (c // 3), c // 3, c // 3], dim=1)
|
||||
|
||||
# loop over samples
|
||||
output = []
|
||||
for i, (f, h, w) in enumerate(grid_sizes.tolist()):
|
||||
seq_len = f * h * w
|
||||
|
||||
# precompute multipliers
|
||||
x_i = torch.view_as_complex(x[i, :s].to(torch.float64).reshape(
|
||||
s, n, -1, 2))
|
||||
freqs_i = torch.cat([
|
||||
freqs[0][:f].view(f, 1, 1, -1).expand(f, h, w, -1),
|
||||
freqs[1][:h].view(1, h, 1, -1).expand(f, h, w, -1),
|
||||
freqs[2][:w].view(1, 1, w, -1).expand(f, h, w, -1)
|
||||
],
|
||||
dim=-1).reshape(seq_len, 1, -1)
|
||||
|
||||
# apply rotary embedding
|
||||
sp_size = get_sequence_parallel_world_size()
|
||||
sp_rank = get_sequence_parallel_rank()
|
||||
freqs_i = pad_freqs(freqs_i, s * sp_size)
|
||||
s_per_rank = s
|
||||
freqs_i_rank = freqs_i[(sp_rank * s_per_rank):((sp_rank + 1) *
|
||||
s_per_rank), :, :]
|
||||
x_i = torch.view_as_real(x_i * freqs_i_rank).flatten(2)
|
||||
x_i = torch.cat([x_i, x[i, s:]])
|
||||
|
||||
# append to collection
|
||||
output.append(x_i)
|
||||
return torch.stack(output).float()
|
||||
|
||||
|
||||
def usp_dit_forward(
|
||||
self,
|
||||
x,
|
||||
t,
|
||||
context,
|
||||
seq_len,
|
||||
clip_fea=None,
|
||||
y=None,
|
||||
):
|
||||
"""
|
||||
x: A list of videos each with shape [C, T, H, W].
|
||||
t: [B].
|
||||
context: A list of text embeddings each with shape [L, C].
|
||||
"""
|
||||
if self.model_type == 'i2v':
|
||||
assert clip_fea is not None and y is not None
|
||||
# params
|
||||
device = self.patch_embedding.weight.device
|
||||
if self.freqs.device != device:
|
||||
self.freqs = self.freqs.to(device)
|
||||
|
||||
if y is not None:
|
||||
x = [torch.cat([u, v], dim=0) for u, v in zip(x, y)]
|
||||
|
||||
# embeddings
|
||||
x = [self.patch_embedding(u.unsqueeze(0)) for u in x]
|
||||
grid_sizes = torch.stack(
|
||||
[torch.tensor(u.shape[2:], dtype=torch.long) for u in x])
|
||||
x = [u.flatten(2).transpose(1, 2) for u in x]
|
||||
seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.long)
|
||||
assert seq_lens.max() <= seq_len
|
||||
x = torch.cat([
|
||||
torch.cat([u, u.new_zeros(1, seq_len - u.size(1), u.size(2))], dim=1)
|
||||
for u in x
|
||||
])
|
||||
|
||||
# time embeddings
|
||||
with amp.autocast(dtype=torch.float32):
|
||||
e = self.time_embedding(
|
||||
sinusoidal_embedding_1d(self.freq_dim, t).float())
|
||||
e0 = self.time_projection(e).unflatten(1, (6, self.dim))
|
||||
assert e.dtype == torch.float32 and e0.dtype == torch.float32
|
||||
|
||||
# context
|
||||
context_lens = None
|
||||
context = self.text_embedding(
|
||||
torch.stack([
|
||||
torch.cat([u, u.new_zeros(self.text_len - u.size(0), u.size(1))])
|
||||
for u in context
|
||||
]))
|
||||
|
||||
if clip_fea is not None:
|
||||
context_clip = self.img_emb(clip_fea) # bs x 257 x dim
|
||||
context = torch.concat([context_clip, context], dim=1)
|
||||
|
||||
# arguments
|
||||
kwargs = dict(
|
||||
e=e0,
|
||||
seq_lens=seq_lens,
|
||||
grid_sizes=grid_sizes,
|
||||
freqs=self.freqs,
|
||||
context=context,
|
||||
context_lens=context_lens)
|
||||
|
||||
# Context Parallel
|
||||
x = torch.chunk(
|
||||
x, get_sequence_parallel_world_size(),
|
||||
dim=1)[get_sequence_parallel_rank()]
|
||||
|
||||
for block in self.blocks:
|
||||
x = block(x, **kwargs)
|
||||
|
||||
# head
|
||||
x = self.head(x, e)
|
||||
|
||||
# Context Parallel
|
||||
x = get_sp_group().all_gather(x, dim=1)
|
||||
|
||||
# unpatchify
|
||||
x = self.unpatchify(x, grid_sizes)
|
||||
return [u.float() for u in x]
|
||||
|
||||
|
||||
def usp_attn_forward(self,
|
||||
x,
|
||||
seq_lens,
|
||||
grid_sizes,
|
||||
freqs,
|
||||
dtype=torch.bfloat16):
|
||||
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
|
||||
half_dtypes = (torch.float16, torch.bfloat16)
|
||||
|
||||
def half(x):
|
||||
return x if x.dtype in half_dtypes else x.to(dtype)
|
||||
|
||||
# query, key, value function
|
||||
def qkv_fn(x):
|
||||
q = self.norm_q(self.q(x)).view(b, s, n, d)
|
||||
k = self.norm_k(self.k(x)).view(b, s, n, d)
|
||||
v = self.v(x).view(b, s, n, d)
|
||||
return q, k, v
|
||||
|
||||
q, k, v = qkv_fn(x)
|
||||
q = rope_apply(q, grid_sizes, freqs)
|
||||
k = rope_apply(k, grid_sizes, freqs)
|
||||
|
||||
# TODO: We should use unpaded q,k,v for attention.
|
||||
# k_lens = seq_lens // get_sequence_parallel_world_size()
|
||||
# if k_lens is not None:
|
||||
# q = torch.cat([u[:l] for u, l in zip(q, k_lens)]).unsqueeze(0)
|
||||
# k = torch.cat([u[:l] for u, l in zip(k, k_lens)]).unsqueeze(0)
|
||||
# v = torch.cat([u[:l] for u, l in zip(v, k_lens)]).unsqueeze(0)
|
||||
|
||||
x = xFuserLongContextAttention()(
|
||||
None,
|
||||
query=half(q),
|
||||
key=half(k),
|
||||
value=half(v),
|
||||
window_size=self.window_size)
|
||||
|
||||
# TODO: padding after attention.
|
||||
# x = torch.cat([x, x.new_zeros(b, s - x.size(1), n, d)], dim=1)
|
||||
|
||||
# output
|
||||
x = x.flatten(2)
|
||||
x = self.o(x)
|
||||
return x
|
||||
@@ -1,5 +1,5 @@
|
||||
from wan.modules.attention import attention
|
||||
from wan.modules.model import (
|
||||
from fastvideo.wan.modules.attention import attention
|
||||
from fastvideo.wan.modules.model import (
|
||||
WanRMSNorm,
|
||||
rope_apply,
|
||||
WanLayerNorm,
|
||||
@@ -16,6 +16,9 @@ import torch.nn as nn
|
||||
import torch
|
||||
import math
|
||||
import torch.distributed as dist
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# wan 1.3B model has a weird channel / head configurations and require max-autotune to work with flexattention
|
||||
# see https://github.com/pytorch/pytorch/issues/133254
|
||||
@@ -106,8 +109,6 @@ class CausalWanSelfAttention(nn.Module):
|
||||
if cache_start is None:
|
||||
cache_start = current_start
|
||||
|
||||
# print(f"x sum: {x.float().sum().item()}, shape: {x.shape}")
|
||||
|
||||
# query, key, value function
|
||||
def qkv_fn(x):
|
||||
q = self.norm_q(self.q(x)).view(b, s, n, d)
|
||||
@@ -207,7 +208,6 @@ class CausalWanSelfAttention(nn.Module):
|
||||
num_new_tokens = roped_query.shape[1]
|
||||
if self.local_attn_size != -1 and (current_end > kv_cache["global_end_index"].item()) and (
|
||||
num_new_tokens + kv_cache["local_end_index"].item() > kv_cache_size):
|
||||
assert False, "Not implemented"
|
||||
# Calculate the number of new tokens added in this step
|
||||
# Shift existing cache content left to discard oldest tokens
|
||||
# Clone the source slice to avoid overlapping memory error
|
||||
@@ -227,6 +227,13 @@ class CausalWanSelfAttention(nn.Module):
|
||||
# Assign new keys/values directly up to current_end
|
||||
local_end_index = kv_cache["local_end_index"].item() + current_end - kv_cache["global_end_index"].item()
|
||||
local_start_index = local_end_index - num_new_tokens
|
||||
logger.info("local_start_index: %s, local_end_index: %s", local_start_index, local_end_index)
|
||||
logger.info("type(kv_cache['k']): %s", type(kv_cache["k"]))
|
||||
logger.info("type(roped_key): %s", type(roped_key))
|
||||
logger.info("type(kv_cache['v']): %s", type(kv_cache["v"]))
|
||||
logger.info("type(v): %s", type(v))
|
||||
kv_cache["k"] = kv_cache["k"].detach()
|
||||
kv_cache["v"] = kv_cache["v"].detach()
|
||||
kv_cache["k"][:, local_start_index:local_end_index] = roped_key
|
||||
kv_cache["v"][:, local_start_index:local_end_index] = v
|
||||
x = attention(
|
||||
@@ -309,14 +316,9 @@ class CausalWanAttentionBlock(nn.Module):
|
||||
num_frames, frame_seqlen = e.shape[1], x.shape[1] // e.shape[1]
|
||||
# assert e.dtype == torch.float32
|
||||
# with amp.autocast(dtype=torch.float32):
|
||||
# print(f"e sum: {e.float().sum().item()}, dtype: {e.dtype}")
|
||||
|
||||
e = (self.modulation.unsqueeze(1) + e).chunk(6, dim=2)
|
||||
# assert e[0].dtype == torch.float32
|
||||
|
||||
# print(f"e[1] sum: {e[1].float().sum().item()}, dtype: {e[1].dtype}")
|
||||
# print(f"e[0] sum: {e[0].float().sum().item()}, dtype: {e[0].dtype}")
|
||||
|
||||
# self-attention
|
||||
y = self.self_attn(
|
||||
(self.norm1(x).unflatten(dim=1, sizes=(num_frames, frame_seqlen)) * (1 + e[1]) + e[0]).flatten(1, 2),
|
||||
@@ -790,6 +792,10 @@ class CausalWanModel(ModelMixin, ConfigMixin):
|
||||
|
||||
# context
|
||||
context_lens = None
|
||||
# logger.info(f"context: {context}")
|
||||
# for u in context:
|
||||
# logger.info(f"u.shape: {u.shape}")
|
||||
context = [u.squeeze(0) for u in context]
|
||||
context = self.text_embedding(
|
||||
torch.stack([
|
||||
torch.cat(
|
||||
|
||||
@@ -143,14 +143,8 @@ class WanSelfAttention(nn.Module):
|
||||
|
||||
q, k, v = qkv_fn(x)
|
||||
|
||||
print(f"query sum: {torch.sum(q.float()).item()}")
|
||||
|
||||
q = rope_apply(q, grid_sizes, freqs)
|
||||
|
||||
print(f"query after rotary embeddings sum: {torch.sum(q.float()).item()}")
|
||||
|
||||
x = flash_attention(
|
||||
q=q,
|
||||
q=rope_apply(q, grid_sizes, freqs),
|
||||
k=rope_apply(k, grid_sizes, freqs),
|
||||
v=v,
|
||||
k_lens=seq_lens,
|
||||
@@ -159,8 +153,6 @@ class WanSelfAttention(nn.Module):
|
||||
# output
|
||||
x = x.flatten(2)
|
||||
x = self.o(x)
|
||||
|
||||
print(f"attn_output sum: {torch.sum(x.float()).item()}")
|
||||
return x
|
||||
|
||||
|
||||
@@ -343,12 +335,9 @@ class WanAttentionBlock(nn.Module):
|
||||
e = (self.modulation + e).chunk(6, dim=1)
|
||||
# assert e[0].dtype == torch.float32
|
||||
|
||||
|
||||
norm_x = self.norm1(x) * (1 + e[1]) + e[0]
|
||||
print(f"norm_hidden_states sum: {torch.sum(norm_x.float()).item()}")
|
||||
# self-attention
|
||||
y = self.self_attn(
|
||||
norm_x, seq_lens, grid_sizes,
|
||||
self.norm1(x) * (1 + e[1]) + e[0], seq_lens, grid_sizes,
|
||||
freqs)
|
||||
# with amp.autocast(dtype=torch.float32):
|
||||
x = x + y * e[2]
|
||||
|
||||
+4
-1
@@ -49,7 +49,10 @@ dependencies = [
|
||||
"av",
|
||||
|
||||
# Preprocessing Dependencies
|
||||
"torchcodec==0.5.0"
|
||||
"torchcodec==0.5.0",
|
||||
|
||||
# SF WAN
|
||||
"easydict", "ftfy"
|
||||
]
|
||||
|
||||
[tool.uv]
|
||||
|
||||
Reference in New Issue
Block a user