Compare commits

...
Author SHA1 Message Date
JerryZhou54 f42eb84119 fix lint 2025-11-20 07:19:09 +00:00
JerryZhou54 61136292db ckpt 2025-11-20 06:26:42 +00:00
JerryZhou54 b7327f05e1 Same with main 2025-11-20 04:23:20 +00:00
JerryZhou54 5840578c56 Support variable image resolution less than 480*832 2025-11-20 04:17:03 +00:00
RandNMR73 228cb95e92 add demo prompts 2025-11-20 04:15:44 +00:00
JerryZhou54 9cee424b4e Fi lint 2025-11-20 04:13:05 +00:00
RandNMR73 50d97ba225 Add inference for MoE SF 2025-11-20 04:12:37 +00:00
3 changed files with 25 additions and 21 deletions
@@ -26,7 +26,7 @@ def main():
)
sampling_param = SamplingParam.from_pretrained("FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers")
sampling_param.num_frames = 73
sampling_param.num_frames = 81
sampling_param.width = 832
sampling_param.height = 480
sampling_param.seed = 1000
@@ -400,6 +400,10 @@ class CausalDMDDenosingStage(DenoisingStage):
start_index += current_num_frames
if boundary_timestep is not None:
num_frames_to_remove = self.num_frames_per_block - 1
latents = latents[:, :, :-num_frames_to_remove, :, :]
batch.latents = latents
return batch
@@ -490,4 +494,4 @@ class CausalDMDDenosingStage(DenoisingStage):
result.add_check(
"negative_prompt_embeds", batch.negative_prompt_embeds, lambda x:
not batch.do_classifier_free_guidance or V.list_not_empty(x))
return result
return result
+19 -19
View File
@@ -5,6 +5,7 @@ Input validation stage for diffusion pipelines.
import torch
import torchvision.transforms.functional as TF
from PIL import Image
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
@@ -13,6 +14,7 @@ from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.base import PipelineStage
from fastvideo.pipelines.stages.validators import (StageValidators,
VerificationResult)
from fastvideo.utils import best_output_size
logger = init_logger(__name__)
@@ -108,27 +110,25 @@ class InputValidationStage(PipelineStage):
or fastvideo_args.pipeline_config.is_causal
) and batch.pil_image is not None:
img = batch.pil_image
# ih, iw = img.height, img.width
# patch_size = fastvideo_args.pipeline_config.dit_config.arch_config.patch_size
# vae_stride = fastvideo_args.pipeline_config.vae_config.arch_config.scale_factor_spatial
# dh, dw = patch_size[1] * vae_stride, patch_size[2] * vae_stride
# max_area = 720 * 1280
# ow, oh = best_output_size(iw, ih, dw, dh, max_area)
ih, iw = img.height, img.width
patch_size = fastvideo_args.pipeline_config.dit_config.arch_config.patch_size
vae_stride = fastvideo_args.pipeline_config.vae_config.arch_config.scale_factor_spatial
dh, dw = patch_size[1] * vae_stride, patch_size[2] * vae_stride
max_area = 480 * 832
ow, oh = best_output_size(iw, ih, dw, dh, max_area)
# scale = max(ow / iw, oh / ih)
# img = img.resize((round(iw * scale), round(ih * scale)),
# Image.LANCZOS)
# logger.info("resized img height: %s, img width: %s", img.height,
# img.width)
scale = max(ow / iw, oh / ih)
img = img.resize((round(iw * scale), round(ih * scale)),
Image.LANCZOS)
# center-crop
x1 = (img.width - ow) // 2
y1 = (img.height - oh) // 2
img = img.crop((x1, y1, x1 + ow, y1 + oh))
assert img.width == ow and img.height == oh
logger.info("final processed img height: %s, img width: %s",
img.height, img.width)
# # center-crop
# x1 = (img.width - ow) // 2
# y1 = (img.height - oh) // 2
# img = img.crop((x1, y1, x1 + ow, y1 + oh))
# assert img.width == ow and img.height == oh
logger.info("img height: %s, img width: %s", img.height, img.width)
oh = img.height
ow = img.width
# to tensor
img = TF.to_tensor(img).sub_(0.5).div_(0.5).to(
self.device).unsqueeze(1)