Support inference with LingBot-World-Base (Cam)

This commit is contained in:
huangkunzhe.hkz
2026-02-05 10:42:33 +08:00
parent a6b026526f
commit 9b86fff645
5 changed files with 1129 additions and 0 deletions
@@ -0,0 +1,56 @@
format: civitai
pipeline: LingBotWorld
transformer_additional_kwargs:
transformer_low_noise_model_subpath: ./low_noise_model
transformer_high_noise_model_subpath: ./high_noise_model
transformer_combination_type: "moe"
# LingBot-World uses higher boundary than Wan2.2 (0.947 vs 0.900)
boundary: 0.947
model_type: i2v
cross_attn_type: cross_attn
dict_mapping:
in_dim: in_channels
dim: hidden_size
vae_kwargs:
vae_type: "AutoencoderKLWan"
vae_subpath: Wan2.1_VAE.pth
temporal_compression_ratio: 4
spatial_compression_ratio: 8
text_encoder_kwargs:
text_encoder_subpath: models_t5_umt5-xxl-enc-bf16.pth
tokenizer_subpath: google/umt5-xxl
text_length: 512
vocab: 256384
dim: 4096
dim_attn: 4096
dim_ffn: 10240
num_heads: 64
num_layers: 24
num_buckets: 32
shared_pos: False
dropout: 0.0
scheduler_kwargs:
scheduler_subpath: null
num_train_timesteps: 1000
# LingBot-World uses higher shift than Wan2.2 (10.0 vs 5.0)
shift: 10.0
use_dynamic_shifting: false
base_shift: 0.5
max_shift: 1.15
base_image_seq_len: 256
max_image_seq_len: 4096
image_encoder_kwargs:
image_encoder_subpath: models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth
# LingBot-World specific inference settings
inference_kwargs:
# Number of inference steps (LingBot-World uses 70, Wan2.2 uses 40)
num_inference_steps: 70
# Guidance scale for both low and high noise models
guidance_scale: 5.0
# Default negative prompt
negative_prompt: "画面突变,色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走,镜头晃动,画面闪烁,模糊,噪点,水印,签名,文字,变形,扭曲,液化,不合逻辑的结构,卡顿,PPT幻灯片感,过暗,欠曝,低对比度,霓虹灯光感,过度锐化,3D渲染感,人物,行人,游客,身体,皮肤,肢体,面部特征,汽车,电线"
+1
View File
@@ -42,6 +42,7 @@ from .wan_text_encoder import WanT5EncoderModel
from .wan_transformer3d import (Wan2_2Transformer3DModel, WanRMSNorm,
WanSelfAttention, WanTransformer3DModel)
from .wan_transformer3d_animate import Wan2_2Transformer3DModel_Animate
from .wan_transformer3d_lingbot import LingBotWanTransformer3DModel
from .wan_transformer3d_s2v import Wan2_2Transformer3DModel_S2V
from .wan_transformer3d_vace import VaceWanTransformer3DModel
from .wan_vae import AutoencoderKLWan, AutoencoderKLWan_
+1
View File
@@ -28,6 +28,7 @@ from .pipeline_wan_phantom import WanFunPhantomPipeline
from .pipeline_wan_vace import WanVacePipeline
from .pipeline_z_image import ZImagePipeline
from .pipeline_z_image_control import ZImageControlPipeline
from .pipeline_lingbot_world import LingBotWorldI2VPipeline
WanFunPipeline = WanPipeline
WanI2VPipeline = WanFunInpaintPipeline
@@ -0,0 +1,840 @@
import inspect
import math
import os
from dataclasses import dataclass
from typing import Any, Callable, Dict, List, Optional, Tuple, Union
import numpy as np
import torch
import torch.nn.functional as TorchF
import torchvision.transforms.functional as TF
from diffusers import FlowMatchEulerDiscreteScheduler
from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback
from diffusers.image_processor import VaeImageProcessor
from diffusers.pipelines.pipeline_utils import DiffusionPipeline
from diffusers.utils import BaseOutput, logging, replace_example_docstring
from diffusers.utils.torch_utils import randn_tensor
from diffusers.video_processor import VideoProcessor
from einops import rearrange
from PIL import Image
from ..models import AutoencoderKLWan, AutoTokenizer, WanT5EncoderModel
from ..models.wan_transformer3d_lingbot import LingBotWanTransformer3DModel
from ..utils.cam_utils import (
compute_relative_poses,
get_Ks_transformed,
get_plucker_embeddings,
interpolate_camera_poses,
)
from ..utils.fm_solvers import FlowDPMSolverMultistepScheduler, get_sampling_sigmas
from ..utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
logger = logging.get_logger(__name__)
EXAMPLE_DOC_STRING = """
Examples:
```python
# LingBot-World I2V with camera control
from videox_fun.pipeline import LingBotWorldI2VPipeline
pipeline = LingBotWorldI2VPipeline(...)
video = pipeline(
prompt="A beautiful landscape",
video=input_video,
mask_video=mask,
action_path="path/to/camera/data", # Optional camera control
)
```
"""
def retrieve_timesteps(
scheduler,
num_inference_steps: Optional[int] = None,
device: Optional[Union[str, torch.device]] = None,
timesteps: Optional[List[int]] = None,
sigmas: Optional[List[float]] = None,
**kwargs,
):
"""Retrieve timesteps from scheduler."""
if timesteps is not None and sigmas is not None:
raise ValueError("Only one of `timesteps` or `sigmas` can be passed.")
if timesteps is not None:
accepts_timesteps = "timesteps" in set(inspect.signature(scheduler.set_timesteps).parameters.keys())
if not accepts_timesteps:
raise ValueError(
f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"
f" timestep schedules."
)
scheduler.set_timesteps(timesteps=timesteps, device=device, **kwargs)
timesteps = scheduler.timesteps
num_inference_steps = len(timesteps)
elif sigmas is not None:
accept_sigmas = "sigmas" in set(inspect.signature(scheduler.set_timesteps).parameters.keys())
if not accept_sigmas:
raise ValueError(
f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"
f" sigmas schedules."
)
scheduler.set_timesteps(sigmas=sigmas, device=device, **kwargs)
timesteps = scheduler.timesteps
num_inference_steps = len(timesteps)
else:
scheduler.set_timesteps(num_inference_steps, device=device, **kwargs)
timesteps = scheduler.timesteps
return timesteps, num_inference_steps
def resize_mask(mask, latent, process_first_frame_only=True):
"""Resize mask to match latent dimensions."""
latent_size = latent.size()
batch_size, channels, num_frames, height, width = mask.shape
if process_first_frame_only:
target_size = list(latent_size[2:])
target_size[0] = 1
first_frame_resized = TorchF.interpolate(
mask[:, :, 0:1, :, :],
size=target_size,
mode='trilinear',
align_corners=False
)
target_size = list(latent_size[2:])
target_size[0] = target_size[0] - 1
if target_size[0] != 0:
remaining_frames_resized = TorchF.interpolate(
mask[:, :, 1:, :, :],
size=target_size,
mode='trilinear',
align_corners=False
)
resized_mask = torch.cat([first_frame_resized, remaining_frames_resized], dim=2)
else:
resized_mask = first_frame_resized
else:
target_size = list(latent_size[2:])
resized_mask = TorchF.interpolate(
mask,
size=target_size,
mode='trilinear',
align_corners=False
)
return resized_mask
@dataclass
class LingBotWorldPipelineOutput(BaseOutput):
"""
Output class for LingBot-World pipelines.
Args:
videos: Generated video tensor with shape (batch_size, num_frames, channels, height, width)
"""
videos: torch.Tensor
class LingBotWorldI2VPipeline(DiffusionPipeline):
"""
Pipeline for image-to-video generation using LingBot-World with optional camera control.
"""
_optional_components = ["transformer_2"]
model_cpu_offload_seq = "text_encoder->transformer_2->transformer->vae"
_callback_tensor_inputs = [
"latents",
"prompt_embeds",
"negative_prompt_embeds",
]
def __init__(
self,
tokenizer: AutoTokenizer,
text_encoder: WanT5EncoderModel,
vae: AutoencoderKLWan,
transformer: LingBotWanTransformer3DModel,
transformer_2: LingBotWanTransformer3DModel = None,
scheduler: FlowMatchEulerDiscreteScheduler = None,
):
super().__init__()
self.register_modules(
tokenizer=tokenizer,
text_encoder=text_encoder,
vae=vae,
transformer=transformer,
transformer_2=transformer_2,
scheduler=scheduler
)
self.video_processor = VideoProcessor(vae_scale_factor=self.vae.spatial_compression_ratio)
self.image_processor = VaeImageProcessor(vae_scale_factor=self.vae.spatial_compression_ratio)
self.mask_processor = VaeImageProcessor(
vae_scale_factor=self.vae.spatial_compression_ratio,
do_normalize=False,
do_binarize=True,
do_convert_grayscale=True
)
def _get_t5_prompt_embeds(
self,
prompt: Union[str, List[str]] = None,
num_videos_per_prompt: int = 1,
max_sequence_length: int = 512,
device: Optional[torch.device] = None,
dtype: Optional[torch.dtype] = None,
):
device = device or self._execution_device
dtype = dtype or self.text_encoder.dtype
prompt = [prompt] if isinstance(prompt, str) else prompt
batch_size = len(prompt)
text_inputs = self.tokenizer(
prompt,
padding="max_length",
max_length=max_sequence_length,
truncation=True,
add_special_tokens=True,
return_tensors="pt",
)
text_input_ids = text_inputs.input_ids
prompt_attention_mask = text_inputs.attention_mask
untruncated_ids = self.tokenizer(prompt, padding="longest", return_tensors="pt").input_ids
if untruncated_ids.shape[-1] >= text_input_ids.shape[-1] and not torch.equal(text_input_ids, untruncated_ids):
removed_text = self.tokenizer.batch_decode(untruncated_ids[:, max_sequence_length - 1 : -1])
logger.warning(
"The following part of your input was truncated because `max_sequence_length` is set to "
f" {max_sequence_length} tokens: {removed_text}"
)
seq_lens = prompt_attention_mask.gt(0).sum(dim=1).long()
prompt_embeds = self.text_encoder(text_input_ids.to(device), attention_mask=prompt_attention_mask.to(device))[0]
prompt_embeds = prompt_embeds.to(dtype=dtype, device=device)
_, seq_len, _ = prompt_embeds.shape
prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt, 1)
prompt_embeds = prompt_embeds.view(batch_size * num_videos_per_prompt, seq_len, -1)
return [u[:v] for u, v in zip(prompt_embeds, seq_lens)]
def encode_prompt(
self,
prompt: Union[str, List[str]],
negative_prompt: Optional[Union[str, List[str]]] = None,
do_classifier_free_guidance: bool = True,
num_videos_per_prompt: int = 1,
prompt_embeds: Optional[torch.Tensor] = None,
negative_prompt_embeds: Optional[torch.Tensor] = None,
max_sequence_length: int = 512,
device: Optional[torch.device] = None,
dtype: Optional[torch.dtype] = None,
):
"""Encodes the prompt into text encoder hidden states."""
device = device or self._execution_device
prompt = [prompt] if isinstance(prompt, str) else prompt
if prompt is not None:
batch_size = len(prompt)
else:
batch_size = prompt_embeds.shape[0]
if prompt_embeds is None:
prompt_embeds = self._get_t5_prompt_embeds(
prompt=prompt,
num_videos_per_prompt=num_videos_per_prompt,
max_sequence_length=max_sequence_length,
device=device,
dtype=dtype,
)
if do_classifier_free_guidance and negative_prompt_embeds is None:
negative_prompt = negative_prompt or ""
negative_prompt = batch_size * [negative_prompt] if isinstance(negative_prompt, str) else negative_prompt
if prompt is not None and type(prompt) is not type(negative_prompt):
raise TypeError(
f"`negative_prompt` should be the same type to `prompt`, but got {type(negative_prompt)} !="
f" {type(prompt)}."
)
elif batch_size != len(negative_prompt):
raise ValueError(
f"`negative_prompt`: {negative_prompt} has batch size {len(negative_prompt)}, but `prompt`:"
f" {prompt} has batch size {batch_size}."
)
negative_prompt_embeds = self._get_t5_prompt_embeds(
prompt=negative_prompt,
num_videos_per_prompt=num_videos_per_prompt,
max_sequence_length=max_sequence_length,
device=device,
dtype=dtype,
)
return prompt_embeds, negative_prompt_embeds
def prepare_latents(
self, batch_size, num_channels_latents, num_frames, height, width, dtype, device, generator, latents=None
):
"""Prepare latent tensors."""
if isinstance(generator, list) and len(generator) != batch_size:
raise ValueError(
f"You have passed a list of generators of length {len(generator)}, but requested an effective batch"
f" size of {batch_size}."
)
shape = (
batch_size,
num_channels_latents,
(num_frames - 1) // self.vae.temporal_compression_ratio + 1,
height // self.vae.spatial_compression_ratio,
width // self.vae.spatial_compression_ratio,
)
if latents is None:
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
else:
latents = latents.to(device)
if hasattr(self.scheduler, "init_noise_sigma"):
latents = latents * self.scheduler.init_noise_sigma
return latents
def prepare_mask_latents(
self, mask, masked_image, batch_size, height, width, dtype, device, generator, do_classifier_free_guidance, noise_aug_strength
):
"""Prepare mask latent variables."""
if mask is not None:
mask = mask.to(device=device, dtype=self.vae.dtype)
bs = 1
new_mask = []
for i in range(0, mask.shape[0], bs):
mask_bs = mask[i : i + bs]
mask_bs = self.vae.encode(mask_bs)[0]
mask_bs = mask_bs.mode()
new_mask.append(mask_bs)
mask = torch.cat(new_mask, dim=0)
if masked_image is not None:
masked_image = masked_image.to(device=device, dtype=self.vae.dtype)
bs = 1
new_mask_pixel_values = []
for i in range(0, masked_image.shape[0], bs):
mask_pixel_values_bs = masked_image[i : i + bs]
mask_pixel_values_bs = self.vae.encode(mask_pixel_values_bs)[0]
mask_pixel_values_bs = mask_pixel_values_bs.mode()
new_mask_pixel_values.append(mask_pixel_values_bs)
masked_image_latents = torch.cat(new_mask_pixel_values, dim=0)
else:
masked_image_latents = None
return mask, masked_image_latents
def prepare_camera_conditioning(
self,
action_path: str,
frame_num: int,
height: int,
width: int,
lat_f: int,
lat_h: int,
lat_w: int,
device: torch.device,
dtype: torch.dtype,
) -> Optional[Dict[str, torch.Tensor]]:
"""
Prepare camera conditioning from action_path.
Args:
action_path: Path to directory containing poses.npy and intrinsics.npy
frame_num: Number of video frames
height: Video height
width: Video width
lat_f: Latent temporal dimension
lat_h: Latent height
lat_w: Latent width
device: Target device
dtype: Target dtype
Returns:
Dictionary containing c2ws_plucker_emb or None if action_path is None
"""
if action_path is None:
return None
# Load camera data
c2ws = np.load(os.path.join(action_path, "poses.npy")) # opencv coordinate
len_c2ws = ((len(c2ws) - 1) // 4) * 4 + 1
frame_num = min(frame_num, len_c2ws)
c2ws = c2ws[:frame_num]
Ks = torch.from_numpy(np.load(os.path.join(action_path, "intrinsics.npy"))).float()
# Transform intrinsics for the target resolution
# The provided intrinsics are for original image size (480p)
Ks = get_Ks_transformed(
Ks,
height_org=480,
width_org=832,
height_resize=height,
width_resize=width,
height_final=height,
width_final=width
)
Ks = Ks[0]
# Interpolate camera poses to latent temporal resolution
len_c2ws = len(c2ws)
c2ws_infer = interpolate_camera_poses(
src_indices=np.linspace(0, len_c2ws - 1, len_c2ws),
src_rot_mat=c2ws[:, :3, :3],
src_trans_vec=c2ws[:, :3, 3],
tgt_indices=np.linspace(0, len_c2ws - 1, int((len_c2ws - 1) // 4) + 1),
)
c2ws_infer = compute_relative_poses(c2ws_infer, framewise=True)
Ks = Ks.repeat(len(c2ws_infer), 1)
c2ws_infer = c2ws_infer.to(device)
Ks = Ks.to(device)
# Generate Plucker embeddings
c2ws_plucker_emb = get_plucker_embeddings(c2ws_infer, Ks, height, width)
c2ws_plucker_emb = rearrange(
c2ws_plucker_emb,
'f (h c1) (w c2) c -> (f h w) (c c1 c2)',
c1=int(height // lat_h),
c2=int(width // lat_w),
)
c2ws_plucker_emb = c2ws_plucker_emb[None, ...] # [1, f*h*w, c]
c2ws_plucker_emb = rearrange(
c2ws_plucker_emb,
'b (f h w) c -> b c f h w',
f=lat_f, h=lat_h, w=lat_w
).to(dtype)
dit_cond_dict = {
"c2ws_plucker_emb": c2ws_plucker_emb.chunk(1, dim=0),
}
return dit_cond_dict
def decode_latents(self, latents: torch.Tensor) -> torch.Tensor:
"""Decode latents to video frames."""
frames = self.vae.decode(latents.to(self.vae.dtype)).sample
frames = (frames / 2 + 0.5).clamp(0, 1)
frames = frames.cpu().float().numpy()
return frames
def prepare_extra_step_kwargs(self, generator, eta):
"""Prepare extra kwargs for the scheduler step."""
accepts_eta = "eta" in set(inspect.signature(self.scheduler.step).parameters.keys())
extra_step_kwargs = {}
if accepts_eta:
extra_step_kwargs["eta"] = eta
accepts_generator = "generator" in set(inspect.signature(self.scheduler.step).parameters.keys())
if accepts_generator:
extra_step_kwargs["generator"] = generator
return extra_step_kwargs
def check_inputs(
self,
prompt,
height,
width,
negative_prompt,
callback_on_step_end_tensor_inputs,
prompt_embeds=None,
negative_prompt_embeds=None,
):
"""Check inputs for validity."""
if height % 8 != 0 or width % 8 != 0:
raise ValueError(f"`height` and `width` have to be divisible by 8 but are {height} and {width}.")
if callback_on_step_end_tensor_inputs is not None and not all(
k in self._callback_tensor_inputs for k in callback_on_step_end_tensor_inputs
):
raise ValueError(
f"`callback_on_step_end_tensor_inputs` has to be in {self._callback_tensor_inputs}, "
f"but found {[k for k in callback_on_step_end_tensor_inputs if k not in self._callback_tensor_inputs]}"
)
if prompt is not None and prompt_embeds is not None:
raise ValueError(
f"Cannot forward both `prompt`: {prompt} and `prompt_embeds`: {prompt_embeds}."
)
elif prompt is None and prompt_embeds is None:
raise ValueError(
"Provide either `prompt` or `prompt_embeds`. Cannot leave both undefined."
)
elif prompt is not None and (not isinstance(prompt, str) and not isinstance(prompt, list)):
raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}")
if prompt is not None and negative_prompt_embeds is not None:
raise ValueError(
f"Cannot forward both `prompt`: {prompt} and `negative_prompt_embeds`: {negative_prompt_embeds}."
)
if negative_prompt is not None and negative_prompt_embeds is not None:
raise ValueError(
f"Cannot forward both `negative_prompt`: {negative_prompt} and `negative_prompt_embeds`:"
f" {negative_prompt_embeds}."
)
if prompt_embeds is not None and negative_prompt_embeds is not None:
if prompt_embeds.shape != negative_prompt_embeds.shape:
raise ValueError(
"`prompt_embeds` and `negative_prompt_embeds` must have the same shape."
)
@property
def guidance_scale(self):
return self._guidance_scale
@property
def num_timesteps(self):
return self._num_timesteps
@property
def attention_kwargs(self):
return self._attention_kwargs
@property
def interrupt(self):
return self._interrupt
@torch.no_grad()
@replace_example_docstring(EXAMPLE_DOC_STRING)
def __call__(
self,
prompt: Optional[Union[str, List[str]]] = None,
negative_prompt: Optional[Union[str, List[str]]] = None,
height: int = 480,
width: int = 832,
video: Union[torch.FloatTensor] = None,
mask_video: Union[torch.FloatTensor] = None,
num_frames: int = 81,
num_inference_steps: int = 70,
timesteps: Optional[List[int]] = None,
guidance_scale: float = 5.0,
num_videos_per_prompt: int = 1,
eta: float = 0.0,
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
latents: Optional[torch.FloatTensor] = None,
prompt_embeds: Optional[torch.FloatTensor] = None,
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
output_type: str = "numpy",
return_dict: bool = False,
callback_on_step_end: Optional[
Union[Callable[[int, int, Dict], None], PipelineCallback, MultiPipelineCallbacks]
] = None,
attention_kwargs: Optional[Dict[str, Any]] = None,
callback_on_step_end_tensor_inputs: List[str] = ["latents"],
max_sequence_length: int = 512,
boundary: float = 0.947,
comfyui_progressbar: bool = False,
shift: int = 10,
action_path: Optional[str] = None,
) -> Union[LingBotWorldPipelineOutput, Tuple]:
"""
Generate video from image with optional camera control.
Args:
prompt: Text prompt for generation
negative_prompt: Negative prompt for exclusion
height: Video height
width: Video width
video: Input video tensor
mask_video: Mask video tensor
num_frames: Number of frames to generate
num_inference_steps: Number of denoising steps
timesteps: Custom timesteps
guidance_scale: Classifier-free guidance scale
num_videos_per_prompt: Number of videos per prompt
eta: Eta for DDIM scheduler
generator: Random generator
latents: Pre-generated latents
prompt_embeds: Pre-computed prompt embeddings
negative_prompt_embeds: Pre-computed negative prompt embeddings
output_type: Output format ("numpy", "latent", etc.)
return_dict: Whether to return a dict
callback_on_step_end: Callback function
attention_kwargs: Additional attention kwargs
callback_on_step_end_tensor_inputs: Tensor inputs for callback
max_sequence_length: Max text sequence length
boundary: Timestep boundary for model switching
comfyui_progressbar: Enable ComfyUI progress bar
shift: Noise schedule shift parameter
action_path: Path to camera data (poses.npy, intrinsics.npy)
Examples:
Returns:
Generated video tensor
"""
if isinstance(callback_on_step_end, (PipelineCallback, MultiPipelineCallbacks)):
callback_on_step_end_tensor_inputs = callback_on_step_end.tensor_inputs
num_videos_per_prompt = 1
# 1. Check inputs
self.check_inputs(
prompt, height, width, negative_prompt,
callback_on_step_end_tensor_inputs, prompt_embeds, negative_prompt_embeds,
)
self._guidance_scale = guidance_scale
self._attention_kwargs = attention_kwargs
self._interrupt = False
# 2. Default call parameters
if prompt is not None and isinstance(prompt, str):
batch_size = 1
elif prompt is not None and isinstance(prompt, list):
batch_size = len(prompt)
else:
batch_size = prompt_embeds.shape[0]
device = self._execution_device
weight_dtype = self.text_encoder.dtype
do_classifier_free_guidance = guidance_scale > 1.0
# 3. Encode input prompt
prompt_embeds, negative_prompt_embeds = self.encode_prompt(
prompt, negative_prompt, do_classifier_free_guidance,
num_videos_per_prompt=num_videos_per_prompt,
prompt_embeds=prompt_embeds,
negative_prompt_embeds=negative_prompt_embeds,
max_sequence_length=max_sequence_length,
device=device,
)
if do_classifier_free_guidance:
in_prompt_embeds = negative_prompt_embeds + prompt_embeds
else:
in_prompt_embeds = prompt_embeds
# 4. Prepare timesteps
if isinstance(self.scheduler, FlowMatchEulerDiscreteScheduler):
timesteps, num_inference_steps = retrieve_timesteps(
self.scheduler, num_inference_steps, device, timesteps, mu=1
)
elif isinstance(self.scheduler, FlowUniPCMultistepScheduler):
self.scheduler.set_timesteps(num_inference_steps, device=device, shift=shift)
timesteps = self.scheduler.timesteps
elif isinstance(self.scheduler, FlowDPMSolverMultistepScheduler):
sampling_sigmas = get_sampling_sigmas(num_inference_steps, shift)
timesteps, _ = retrieve_timesteps(self.scheduler, device=device, sigmas=sampling_sigmas)
else:
timesteps, num_inference_steps = retrieve_timesteps(
self.scheduler, num_inference_steps, device, timesteps
)
self._num_timesteps = len(timesteps)
if comfyui_progressbar:
from comfy.utils import ProgressBar
pbar = ProgressBar(num_inference_steps + 2)
# 5. Prepare latents
if video is not None:
video_length = video.shape[2]
init_video = self.image_processor.preprocess(
rearrange(video, "b c f h w -> (b f) c h w"), height=height, width=width
)
init_video = init_video.to(dtype=torch.float32)
init_video = rearrange(init_video, "(b f) c h w -> b c f h w", f=video_length)
else:
init_video = None
latent_channels = self.vae.config.latent_channels
latents = self.prepare_latents(
batch_size * num_videos_per_prompt,
latent_channels,
num_frames,
height,
width,
weight_dtype,
device,
generator,
latents,
)
if comfyui_progressbar:
pbar.update(1)
# 6. Prepare mask latent variables
if init_video is not None:
if (mask_video == 255).all():
mask_latents = torch.tile(
torch.zeros_like(latents)[:, :1].to(device, weight_dtype), [1, 4, 1, 1, 1]
)
masked_video_latents = torch.zeros_like(latents).to(device, weight_dtype)
if self.vae.spatial_compression_ratio >= 16:
mask = torch.ones_like(latents).to(device, weight_dtype)[:, :1]
else:
bs, _, video_length, height, width = video.size()
mask_condition = self.mask_processor.preprocess(
rearrange(mask_video, "b c f h w -> (b f) c h w"), height=height, width=width
)
mask_condition = mask_condition.to(dtype=torch.float32)
mask_condition = rearrange(mask_condition, "(b f) c h w -> b c f h w", f=video_length)
masked_video = init_video * (torch.tile(mask_condition, [1, 3, 1, 1, 1]) < 0.5)
_, masked_video_latents = self.prepare_mask_latents(
None, masked_video, batch_size, height, width,
weight_dtype, device, generator, do_classifier_free_guidance, noise_aug_strength=None,
)
mask_condition = torch.concat([
torch.repeat_interleave(mask_condition[:, :, 0:1], repeats=4, dim=2),
mask_condition[:, :, 1:]
], dim=2)
mask_condition = mask_condition.view(bs, mask_condition.shape[2] // 4, 4, height, width)
mask_condition = mask_condition.transpose(1, 2)
mask_latents = resize_mask(1 - mask_condition, masked_video_latents, True).to(device, weight_dtype)
if self.vae.spatial_compression_ratio >= 16:
mask = TorchF.interpolate(
mask_condition[:, :1], size=latents.size()[-3:], mode='trilinear', align_corners=True
).to(device, weight_dtype)
if not mask[:, :, 0, :, :].any():
mask[:, :, 1:, :, :] = 1
latents = (1 - mask) * masked_video_latents + mask * latents
if comfyui_progressbar:
pbar.update(1)
# 7. Prepare camera conditioning
lat_f = (num_frames - 1) // self.vae.temporal_compression_ratio + 1
lat_h = height // self.vae.spatial_compression_ratio
lat_w = width // self.vae.spatial_compression_ratio
dit_cond_dict = self.prepare_camera_conditioning(
action_path=action_path,
frame_num=num_frames,
height=height,
width=width,
lat_f=lat_f,
lat_h=lat_h,
lat_w=lat_w,
device=device,
dtype=weight_dtype,
)
# 8. Prepare extra step kwargs
extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta)
target_shape = (
self.vae.latent_channels,
lat_f,
width // self.vae.spatial_compression_ratio,
height // self.vae.spatial_compression_ratio
)
seq_len = math.ceil(
(target_shape[2] * target_shape[3]) /
(self.transformer.config.patch_size[1] * self.transformer.config.patch_size[2]) * target_shape[1]
)
# 9. Denoising loop
num_warmup_steps = max(len(timesteps) - num_inference_steps * self.scheduler.order, 0)
self.transformer.num_inference_steps = num_inference_steps
with self.progress_bar(total=num_inference_steps) as progress_bar:
for i, t in enumerate(timesteps):
self.transformer.current_steps = i
if self.interrupt:
continue
latent_model_input = torch.cat([latents] * 2) if do_classifier_free_guidance else latents
if hasattr(self.scheduler, "scale_model_input"):
latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)
if init_video is not None:
mask_input = torch.cat([mask_latents] * 2) if do_classifier_free_guidance else mask_latents
masked_video_latents_input = (
torch.cat([masked_video_latents] * 2) if do_classifier_free_guidance else masked_video_latents
)
y = torch.cat([mask_input, masked_video_latents_input], dim=1).to(device, weight_dtype)
# Timestep handling
if self.vae.spatial_compression_ratio >= 16 and init_video is not None:
temp_ts = ((mask[0][0][:, ::2, ::2]) * t).flatten()
temp_ts = torch.cat([
temp_ts,
temp_ts.new_ones(seq_len - temp_ts.size(0)) * t
])
temp_ts = temp_ts.unsqueeze(0)
timestep = temp_ts.expand(latent_model_input.shape[0], temp_ts.size(1))
else:
timestep = t.expand(latent_model_input.shape[0])
# Select transformer based on timestep
if self.transformer_2 is not None:
if t >= boundary * self.scheduler.config.num_train_timesteps:
local_transformer = self.transformer_2
else:
local_transformer = self.transformer
else:
local_transformer = self.transformer
# Predict noise
with torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=device):
noise_pred = local_transformer(
x=latent_model_input,
context=in_prompt_embeds,
t=timestep,
seq_len=seq_len,
y=y,
dit_cond_dict=dit_cond_dict,
)
# Perform guidance
if do_classifier_free_guidance:
if self.transformer_2 is not None and isinstance(self.guidance_scale, (list, tuple)):
sample_guide_scale = (
self.guidance_scale[1]
if t >= self.transformer_2.config.boundary * self.scheduler.config.num_train_timesteps
else self.guidance_scale[0]
)
else:
sample_guide_scale = self.guidance_scale
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
noise_pred = noise_pred_uncond + sample_guide_scale * (noise_pred_text - noise_pred_uncond)
# Compute previous noisy sample
latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs, return_dict=False)[0]
if self.vae.spatial_compression_ratio >= 16 and not mask[:, :, 0, :, :].any():
latents = (1 - mask) * masked_video_latents + mask * latents
if callback_on_step_end is not None:
callback_kwargs = {}
for k in callback_on_step_end_tensor_inputs:
callback_kwargs[k] = locals()[k]
callback_outputs = callback_on_step_end(self, i, t, callback_kwargs)
latents = callback_outputs.pop("latents", latents)
prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
negative_prompt_embeds = callback_outputs.pop("negative_prompt_embeds", negative_prompt_embeds)
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
progress_bar.update()
if comfyui_progressbar:
pbar.update(1)
# 10. Decode latents
if output_type == "numpy":
video = self.decode_latents(latents)
elif not output_type == "latent":
video = self.decode_latents(latents)
video = self.video_processor.postprocess_video(video=video, output_type=output_type)
else:
video = latents
# Offload models
self.maybe_free_model_hooks()
if not return_dict:
video = torch.from_numpy(video)
return LingBotWorldPipelineOutput(videos=video)
+231
View File
@@ -0,0 +1,231 @@
"""Borrowed from LingBot-World/wan/utils/cam_utils.py.
"""
import numpy as np
import torch
from scipy.interpolate import interp1d
from scipy.spatial.transform import Rotation, Slerp
def interpolate_camera_poses(
src_indices: np.ndarray,
src_rot_mat: np.ndarray,
src_trans_vec: np.ndarray,
tgt_indices: np.ndarray,
) -> torch.Tensor:
"""
Interpolate camera poses using linear interpolation for translation
and Slerp for rotation.
Args:
src_indices: Source frame indices
src_rot_mat: Source rotation matrices, shape [N, 3, 3]
src_trans_vec: Source translation vectors, shape [N, 3]
tgt_indices: Target frame indices for interpolation
Returns:
Interpolated camera poses as 4x4 matrices, shape [M, 4, 4]
"""
# interpolate translation
interp_func_trans = interp1d(
src_indices,
src_trans_vec,
axis=0,
kind='linear',
bounds_error=False,
fill_value="extrapolate",
)
interpolated_trans_vec = interp_func_trans(tgt_indices)
# interpolate rotation
src_quat_vec = Rotation.from_matrix(src_rot_mat)
# ensure there is no sudden change in qw
quats = src_quat_vec.as_quat().copy() # [N, 4]
for i in range(1, len(quats)):
if np.dot(quats[i], quats[i-1]) < 0:
quats[i] = -quats[i]
src_quat_vec = Rotation.from_quat(quats)
slerp_func_rot = Slerp(src_indices, src_quat_vec)
interpolated_rot_quat = slerp_func_rot(tgt_indices)
interpolated_rot_mat = interpolated_rot_quat.as_matrix()
poses = np.zeros((len(tgt_indices), 4, 4))
poses[:, :3, :3] = interpolated_rot_mat
poses[:, :3, 3] = interpolated_trans_vec
poses[:, 3, 3] = 1.0
return torch.from_numpy(poses).float()
def SE3_inverse(T: torch.Tensor) -> torch.Tensor:
"""
Compute the inverse of SE3 transformation matrices.
Args:
T: SE3 transformation matrices, shape [B, 4, 4]
Returns:
Inverse transformation matrices, shape [B, 4, 4]
"""
Rot = T[:, :3, :3] # [B, 3, 3]
trans = T[:, :3, 3:] # [B, 3, 1]
R_inv = Rot.transpose(-1, -2)
t_inv = -torch.bmm(R_inv, trans)
T_inv = torch.eye(4, device=T.device, dtype=T.dtype)[None, :, :].repeat(T.shape[0], 1, 1)
T_inv[:, :3, :3] = R_inv
T_inv[:, :3, 3:] = t_inv
return T_inv
def compute_relative_poses(
c2ws_mat: torch.Tensor,
framewise: bool = False,
normalize_trans: bool = True,
) -> torch.Tensor:
"""
Compute relative poses from camera-to-world matrices.
Args:
c2ws_mat: Camera-to-world matrices, shape [F, 4, 4]
framewise: If True, compute frame-to-frame relative poses
normalize_trans: If True, normalize translations to unit max norm
Returns:
Relative pose matrices, shape [F, 4, 4]
"""
ref_w2cs = SE3_inverse(c2ws_mat[0:1])
relative_poses = torch.matmul(ref_w2cs, c2ws_mat)
# ensure identity matrix for 1st frame
relative_poses[0] = torch.eye(4, device=c2ws_mat.device, dtype=c2ws_mat.dtype)
if framewise:
# compute pose between i and i+1
relative_poses_framewise = torch.bmm(SE3_inverse(relative_poses[:-1]), relative_poses[1:])
relative_poses[1:] = relative_poses_framewise
if normalize_trans:
# note refer to camctrl2: "we scale the coordinate inputs to roughly 1 standard deviation to simplify model learning."
translations = relative_poses[:, :3, 3] # [f, 3]
max_norm = torch.norm(translations, dim=-1).max()
# only normalize when moving
if max_norm > 0:
relative_poses[:, :3, 3] = translations / max_norm
return relative_poses
@torch.no_grad()
def create_meshgrid(
n_frames: int,
height: int,
width: int,
bias: float = 0.5,
device='cuda',
dtype=torch.float32
) -> torch.Tensor:
"""
Create a meshgrid for ray generation.
Args:
n_frames: Number of frames
height: Image height
width: Image width
bias: Pixel center bias (0.5 for pixel center)
device: Torch device
dtype: Torch dtype
Returns:
Grid coordinates, shape [F, H*W, 2]
"""
x_range = torch.arange(width, device=device, dtype=dtype)
y_range = torch.arange(height, device=device, dtype=dtype)
grid_y, grid_x = torch.meshgrid(y_range, x_range, indexing='ij')
grid_xy = torch.stack([grid_x, grid_y], dim=-1).view([-1, 2]) + bias # [h*w, 2]
grid_xy = grid_xy[None, ...].repeat(n_frames, 1, 1) # [f, h*w, 2]
return grid_xy
def get_plucker_embeddings(
c2ws_mat: torch.Tensor,
Ks: torch.Tensor,
height: int,
width: int,
) -> torch.Tensor:
"""
Generate Plucker ray embeddings from camera parameters.
Args:
c2ws_mat: Camera-to-world matrices, shape [F, 4, 4]
Ks: Camera intrinsics [fx, fy, cx, cy], shape [F, 4]
height: Image height
width: Image width
Returns:
Plucker embeddings (ray_origin + ray_direction), shape [F, H, W, 6]
"""
n_frames = c2ws_mat.shape[0]
grid_xy = create_meshgrid(n_frames, height, width, device=c2ws_mat.device, dtype=c2ws_mat.dtype) # [f, h*w, 2]
fx, fy, cx, cy = Ks.chunk(4, dim=-1) # [f, 1]
i = grid_xy[..., 0] # [f, h*w]
j = grid_xy[..., 1] # [f, h*w]
zs = torch.ones_like(i) # [f, h*w]
xs = (i - cx) / fx * zs
ys = (j - cy) / fy * zs
directions = torch.stack([xs, ys, zs], dim=-1) # [f, h*w, 3]
directions = directions / directions.norm(dim=-1, keepdim=True) # [f, h*w, 3]
rays_d = directions @ c2ws_mat[:, :3, :3].transpose(-1, -2) # [f, h*w, 3]
rays_o = c2ws_mat[:, :3, 3] # [f, 3]
rays_o = rays_o[:, None, :].expand_as(rays_d) # [f, h*w, 3]
# Plucker coordinates: origin + direction (refer to apt2)
plucker_embeddings = torch.cat([rays_o, rays_d], dim=-1) # [f, h*w, 6]
plucker_embeddings = plucker_embeddings.view([n_frames, height, width, 6]) # [f, h, w, 6]
return plucker_embeddings
def get_Ks_transformed(
Ks: torch.Tensor,
height_org: int,
width_org: int,
height_resize: int,
width_resize: int,
height_final: int,
width_final: int,
) -> torch.Tensor:
"""
Transform camera intrinsics for image resizing and cropping.
Args:
Ks: Original camera intrinsics [fx, fy, cx, cy], shape [F, 4]
height_org: Original image height
width_org: Original image width
height_resize: Resized image height
width_resize: Resized image width
height_final: Final cropped image height
width_final: Final cropped image width
Returns:
Transformed camera intrinsics, shape [F, 4]
"""
fx, fy, cx, cy = Ks.chunk(4, dim=-1) # [f, 1]
scale_x = width_resize / width_org
scale_y = height_resize / height_org
fx_resize = fx * scale_x
fy_resize = fy * scale_y
cx_resize = cx * scale_x
cy_resize = cy * scale_y
crop_offset_x = (width_resize - width_final) / 2
crop_offset_y = (height_resize - height_final) / 2
cx_final = cx_resize - crop_offset_x
cy_final = cy_resize - crop_offset_y
Ks_transformed = torch.zeros_like(Ks)
Ks_transformed[:, 0:1] = fx_resize
Ks_transformed[:, 1:2] = fy_resize
Ks_transformed[:, 2:3] = cx_final
Ks_transformed[:, 3:4] = cy_final
return Ks_transformed