545 lines
21 KiB
Python
545 lines
21 KiB
Python
# -*- coding: utf-8 -*-
|
|
|
|
# Copyright 2024 The CogVideoX team, Tsinghua University & ZhipuAI and The HuggingFace Team.
|
|
# All rights reserved.
|
|
#
|
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
# you may not use this file except in compliance with the License.
|
|
# You may obtain a copy of the License at
|
|
#
|
|
# http://www.apache.org/licenses/LICENSE-2.0
|
|
#
|
|
# Unless required by applicable law or agreed to in writing, software
|
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
# See the License for the specific language governing permissions and
|
|
# limitations under the License.
|
|
|
|
import math
|
|
from typing import Optional, Tuple, Union, List
|
|
|
|
import numpy as np
|
|
import torch
|
|
import torch.nn as nn
|
|
import torch.nn.functional as F
|
|
|
|
|
|
ACTIVATION_FUNCTIONS = {
|
|
"swish": nn.SiLU(),
|
|
"silu": nn.SiLU(),
|
|
"mish": nn.Mish(),
|
|
"gelu": nn.GELU(),
|
|
"relu": nn.ReLU(),
|
|
}
|
|
|
|
|
|
def get_activation(act_fn: str) -> nn.Module:
|
|
"""Helper function to get activation function from string.
|
|
|
|
Args:
|
|
act_fn (str): Name of activation function.
|
|
|
|
Returns:
|
|
nn.Module: Activation function.
|
|
"""
|
|
|
|
act_fn = act_fn.lower()
|
|
if act_fn in ACTIVATION_FUNCTIONS:
|
|
return ACTIVATION_FUNCTIONS[act_fn]
|
|
else:
|
|
raise ValueError(f"Unsupported activation function: {act_fn}")
|
|
|
|
|
|
class FP32SiLU(nn.Module):
|
|
r"""
|
|
SiLU activation function with input upcasted to torch.float32.
|
|
"""
|
|
|
|
def __init__(self):
|
|
super().__init__()
|
|
|
|
def forward(self, inputs: torch.Tensor) -> torch.Tensor:
|
|
return F.silu(inputs.float(), inplace=False).to(inputs.dtype)
|
|
|
|
|
|
class GELU(nn.Module):
|
|
r"""
|
|
GELU activation function with tanh approximation support with `approximate="tanh"`.
|
|
|
|
Parameters:
|
|
dim_in (`int`): The number of channels in the input.
|
|
dim_out (`int`): The number of channels in the output.
|
|
approximate (`str`, *optional*, defaults to `"none"`): If `"tanh"`, use tanh approximation.
|
|
bias (`bool`, defaults to True): Whether to use a bias in the linear layer.
|
|
"""
|
|
|
|
def __init__(self, dim_in: int, dim_out: int, approximate: str = "none", bias: bool = True):
|
|
super().__init__()
|
|
self.proj = nn.Linear(dim_in, dim_out, bias=bias)
|
|
self.approximate = approximate
|
|
|
|
def gelu(self, gate: torch.Tensor) -> torch.Tensor:
|
|
if gate.device.type != "mps":
|
|
return F.gelu(gate, approximate=self.approximate)
|
|
# mps: gelu is not implemented for float16
|
|
return F.gelu(gate.to(dtype=torch.float32), approximate=self.approximate).to(dtype=gate.dtype)
|
|
|
|
def forward(self, hidden_states):
|
|
hidden_states = self.proj(hidden_states)
|
|
hidden_states = self.gelu(hidden_states)
|
|
return hidden_states
|
|
|
|
|
|
class GEGLU(nn.Module):
|
|
r"""
|
|
A [variant](https://arxiv.org/abs/2002.05202) of the gated linear unit activation function.
|
|
|
|
Parameters:
|
|
dim_in (`int`): The number of channels in the input.
|
|
dim_out (`int`): The number of channels in the output.
|
|
bias (`bool`, defaults to True): Whether to use a bias in the linear layer.
|
|
"""
|
|
|
|
def __init__(self, dim_in: int, dim_out: int, bias: bool = True):
|
|
super().__init__()
|
|
self.proj = nn.Linear(dim_in, dim_out * 2, bias=bias)
|
|
|
|
def gelu(self, gate: torch.Tensor) -> torch.Tensor:
|
|
if gate.device.type != "mps":
|
|
return F.gelu(gate)
|
|
# mps: gelu is not implemented for float16
|
|
return F.gelu(gate.to(dtype=torch.float32)).to(dtype=gate.dtype)
|
|
|
|
def forward(self, hidden_states, *args, **kwargs):
|
|
if len(args) > 0 or kwargs.get("scale", None) is not None:
|
|
deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`."
|
|
print("scale", "1.0.0", deprecation_message)
|
|
hidden_states = self.proj(hidden_states)
|
|
hidden_states, gate = hidden_states.chunk(2, dim=-1)
|
|
return hidden_states * self.gelu(gate)
|
|
|
|
|
|
class SwiGLU(nn.Module):
|
|
r"""
|
|
A [variant](https://arxiv.org/abs/2002.05202) of the gated linear unit activation function. It's similar to `GEGLU`
|
|
but uses SiLU / Swish instead of GeLU.
|
|
|
|
Parameters:
|
|
dim_in (`int`): The number of channels in the input.
|
|
dim_out (`int`): The number of channels in the output.
|
|
bias (`bool`, defaults to True): Whether to use a bias in the linear layer.
|
|
"""
|
|
|
|
def __init__(self, dim_in: int, dim_out: int, bias: bool = True):
|
|
super().__init__()
|
|
self.proj = nn.Linear(dim_in, dim_out * 2, bias=bias)
|
|
self.activation = nn.SiLU()
|
|
|
|
def forward(self, hidden_states):
|
|
hidden_states = self.proj(hidden_states)
|
|
hidden_states, gate = hidden_states.chunk(2, dim=-1)
|
|
return hidden_states * self.activation(gate)
|
|
|
|
|
|
class ApproximateGELU(nn.Module):
|
|
r"""
|
|
The approximate form of the Gaussian Error Linear Unit (GELU). For more details, see section 2 of this
|
|
[paper](https://arxiv.org/abs/1606.08415).
|
|
|
|
Parameters:
|
|
dim_in (`int`): The number of channels in the input.
|
|
dim_out (`int`): The number of channels in the output.
|
|
bias (`bool`, defaults to True): Whether to use a bias in the linear layer.
|
|
"""
|
|
|
|
def __init__(self, dim_in: int, dim_out: int, bias: bool = True):
|
|
super().__init__()
|
|
self.proj = nn.Linear(dim_in, dim_out, bias=bias)
|
|
|
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
|
x = self.proj(x)
|
|
return x * torch.sigmoid(1.702 * x)
|
|
|
|
|
|
def randn_tensor(
|
|
shape: Union[Tuple, List],
|
|
generator: Optional[Union[List["torch.Generator"], "torch.Generator"]] = None,
|
|
device: Optional["torch.device"] = None,
|
|
dtype: Optional["torch.dtype"] = None,
|
|
layout: Optional["torch.layout"] = None,
|
|
):
|
|
"""A helper function to create random tensors on the desired `device` with the desired `dtype`. When
|
|
passing a list of generators, you can seed each batch size individually. If CPU generators are passed, the tensor
|
|
is always created on the CPU.
|
|
"""
|
|
# device on which tensor is created defaults to device
|
|
rand_device = device
|
|
batch_size = shape[0]
|
|
|
|
layout = layout or torch.strided
|
|
device = device or torch.device("cpu")
|
|
|
|
if generator is not None:
|
|
gen_device_type = generator.device.type if not isinstance(generator, list) else generator[0].device.type
|
|
if gen_device_type != device.type and gen_device_type == "cpu":
|
|
rand_device = "cpu"
|
|
if device != "mps":
|
|
print(
|
|
f"The passed generator was created on 'cpu' even though a tensor on {device} was expected."
|
|
f" Tensors will be created on 'cpu' and then moved to {device}. Note that one can probably"
|
|
f" slighly speed up this function by passing a generator that was created on the {device} device."
|
|
)
|
|
elif gen_device_type != device.type and gen_device_type == "cuda":
|
|
raise ValueError(f"Cannot generate a {device} tensor from a generator of type {gen_device_type}.")
|
|
|
|
# make sure generator list of length 1 is treated like a non-list
|
|
if isinstance(generator, list) and len(generator) == 1:
|
|
generator = generator[0]
|
|
|
|
if isinstance(generator, list):
|
|
shape = (1,) + shape[1:]
|
|
latents = [
|
|
torch.randn(shape, generator=generator[i], device=rand_device, dtype=dtype, layout=layout)
|
|
for i in range(batch_size)
|
|
]
|
|
latents = torch.cat(latents, dim=0).to(device)
|
|
else:
|
|
latents = torch.randn(shape, generator=generator, device=rand_device, dtype=dtype, layout=layout).to(device)
|
|
|
|
return latents
|
|
|
|
|
|
def get_timestep_embedding(
|
|
timesteps: torch.Tensor,
|
|
embedding_dim: int,
|
|
flip_sin_to_cos: bool = False,
|
|
downscale_freq_shift: float = 1,
|
|
scale: float = 1,
|
|
max_period: int = 10000,
|
|
):
|
|
"""
|
|
This matches the implementation in Denoising Diffusion Probabilistic Models: Create sinusoidal timestep embeddings.
|
|
|
|
Args
|
|
timesteps (torch.Tensor):
|
|
a 1-D Tensor of N indices, one per batch element. These may be fractional.
|
|
embedding_dim (int):
|
|
the dimension of the output.
|
|
flip_sin_to_cos (bool):
|
|
Whether the embedding order should be `cos, sin` (if True) or `sin, cos` (if False)
|
|
downscale_freq_shift (float):
|
|
Controls the delta between frequencies between dimensions
|
|
scale (float):
|
|
Scaling factor applied to the embeddings.
|
|
max_period (int):
|
|
Controls the maximum frequency of the embeddings
|
|
Returns
|
|
torch.Tensor: an [N x dim] Tensor of positional embeddings.
|
|
"""
|
|
assert len(timesteps.shape) == 1, "Timesteps should be a 1d-array"
|
|
|
|
half_dim = embedding_dim // 2
|
|
exponent = -math.log(max_period) * torch.arange(
|
|
start=0, end=half_dim, dtype=torch.float32, device=timesteps.device
|
|
)
|
|
exponent = exponent / (half_dim - downscale_freq_shift)
|
|
|
|
emb = torch.exp(exponent)
|
|
emb = timesteps[:, None].float() * emb[None, :]
|
|
|
|
# scale embeddings
|
|
emb = scale * emb
|
|
|
|
# concat sine and cosine embeddings
|
|
emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=-1)
|
|
|
|
# flip sine and cosine embeddings
|
|
if flip_sin_to_cos:
|
|
emb = torch.cat([emb[:, half_dim:], emb[:, :half_dim]], dim=-1)
|
|
|
|
# zero pad
|
|
if embedding_dim % 2 == 1:
|
|
emb = torch.nn.functional.pad(emb, (0, 1, 0, 0))
|
|
return emb
|
|
|
|
|
|
|
|
def get_1d_sincos_pos_embed_from_grid(embed_dim, pos):
|
|
"""
|
|
embed_dim: output dimension for each position pos: a list of positions to be encoded: size (M,) out: (M, D)
|
|
"""
|
|
if embed_dim % 2 != 0:
|
|
raise ValueError("embed_dim must be divisible by 2")
|
|
|
|
omega = np.arange(embed_dim // 2, dtype=np.float64)
|
|
omega /= embed_dim / 2.0
|
|
omega = 1.0 / 10000**omega # (D/2,)
|
|
|
|
pos = pos.reshape(-1) # (M,)
|
|
out = np.einsum("m,d->md", pos, omega) # (M, D/2), outer product
|
|
|
|
emb_sin = np.sin(out) # (M, D/2)
|
|
emb_cos = np.cos(out) # (M, D/2)
|
|
|
|
emb = np.concatenate([emb_sin, emb_cos], axis=1) # (M, D)
|
|
return emb
|
|
|
|
|
|
def get_2d_sincos_pos_embed_from_grid(embed_dim, grid):
|
|
if embed_dim % 2 != 0:
|
|
raise ValueError("embed_dim must be divisible by 2")
|
|
|
|
# use half of dimensions to encode grid_h
|
|
emb_h = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[0]) # (H*W, D/2)
|
|
emb_w = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[1]) # (H*W, D/2)
|
|
|
|
emb = np.concatenate([emb_h, emb_w], axis=1) # (H*W, D)
|
|
return emb
|
|
|
|
def get_3d_sincos_pos_embed(
|
|
embed_dim: int,
|
|
spatial_size: Union[int, Tuple[int, int]],
|
|
temporal_size: int,
|
|
spatial_interpolation_scale: float = 1.0,
|
|
temporal_interpolation_scale: float = 1.0,
|
|
) -> np.ndarray:
|
|
r"""
|
|
Args:
|
|
embed_dim (`int`):
|
|
spatial_size (`int` or `Tuple[int, int]`):
|
|
temporal_size (`int`):
|
|
spatial_interpolation_scale (`float`, defaults to 1.0):
|
|
temporal_interpolation_scale (`float`, defaults to 1.0):
|
|
"""
|
|
if embed_dim % 4 != 0:
|
|
raise ValueError("`embed_dim` must be divisible by 4")
|
|
if isinstance(spatial_size, int):
|
|
spatial_size = (spatial_size, spatial_size)
|
|
|
|
embed_dim_spatial = 3 * embed_dim // 4
|
|
embed_dim_temporal = embed_dim // 4
|
|
|
|
# 1. Spatial
|
|
grid_h = np.arange(spatial_size[1], dtype=np.float32) / spatial_interpolation_scale
|
|
grid_w = np.arange(spatial_size[0], dtype=np.float32) / spatial_interpolation_scale
|
|
grid = np.meshgrid(grid_w, grid_h) # here w goes first
|
|
grid = np.stack(grid, axis=0)
|
|
|
|
grid = grid.reshape([2, 1, spatial_size[1], spatial_size[0]])
|
|
pos_embed_spatial = get_2d_sincos_pos_embed_from_grid(embed_dim_spatial, grid)
|
|
|
|
# 2. Temporal
|
|
grid_t = np.arange(temporal_size, dtype=np.float32) / temporal_interpolation_scale
|
|
pos_embed_temporal = get_1d_sincos_pos_embed_from_grid(embed_dim_temporal, grid_t)
|
|
|
|
# 3. Concat
|
|
pos_embed_spatial = pos_embed_spatial[np.newaxis, :, :]
|
|
pos_embed_spatial = np.repeat(pos_embed_spatial, temporal_size, axis=0) # [T, H*W, D // 4 * 3]
|
|
|
|
pos_embed_temporal = pos_embed_temporal[:, np.newaxis, :]
|
|
pos_embed_temporal = np.repeat(pos_embed_temporal, spatial_size[0] * spatial_size[1], axis=1) # [T, H*W, D // 4]
|
|
|
|
pos_embed = np.concatenate([pos_embed_temporal, pos_embed_spatial], axis=-1) # [T, H*W, D]
|
|
return pos_embed
|
|
|
|
|
|
def apply_rotary_emb(
|
|
x: torch.Tensor,
|
|
freqs_cis: Union[torch.Tensor, Tuple[torch.Tensor]],
|
|
use_real: bool = True,
|
|
use_real_unbind_dim: int = -1,
|
|
) -> Tuple[torch.Tensor, torch.Tensor]:
|
|
"""
|
|
Apply rotary embeddings to input tensors using the given frequency tensor. This function applies rotary embeddings
|
|
to the given query or key 'x' tensors using the provided frequency tensor 'freqs_cis'. The input tensors are
|
|
reshaped as complex numbers, and the frequency tensor is reshaped for broadcasting compatibility. The resulting
|
|
tensors contain rotary embeddings and are returned as real tensors.
|
|
|
|
Args:
|
|
x (`torch.Tensor`):
|
|
Query or key tensor to apply rotary embeddings. [B, H, S, D] xk (torch.Tensor): Key tensor to apply
|
|
freqs_cis (`Tuple[torch.Tensor]`): Precomputed frequency tensor for complex exponentials. ([S, D], [S, D],)
|
|
|
|
Returns:
|
|
Tuple[torch.Tensor, torch.Tensor]: Tuple of modified query tensor and key tensor with rotary embeddings.
|
|
"""
|
|
if use_real:
|
|
cos, sin = freqs_cis # [S, D]
|
|
cos = cos[None, None]
|
|
sin = sin[None, None]
|
|
cos, sin = cos.to(x.device), sin.to(x.device)
|
|
|
|
if use_real_unbind_dim == -1:
|
|
# Used for flux, cogvideox, hunyuan-dit
|
|
x_real, x_imag = x.reshape(*x.shape[:-1], -1, 2).unbind(-1) # [B, S, H, D//2]
|
|
x_rotated = torch.stack([-x_imag, x_real], dim=-1).flatten(3)
|
|
elif use_real_unbind_dim == -2:
|
|
# Used for Stable Audio
|
|
x_real, x_imag = x.reshape(*x.shape[:-1], 2, -1).unbind(-2) # [B, S, H, D//2]
|
|
x_rotated = torch.cat([-x_imag, x_real], dim=-1)
|
|
else:
|
|
raise ValueError(f"`use_real_unbind_dim={use_real_unbind_dim}` but should be -1 or -2.")
|
|
|
|
out = (x.float() * cos + x_rotated.float() * sin).to(x.dtype)
|
|
|
|
return out
|
|
else:
|
|
# used for lumina
|
|
x_rotated = torch.view_as_complex(x.float().reshape(*x.shape[:-1], -1, 2))
|
|
freqs_cis = freqs_cis.unsqueeze(2)
|
|
x_out = torch.view_as_real(x_rotated * freqs_cis).flatten(3)
|
|
|
|
return x_out.type_as(x)
|
|
|
|
|
|
def get_1d_rotary_pos_embed(
|
|
dim: int,
|
|
pos: Union[np.ndarray, int],
|
|
theta: float = 10000.0,
|
|
use_real=False,
|
|
linear_factor=1.0,
|
|
ntk_factor=1.0,
|
|
repeat_interleave_real=True,
|
|
freqs_dtype=torch.float32, # torch.float32, torch.float64 (flux)
|
|
):
|
|
"""
|
|
Precompute the frequency tensor for complex exponentials (cis) with given dimensions.
|
|
|
|
This function calculates a frequency tensor with complex exponentials using the given dimension 'dim' and the end
|
|
index 'end'. The 'theta' parameter scales the frequencies. The returned tensor contains complex values in complex64
|
|
data type.
|
|
|
|
Args:
|
|
dim (`int`): Dimension of the frequency tensor.
|
|
pos (`np.ndarray` or `int`): Position indices for the frequency tensor. [S] or scalar
|
|
theta (`float`, *optional*, defaults to 10000.0):
|
|
Scaling factor for frequency computation. Defaults to 10000.0.
|
|
use_real (`bool`, *optional*):
|
|
If True, return real part and imaginary part separately. Otherwise, return complex numbers.
|
|
linear_factor (`float`, *optional*, defaults to 1.0):
|
|
Scaling factor for the context extrapolation. Defaults to 1.0.
|
|
ntk_factor (`float`, *optional*, defaults to 1.0):
|
|
Scaling factor for the NTK-Aware RoPE. Defaults to 1.0.
|
|
repeat_interleave_real (`bool`, *optional*, defaults to `True`):
|
|
If `True` and `use_real`, real part and imaginary part are each interleaved with themselves to reach `dim`.
|
|
Otherwise, they are concateanted with themselves.
|
|
freqs_dtype (`torch.float32` or `torch.float64`, *optional*, defaults to `torch.float32`):
|
|
the dtype of the frequency tensor.
|
|
Returns:
|
|
`torch.Tensor`: Precomputed frequency tensor with complex exponentials. [S, D/2]
|
|
"""
|
|
assert dim % 2 == 0
|
|
|
|
if isinstance(pos, int):
|
|
pos = torch.arange(pos)
|
|
if isinstance(pos, np.ndarray):
|
|
pos = torch.from_numpy(pos) # type: ignore # [S]
|
|
|
|
theta = theta * ntk_factor
|
|
freqs = (
|
|
1.0
|
|
/ (theta ** (torch.arange(0, dim, 2, dtype=freqs_dtype, device=pos.device)[: (dim // 2)] / dim))
|
|
/ linear_factor
|
|
) # [D/2]
|
|
freqs = torch.outer(pos, freqs) # type: ignore # [S, D/2]
|
|
if use_real and repeat_interleave_real:
|
|
# flux, hunyuan-dit, cogvideox
|
|
freqs_cos = freqs.cos().repeat_interleave(2, dim=1).float() # [S, D]
|
|
freqs_sin = freqs.sin().repeat_interleave(2, dim=1).float() # [S, D]
|
|
return freqs_cos, freqs_sin
|
|
elif use_real:
|
|
# stable audio
|
|
freqs_cos = torch.cat([freqs.cos(), freqs.cos()], dim=-1).float() # [S, D]
|
|
freqs_sin = torch.cat([freqs.sin(), freqs.sin()], dim=-1).float() # [S, D]
|
|
return freqs_cos, freqs_sin
|
|
else:
|
|
# lumina
|
|
freqs_cis = torch.polar(torch.ones_like(freqs), freqs) # complex64 # [S, D/2]
|
|
return freqs_cis
|
|
|
|
|
|
def get_3d_rotary_pos_embed(
|
|
embed_dim, crops_coords, grid_size, temporal_size, theta: int = 10000, use_real: bool = True
|
|
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
|
"""
|
|
RoPE for video tokens with 3D structure.
|
|
|
|
Args:
|
|
embed_dim: (`int`):
|
|
The embedding dimension size, corresponding to hidden_size_head.
|
|
crops_coords (`Tuple[int]`):
|
|
The top-left and bottom-right coordinates of the crop.
|
|
grid_size (`Tuple[int]`):
|
|
The grid size of the spatial positional embedding (height, width).
|
|
temporal_size (`int`):
|
|
The size of the temporal dimension.
|
|
theta (`float`):
|
|
Scaling factor for frequency computation.
|
|
|
|
Returns:
|
|
`torch.Tensor`: positional embedding with shape `(temporal_size * grid_size[0] * grid_size[1], embed_dim/2)`.
|
|
"""
|
|
if use_real is not True:
|
|
raise ValueError(" `use_real = False` is not currently supported for get_3d_rotary_pos_embed")
|
|
start, stop = crops_coords
|
|
grid_size_h, grid_size_w = grid_size
|
|
grid_h = np.linspace(start[0], stop[0], grid_size_h, endpoint=False, dtype=np.float32)
|
|
grid_w = np.linspace(start[1], stop[1], grid_size_w, endpoint=False, dtype=np.float32)
|
|
grid_t = np.linspace(0, temporal_size, temporal_size, endpoint=False, dtype=np.float32)
|
|
|
|
# Compute dimensions for each axis
|
|
dim_t = embed_dim // 4
|
|
dim_h = embed_dim // 8 * 3
|
|
dim_w = embed_dim // 8 * 3
|
|
|
|
# Temporal frequencies
|
|
freqs_t = get_1d_rotary_pos_embed(dim_t, grid_t, use_real=True)
|
|
# Spatial frequencies for height and width
|
|
freqs_h = get_1d_rotary_pos_embed(dim_h, grid_h, use_real=True)
|
|
freqs_w = get_1d_rotary_pos_embed(dim_w, grid_w, use_real=True)
|
|
|
|
# BroadCast and concatenate temporal and spaial frequencie (height and width) into a 3d tensor
|
|
def combine_time_height_width(freqs_t, freqs_h, freqs_w):
|
|
freqs_t = freqs_t[:, None, None, :].expand(
|
|
-1, grid_size_h, grid_size_w, -1
|
|
) # temporal_size, grid_size_h, grid_size_w, dim_t
|
|
freqs_h = freqs_h[None, :, None, :].expand(
|
|
temporal_size, -1, grid_size_w, -1
|
|
) # temporal_size, grid_size_h, grid_size_2, dim_h
|
|
freqs_w = freqs_w[None, None, :, :].expand(
|
|
temporal_size, grid_size_h, -1, -1
|
|
) # temporal_size, grid_size_h, grid_size_2, dim_w
|
|
|
|
freqs = torch.cat(
|
|
[freqs_t, freqs_h, freqs_w], dim=-1
|
|
) # temporal_size, grid_size_h, grid_size_w, (dim_t + dim_h + dim_w)
|
|
freqs = freqs.view(
|
|
temporal_size * grid_size_h * grid_size_w, -1
|
|
) # (temporal_size * grid_size_h * grid_size_w), (dim_t + dim_h + dim_w)
|
|
return freqs
|
|
|
|
t_cos, t_sin = freqs_t # both t_cos and t_sin has shape: temporal_size, dim_t
|
|
h_cos, h_sin = freqs_h # both h_cos and h_sin has shape: grid_size_h, dim_h
|
|
w_cos, w_sin = freqs_w # both w_cos and w_sin has shape: grid_size_w, dim_w
|
|
cos = combine_time_height_width(t_cos, h_cos, w_cos)
|
|
sin = combine_time_height_width(t_sin, h_sin, w_sin)
|
|
return cos, sin
|
|
|
|
|
|
def get_resize_crop_region_for_grid(src, tgt_width, tgt_height):
|
|
tw = tgt_width
|
|
th = tgt_height
|
|
h, w = src
|
|
r = h / w
|
|
if r > (th / tw):
|
|
resize_height = th
|
|
resize_width = int(round(th / h * w))
|
|
else:
|
|
resize_width = tw
|
|
resize_height = int(round(tw / w * h))
|
|
|
|
crop_top = int(round((th - resize_height) / 2.0))
|
|
crop_left = int(round((tw - resize_width) / 2.0))
|
|
|
|
return (crop_top, crop_left), (crop_top + resize_height, crop_left + resize_width)
|