Compare commits
13
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
f42eb84119 | ||
|
|
61136292db | ||
|
|
b7327f05e1 | ||
|
|
5840578c56 | ||
|
|
228cb95e92 | ||
|
|
9cee424b4e | ||
|
|
50d97ba225 | ||
|
|
681f1583f9 | ||
|
|
e3b4564d5a | ||
|
|
c0d03fc43d | ||
|
|
404ee8538e | ||
|
|
e57ac59462 | ||
|
|
8c55fdaf7e |
@@ -15,6 +15,7 @@ FastVideo features an end-to-end unified pipeline for accelerating diffusion mod
|
||||
</div>
|
||||
|
||||
## NEWS
|
||||
- ```2025/11/19```: Release [CausalWan2.2 I2V A14B Preview](https://huggingface.co/FastVideo/CausalWan2.2-I2V-A14B-Preview-Diffusers) models, [Blog](https://hao-ai-lab.github.io/blogs/fastvideo_causalwan_preview/) and [Inference Code!](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_self_forcing_causal_wan2_2_i2v.py)
|
||||
- ```2025/08/04```: Release [FastWan](https://hao-ai-lab.github.io/FastVideo/distillation/dmd.html) models and [Sparse-Distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/).
|
||||
- ```2025/06/14```: Release finetuning and inference code for [VSA](https://arxiv.org/pdf/2505.13389)
|
||||
- ```2025/04/24```: [FastVideo V1](https://hao-ai-lab.github.io/blogs/fastvideo/) is released!
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -6,10 +6,8 @@ import time
|
||||
import json
|
||||
from pathlib import Path
|
||||
import tempfile
|
||||
from io import BytesIO
|
||||
|
||||
import gradio as gr
|
||||
from PIL import Image
|
||||
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
@@ -87,39 +85,28 @@ def save_video_from_base64(video_data: str, output_dir: str, prompt: str) -> str
|
||||
return None
|
||||
|
||||
|
||||
def encode_image_to_base64(image_input) -> str:
|
||||
"""Encode an image file path or in-memory image to a base64 string."""
|
||||
if image_input is None:
|
||||
def encode_image_to_base64(image_path: str) -> str:
|
||||
"""Encode an image file to base64 string."""
|
||||
if not image_path or not os.path.exists(image_path):
|
||||
return None
|
||||
|
||||
mime_types = {
|
||||
'.jpg': 'image/jpeg',
|
||||
'.jpeg': 'image/jpeg',
|
||||
'.png': 'image/png',
|
||||
'.gif': 'image/gif',
|
||||
'.webp': 'image/webp',
|
||||
}
|
||||
|
||||
try:
|
||||
if isinstance(image_input, str):
|
||||
if not os.path.exists(image_input):
|
||||
return None
|
||||
|
||||
with open(image_input, 'rb') as f:
|
||||
image_bytes = f.read()
|
||||
|
||||
ext = os.path.splitext(image_input)[1].lower()
|
||||
mime_type = mime_types.get(ext, 'image/jpeg')
|
||||
elif isinstance(image_input, Image.Image):
|
||||
buffer = BytesIO()
|
||||
image_to_save = image_input.convert("RGB")
|
||||
image_to_save.save(buffer, format="PNG")
|
||||
image_bytes = buffer.getvalue()
|
||||
mime_type = 'image/png'
|
||||
else:
|
||||
return None
|
||||
with open(image_path, 'rb') as f:
|
||||
image_bytes = f.read()
|
||||
|
||||
image_base64 = base64.b64encode(image_bytes).decode('utf-8')
|
||||
|
||||
# Determine image type from extension
|
||||
ext = os.path.splitext(image_path)[1].lower()
|
||||
mime_types = {
|
||||
'.jpg': 'image/jpeg',
|
||||
'.jpeg': 'image/jpeg',
|
||||
'.png': 'image/png',
|
||||
'.gif': 'image/gif',
|
||||
'.webp': 'image/webp',
|
||||
}
|
||||
mime_type = mime_types.get(ext, 'image/jpeg')
|
||||
|
||||
return f"data:{mime_type};base64,{image_base64}"
|
||||
|
||||
except Exception as e:
|
||||
@@ -439,7 +426,7 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
|
||||
gr.Markdown("**Please make sure you upload a 480x832 image**")
|
||||
input_image = gr.Image(
|
||||
label="",
|
||||
type="pil",
|
||||
type="filepath",
|
||||
height=400,
|
||||
)
|
||||
|
||||
@@ -525,17 +512,7 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
|
||||
if example_label and example_label in example_labels:
|
||||
index = example_labels.index(example_label)
|
||||
selected_prompt = examples[index]
|
||||
selected_image_path = example_images[index] if index < len(example_images) else None
|
||||
|
||||
if selected_image_path and os.path.exists(selected_image_path):
|
||||
try:
|
||||
with Image.open(selected_image_path) as img:
|
||||
selected_image = img.convert("RGB")
|
||||
except Exception:
|
||||
selected_image = None
|
||||
else:
|
||||
selected_image = None
|
||||
|
||||
selected_image = example_images[index] if index < len(example_images) else None
|
||||
return selected_prompt, selected_image
|
||||
return "", None
|
||||
|
||||
@@ -756,7 +733,6 @@ def main():
|
||||
app,
|
||||
demo,
|
||||
path="/gradio",
|
||||
root_path="/gradio",
|
||||
allowed_paths=[
|
||||
os.path.abspath("outputs"),
|
||||
os.path.abspath("fastvideo-logos"),
|
||||
|
||||
@@ -45,6 +45,7 @@ class PipelineConfig:
|
||||
embedded_cfg_scale: float = 6.0
|
||||
flow_shift: float | None = None
|
||||
disable_autocast: bool = False
|
||||
is_causal: bool = False
|
||||
|
||||
# Model configuration
|
||||
dit_config: DiTConfig = field(default_factory=DiTConfig)
|
||||
|
||||
@@ -101,12 +101,19 @@ class BaseLayerWithLoRA(nn.Module):
|
||||
def set_lora_weights(self,
|
||||
A: torch.Tensor,
|
||||
B: torch.Tensor,
|
||||
lora_alpha: float | None = None,
|
||||
training_mode: bool = False,
|
||||
lora_path: str | None = None) -> None:
|
||||
self.lora_A = torch.nn.Parameter(
|
||||
A) # share storage with weights in the pipeline
|
||||
self.lora_B = torch.nn.Parameter(B)
|
||||
self.disable_lora = False
|
||||
|
||||
# Store rank and alpha directly
|
||||
rank = A.shape[0] # rank is the first dimension of A
|
||||
self.lora_rank = rank
|
||||
self.lora_alpha = int(lora_alpha) if lora_alpha is not None else rank
|
||||
|
||||
if not training_mode:
|
||||
self.merge_lora_weights()
|
||||
self.lora_path = lora_path
|
||||
@@ -134,8 +141,13 @@ class BaseLayerWithLoRA(nn.Module):
|
||||
current_device = self.base_layer.weight.data.device
|
||||
data = self.base_layer.weight.data.to(
|
||||
get_local_torch_device()).full_tensor()
|
||||
data += (self.slice_lora_b_weights(self.lora_B).to(data)
|
||||
@ self.slice_lora_a_weights(self.lora_A).to(data))
|
||||
|
||||
# Apply LoRA with alpha scaling
|
||||
lora_delta = (self.slice_lora_b_weights(self.lora_B).to(data)
|
||||
@ self.slice_lora_a_weights(self.lora_A).to(data))
|
||||
if self.lora_alpha and self.lora_rank and self.lora_alpha != self.lora_rank:
|
||||
lora_delta *= (self.lora_alpha / self.lora_rank)
|
||||
data += lora_delta
|
||||
unsharded_base_layer.weight = nn.Parameter(data.to(current_device))
|
||||
if isinstance(getattr(self.base_layer, "bias", None), DTensor):
|
||||
unsharded_base_layer.bias = nn.Parameter(
|
||||
@@ -154,8 +166,13 @@ class BaseLayerWithLoRA(nn.Module):
|
||||
else:
|
||||
current_device = self.base_layer.weight.data.device
|
||||
data = self.base_layer.weight.data.to(get_local_torch_device())
|
||||
data += \
|
||||
(self.slice_lora_b_weights(self.lora_B.to(data)) @ self.slice_lora_a_weights(self.lora_A.to(data)))
|
||||
|
||||
# Apply LoRA with alpha scaling
|
||||
lora_delta = (self.slice_lora_b_weights(self.lora_B.to(data))
|
||||
@ self.slice_lora_a_weights(self.lora_A.to(data)))
|
||||
if self.lora_alpha and self.lora_rank and self.lora_alpha != self.lora_rank:
|
||||
lora_delta *= (self.lora_alpha / self.lora_rank)
|
||||
data += lora_delta
|
||||
self.base_layer.weight.data = data.to(current_device,
|
||||
non_blocking=True)
|
||||
|
||||
|
||||
@@ -28,7 +28,8 @@ class LoRAPipeline(ComposedPipelineBase):
|
||||
TODO: support training.
|
||||
"""
|
||||
lora_adapters: dict[str, dict[str, torch.Tensor]] = defaultdict(
|
||||
dict) # state dicts of loaded lora adapters
|
||||
dict
|
||||
) # state dicts of loaded lora adapters (includes lora_A, lora_B, and lora_alpha)
|
||||
cur_adapter_name: str = ""
|
||||
cur_adapter_path: str = ""
|
||||
lora_layers: dict[str, BaseLayerWithLoRA] = {}
|
||||
@@ -183,11 +184,26 @@ class LoRAPipeline(ComposedPipelineBase):
|
||||
lora_param_names_mapping_fn = get_param_names_mapping(
|
||||
self.modules["transformer"].lora_param_names_mapping)
|
||||
|
||||
# Extract alpha values and weights in a single pass
|
||||
to_merge_params: defaultdict[Hashable,
|
||||
dict[Any, Any]] = defaultdict(dict)
|
||||
for name, weight in lora_state_dict.items():
|
||||
# Extract weights (lora_A, lora_B, and lora_alpha)
|
||||
name = name.replace("diffusion_model.", "")
|
||||
name = name.replace(".weight", "")
|
||||
|
||||
if "lora_alpha" in name:
|
||||
# Store alpha with minimal mapping - same processing as lora_A/lora_B
|
||||
# but store in lora_adapters with ".lora_alpha" suffix
|
||||
layer_name = name.replace(".lora_alpha", "")
|
||||
layer_name, _, _ = lora_param_names_mapping_fn(layer_name)
|
||||
target_name, _, _ = param_names_mapping_fn(layer_name)
|
||||
# Store alpha alongside weights with same target_name base
|
||||
alpha_key = target_name + ".lora_alpha"
|
||||
self.lora_adapters[lora_nickname][alpha_key] = weight.item(
|
||||
) if weight.numel() == 1 else float(weight.mean())
|
||||
continue
|
||||
|
||||
name, _, _ = lora_param_names_mapping_fn(name)
|
||||
target_name, merge_index, num_params_to_merge = param_names_mapping_fn(
|
||||
name)
|
||||
@@ -225,11 +241,20 @@ class LoRAPipeline(ComposedPipelineBase):
|
||||
for name, layer in self.lora_layers.items():
|
||||
lora_A_name = name + ".lora_A"
|
||||
lora_B_name = name + ".lora_B"
|
||||
lora_alpha_name = name + ".lora_alpha"
|
||||
if lora_A_name in self.lora_adapters[lora_nickname]\
|
||||
and lora_B_name in self.lora_adapters[lora_nickname]:
|
||||
# Get alpha value for this layer (defaults to None if not present)
|
||||
lora_A = self.lora_adapters[lora_nickname][lora_A_name]
|
||||
lora_B = self.lora_adapters[lora_nickname][lora_B_name]
|
||||
# Simple lookup - alpha stored with same naming scheme as lora_A/lora_B
|
||||
alpha = self.lora_adapters[lora_nickname].get(
|
||||
lora_alpha_name) if adapter_updated else None
|
||||
|
||||
layer.set_lora_weights(
|
||||
self.lora_adapters[lora_nickname][lora_A_name],
|
||||
self.lora_adapters[lora_nickname][lora_B_name],
|
||||
lora_A,
|
||||
lora_B,
|
||||
lora_alpha=alpha,
|
||||
training_mode=self.fastvideo_args.training_mode,
|
||||
lora_path=lora_path)
|
||||
adapted_count += 1
|
||||
|
||||
@@ -85,10 +85,10 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
|
||||
if fastvideo_args.pipeline_config.dit_config.boundary_ratio is not None:
|
||||
boundary_timestep = fastvideo_args.pipeline_config.dit_config.boundary_ratio * self.scheduler.num_train_timesteps
|
||||
high_noise_timesteps = timesteps[timesteps >= boundary_timestep]
|
||||
else:
|
||||
boundary_timestep = None
|
||||
|
||||
high_noise_timesteps = timesteps[timesteps >= boundary_timestep]
|
||||
high_noise_timesteps = None
|
||||
|
||||
# Image kwargs (kept empty unless caller provides compatible args)
|
||||
image_kwargs: dict = {}
|
||||
@@ -142,10 +142,18 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
pos_start_base = 0
|
||||
|
||||
# Determine block sizes
|
||||
block_sizes = [self.num_frames_per_block] * 7
|
||||
block_sizes[0] = 1
|
||||
if t % self.num_frames_per_block != 0:
|
||||
raise ValueError(
|
||||
"num_frames must be divisible by num_frames_per_block for causal DMD denoising"
|
||||
)
|
||||
num_blocks = t // self.num_frames_per_block
|
||||
block_sizes = [self.num_frames_per_block] * num_blocks
|
||||
start_index = 0
|
||||
|
||||
# For now hardcode the first block to be 1 frame assuming the model is Wan2.2-MoE
|
||||
if boundary_timestep is not None:
|
||||
block_sizes[0] = 1
|
||||
|
||||
first_frame_latent = None
|
||||
if batch.pil_image is not None:
|
||||
# Causal video gen directly replaces the first frame of the latent with
|
||||
@@ -392,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
|
||||
|
||||
@@ -482,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
|
||||
@@ -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)
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
|
||||
import pytest
|
||||
|
||||
@@ -52,12 +53,38 @@ LORA_CONFIGS = [
|
||||
"negative_prompt": "bad quality video,色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走",
|
||||
"ssim_threshold": 0.79
|
||||
}
|
||||
# TODO: Add a LoRA with lora_alpha values to test alpha scaling
|
||||
#
|
||||
# Context: This change is mainly for an in-progress ticket porting over LongCat-Video,
|
||||
# where they used an alpha value that is two times smaller than their rank. This fix
|
||||
# ensures that LoRA weights are correctly scaled by the alpha/rank ratio when merged.
|
||||
#
|
||||
# Issue: Currently, we cannot add a test for LoRA adapters with alpha values because:
|
||||
# - The existing public LoRAs for Wan-AI/Wan2.1-T2V-1.3B-Diffusers don't store lora_alpha
|
||||
# - No publicly available LoRA for this model includes lora_alpha tensors in their weights
|
||||
# - This is why the alpha/rank scaling bug wasn't caught by existing tests
|
||||
#
|
||||
# The fix has been validated with:
|
||||
# - LongCat-Video distilled LoRA (which includes alpha values)
|
||||
# - Manual testing shows correct alpha/rank scaling behavior
|
||||
# - Backward compatibility confirmed with LoRAs without alpha values
|
||||
#
|
||||
# Future work:
|
||||
# - Add a synthetic LoRA test fixture with alpha values when feasible
|
||||
# - Or wait for public Wan LoRAs with alpha to become available
|
||||
]
|
||||
|
||||
MODEL_TO_PARAMS = {
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WAN_LORA_PARAMS,
|
||||
}
|
||||
|
||||
def _sanitize_filename_component(name: str) -> str:
|
||||
"""Sanitize filename to remove invalid characters (same logic as VideoGenerator)"""
|
||||
sanitized = re.sub(r'[\\/:*?"<>|]', '', name)
|
||||
sanitized = sanitized.strip().strip('.')
|
||||
sanitized = re.sub(r'\s+', ' ', sanitized)
|
||||
return sanitized or "video"
|
||||
|
||||
@pytest.mark.parametrize("model_id", list(MODEL_TO_PARAMS.keys()))
|
||||
def test_merge_lora_weights(model_id):
|
||||
lora_config = LORA_CONFIGS[0] # test only one
|
||||
@@ -137,14 +164,16 @@ def test_lora_inference_similarity(ATTENTION_BACKEND, model_id):
|
||||
generation_kwargs["negative_prompt"] = lora_config["negative_prompt"]
|
||||
|
||||
generator.set_lora_adapter(lora_nickname=lora_nickname, lora_path=lora_path)
|
||||
# Sanitize the filename before adding .mp4 extension to match VideoGenerator's behavior
|
||||
output_video_name = f"{lora_path.split('/')[-1]}_{prompt[:50]}"
|
||||
generation_kwargs["output_path"] = output_dir
|
||||
generation_kwargs["output_video_name"] = output_video_name
|
||||
output_video_name = _sanitize_filename_component(output_video_name)
|
||||
generated_video_path = os.path.join(output_dir, f"{output_video_name}.mp4")
|
||||
generation_kwargs["output_path"] = generated_video_path
|
||||
|
||||
generator.generate_video(prompt, **generation_kwargs)
|
||||
|
||||
assert os.path.exists(
|
||||
output_dir), f"Output video was not generated at {output_dir}"
|
||||
generated_video_path), f"Output video was not generated at {generated_video_path}"
|
||||
|
||||
reference_folder = os.path.join(script_dir, 'L40S_reference_videos', model_id.split('/')[-1], ATTENTION_BACKEND)
|
||||
|
||||
@@ -153,13 +182,25 @@ def test_lora_inference_similarity(ATTENTION_BACKEND, model_id):
|
||||
raise FileNotFoundError(
|
||||
f"Reference video folder does not exist: {reference_folder}")
|
||||
|
||||
# Find the matching reference video for the switched LoRA
|
||||
# Find the matching reference video - try exact match first, then fuzzy match
|
||||
# The reference might have different sanitization (e.g., trailing spaces)
|
||||
reference_video_name = None
|
||||
|
||||
unsanitized_prefix = f"{lora_path.split('/')[-1]}_{prompt[:50]}"
|
||||
|
||||
for filename in os.listdir(reference_folder):
|
||||
# Check if the filename starts with the expected output_video_name and ends with .mp4
|
||||
if filename.startswith(output_video_name) and filename.endswith('.mp4'):
|
||||
reference_video_name = filename # Remove .mp4 extension to match the logic below
|
||||
if not filename.endswith('.mp4'):
|
||||
continue
|
||||
|
||||
# Try exact match with sanitized name
|
||||
if filename.startswith(output_video_name):
|
||||
reference_video_name = filename
|
||||
break
|
||||
|
||||
# Try match with unsanitized prefix (for legacy reference videos)
|
||||
# Remove .mp4 and compare the base names after sanitization
|
||||
base_filename = filename[:-4] # Remove .mp4
|
||||
if _sanitize_filename_component(base_filename) == output_video_name:
|
||||
reference_video_name = filename
|
||||
break
|
||||
|
||||
if not reference_video_name:
|
||||
@@ -167,7 +208,6 @@ def test_lora_inference_similarity(ATTENTION_BACKEND, model_id):
|
||||
raise FileNotFoundError(f"Reference video missing for adapter {lora_path}")
|
||||
|
||||
reference_video_path = os.path.join(reference_folder, reference_video_name)
|
||||
generated_video_path = os.path.join(output_dir, output_video_name + ".mp4")
|
||||
|
||||
logger.info(
|
||||
f"Computing SSIM between {reference_video_path} and {generated_video_path}"
|
||||
|
||||
@@ -729,9 +729,6 @@ class DistillationPipeline(TrainingPipeline):
|
||||
self.num_train_timestep, [1],
|
||||
device=self.device,
|
||||
dtype=torch.long)
|
||||
world_group = get_world_group()
|
||||
if world_group.world_size > 1:
|
||||
world_group.broadcast(timestep, src=0)
|
||||
|
||||
timestep = shift_timestep(
|
||||
timestep,
|
||||
@@ -844,9 +841,6 @@ class DistillationPipeline(TrainingPipeline):
|
||||
self.num_train_timestep, [1],
|
||||
device=self.device,
|
||||
dtype=torch.long)
|
||||
world_group = get_world_group()
|
||||
if world_group.world_size > 1:
|
||||
world_group.broadcast(fake_score_timestep, src=0)
|
||||
|
||||
fake_score_timestep = shift_timestep(
|
||||
fake_score_timestep,
|
||||
|
||||
@@ -467,6 +467,10 @@ class WorkerMultiprocProc:
|
||||
"output_batch": output_batch.output.cpu(),
|
||||
"logging_info": logging_info
|
||||
})
|
||||
else:
|
||||
result = self.worker.execute_method(
|
||||
method, *args, **kwargs)
|
||||
self.pipe.send(result)
|
||||
else:
|
||||
result = self.worker.execute_method(method, *args, **kwargs)
|
||||
self.pipe.send(result)
|
||||
|
||||
@@ -14,6 +14,7 @@ edit_uri: edit/main/docs/
|
||||
# Configuration
|
||||
theme:
|
||||
name: material
|
||||
favicon: assets/logos/icon_simple.svg
|
||||
palette:
|
||||
- scheme: default
|
||||
toggle:
|
||||
|
||||
Reference in New Issue
Block a user