Support inference with LingBot-World-Base (Cam)
This commit is contained in:
@@ -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渲染感,人物,行人,游客,身体,皮肤,肢体,面部特征,汽车,电线"
|
||||
@@ -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_
|
||||
|
||||
@@ -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)
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user