初始化项目,由ModelHub XC社区提供模型
Model: EmpathicRobotics/vla-1.7b-qwen3-v2 Source: Original Platform
This commit is contained in:
63
tools/decode/vendor/cosmos_tokenizer/modules/__init__.py
vendored
Normal file
63
tools/decode/vendor/cosmos_tokenizer/modules/__init__.py
vendored
Normal file
@@ -0,0 +1,63 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# 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.
|
||||
from enum import Enum
|
||||
|
||||
from cosmos_tokenizer.modules.distributions import (
|
||||
GaussianDistribution,
|
||||
IdentityDistribution,
|
||||
)
|
||||
from cosmos_tokenizer.modules.layers2d import Decoder, Encoder
|
||||
from cosmos_tokenizer.modules.layers3d import (
|
||||
DecoderBase,
|
||||
DecoderFactorized,
|
||||
EncoderBase,
|
||||
EncoderFactorized,
|
||||
)
|
||||
from cosmos_tokenizer.modules.quantizers import (
|
||||
FSQuantizer,
|
||||
LFQuantizer,
|
||||
ResidualFSQuantizer,
|
||||
VectorQuantizer,
|
||||
)
|
||||
|
||||
|
||||
class EncoderType(Enum):
|
||||
Default = Encoder
|
||||
|
||||
|
||||
class DecoderType(Enum):
|
||||
Default = Decoder
|
||||
|
||||
|
||||
class Encoder3DType(Enum):
|
||||
BASE = EncoderBase
|
||||
FACTORIZED = EncoderFactorized
|
||||
|
||||
|
||||
class Decoder3DType(Enum):
|
||||
BASE = DecoderBase
|
||||
FACTORIZED = DecoderFactorized
|
||||
|
||||
|
||||
class ContinuousFormulation(Enum):
|
||||
VAE = GaussianDistribution
|
||||
AE = IdentityDistribution
|
||||
|
||||
|
||||
class DiscreteQuantizer(Enum):
|
||||
VQ = VectorQuantizer
|
||||
LFQ = LFQuantizer
|
||||
FSQ = FSQuantizer
|
||||
RESFSQ = ResidualFSQuantizer
|
||||
41
tools/decode/vendor/cosmos_tokenizer/modules/distributions.py
vendored
Normal file
41
tools/decode/vendor/cosmos_tokenizer/modules/distributions.py
vendored
Normal file
@@ -0,0 +1,41 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# 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.
|
||||
"""The distribution modes to use for continuous image tokenizers."""
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
class IdentityDistribution(torch.nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
|
||||
def forward(self, parameters):
|
||||
return parameters, (torch.tensor([0.0]), torch.tensor([0.0]))
|
||||
|
||||
|
||||
class GaussianDistribution(torch.nn.Module):
|
||||
def __init__(self, min_logvar: float = -30.0, max_logvar: float = 20.0):
|
||||
super().__init__()
|
||||
self.min_logvar = min_logvar
|
||||
self.max_logvar = max_logvar
|
||||
|
||||
def sample(self, mean, logvar):
|
||||
std = torch.exp(0.5 * logvar)
|
||||
return mean + std * torch.randn_like(mean)
|
||||
|
||||
def forward(self, parameters):
|
||||
mean, logvar = torch.chunk(parameters, 2, dim=1)
|
||||
logvar = torch.clamp(logvar, self.min_logvar, self.max_logvar)
|
||||
return self.sample(mean, logvar), (mean, logvar)
|
||||
368
tools/decode/vendor/cosmos_tokenizer/modules/layers2d.py
vendored
Normal file
368
tools/decode/vendor/cosmos_tokenizer/modules/layers2d.py
vendored
Normal file
@@ -0,0 +1,368 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# 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.
|
||||
"""The model definition for Continuous 2D layers
|
||||
|
||||
Adapted from: https://github.com/CompVis/stable-diffusion/blob/
|
||||
21f890f9da3cfbeaba8e2ac3c425ee9e998d5229/ldm/modules/diffusionmodules/model.py
|
||||
|
||||
[Copyright (c) 2022 Robin Rombach and Patrick Esser and contributors]
|
||||
https://github.com/CompVis/stable-diffusion/blob/
|
||||
21f890f9da3cfbeaba8e2ac3c425ee9e998d5229/LICENSE
|
||||
"""
|
||||
|
||||
import math
|
||||
|
||||
import numpy as np
|
||||
|
||||
# pytorch_diffusion + derived encoder decoder
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from loguru import logger as logging
|
||||
|
||||
from cosmos_tokenizer.modules.patching import Patcher, UnPatcher
|
||||
from cosmos_tokenizer.modules.utils import Normalize, nonlinearity
|
||||
|
||||
|
||||
class Upsample(nn.Module):
|
||||
def __init__(self, in_channels: int):
|
||||
super().__init__()
|
||||
self.conv = nn.Conv2d(
|
||||
in_channels, in_channels, kernel_size=3, stride=1, padding=1
|
||||
)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
x = x.repeat_interleave(2, dim=2).repeat_interleave(2, dim=3)
|
||||
return self.conv(x)
|
||||
|
||||
|
||||
class Downsample(nn.Module):
|
||||
def __init__(self, in_channels: int):
|
||||
super().__init__()
|
||||
self.conv = nn.Conv2d(
|
||||
in_channels, in_channels, kernel_size=3, stride=2, padding=0
|
||||
)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
pad = (0, 1, 0, 1)
|
||||
x = F.pad(x, pad, mode="constant", value=0)
|
||||
return self.conv(x)
|
||||
|
||||
|
||||
class ResnetBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
in_channels: int,
|
||||
out_channels: int = None,
|
||||
dropout: float,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__()
|
||||
self.in_channels = in_channels
|
||||
out_channels = in_channels if out_channels is None else out_channels
|
||||
|
||||
self.norm1 = Normalize(in_channels)
|
||||
self.conv1 = nn.Conv2d(
|
||||
in_channels, out_channels, kernel_size=3, stride=1, padding=1
|
||||
)
|
||||
self.norm2 = Normalize(out_channels)
|
||||
self.dropout = nn.Dropout(dropout)
|
||||
self.conv2 = nn.Conv2d(
|
||||
out_channels, out_channels, kernel_size=3, stride=1, padding=1
|
||||
)
|
||||
self.nin_shortcut = (
|
||||
nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=1, padding=0)
|
||||
if in_channels != out_channels
|
||||
else nn.Identity()
|
||||
)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
h = x
|
||||
h = self.norm1(h)
|
||||
h = nonlinearity(h)
|
||||
h = self.conv1(h)
|
||||
|
||||
h = self.norm2(h)
|
||||
h = nonlinearity(h)
|
||||
h = self.dropout(h)
|
||||
h = self.conv2(h)
|
||||
|
||||
x = self.nin_shortcut(x)
|
||||
|
||||
return x + h
|
||||
|
||||
|
||||
class AttnBlock(nn.Module):
|
||||
def __init__(self, in_channels: int):
|
||||
super().__init__()
|
||||
|
||||
self.norm = Normalize(in_channels)
|
||||
self.q = nn.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0)
|
||||
self.k = nn.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0)
|
||||
self.v = nn.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0)
|
||||
self.proj_out = nn.Conv2d(
|
||||
in_channels, in_channels, kernel_size=1, stride=1, padding=0
|
||||
)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
# TODO (freda): Consider reusing implementations in Attn `imaginaire`,
|
||||
# since than one is gonna be based on TransformerEngine's attn op,
|
||||
# w/c could ease CP implementations.
|
||||
h_ = x
|
||||
h_ = self.norm(h_)
|
||||
q = self.q(h_)
|
||||
k = self.k(h_)
|
||||
v = self.v(h_)
|
||||
|
||||
# compute attention
|
||||
b, c, h, w = q.shape
|
||||
q = q.reshape(b, c, h * w)
|
||||
q = q.permute(0, 2, 1)
|
||||
k = k.reshape(b, c, h * w)
|
||||
w_ = torch.bmm(q, k)
|
||||
w_ = w_ * (int(c) ** (-0.5))
|
||||
w_ = F.softmax(w_, dim=2)
|
||||
|
||||
# attend to values
|
||||
v = v.reshape(b, c, h * w)
|
||||
w_ = w_.permute(0, 2, 1)
|
||||
h_ = torch.bmm(v, w_)
|
||||
h_ = h_.reshape(b, c, h, w)
|
||||
|
||||
h_ = self.proj_out(h_)
|
||||
|
||||
return x + h_
|
||||
|
||||
|
||||
class Encoder(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int,
|
||||
channels: int,
|
||||
channels_mult: list[int],
|
||||
num_res_blocks: int,
|
||||
attn_resolutions: list[int],
|
||||
dropout: float,
|
||||
resolution: int,
|
||||
z_channels: int,
|
||||
spatial_compression: int,
|
||||
**ignore_kwargs,
|
||||
):
|
||||
super().__init__()
|
||||
self.num_resolutions = len(channels_mult)
|
||||
self.num_res_blocks = num_res_blocks
|
||||
|
||||
# Patcher.
|
||||
patch_size = ignore_kwargs.get("patch_size", 1)
|
||||
self.patcher = Patcher(
|
||||
patch_size, ignore_kwargs.get("patch_method", "rearrange")
|
||||
)
|
||||
in_channels = in_channels * patch_size * patch_size
|
||||
|
||||
# calculate the number of downsample operations
|
||||
self.num_downsamples = int(math.log2(spatial_compression)) - int(
|
||||
math.log2(patch_size)
|
||||
)
|
||||
assert (
|
||||
self.num_downsamples <= self.num_resolutions
|
||||
), f"we can only downsample {self.num_resolutions} times at most"
|
||||
|
||||
# downsampling
|
||||
self.conv_in = torch.nn.Conv2d(
|
||||
in_channels, channels, kernel_size=3, stride=1, padding=1
|
||||
)
|
||||
|
||||
curr_res = resolution // patch_size
|
||||
in_ch_mult = (1,) + tuple(channels_mult)
|
||||
self.in_ch_mult = in_ch_mult
|
||||
self.down = nn.ModuleList()
|
||||
for i_level in range(self.num_resolutions):
|
||||
block = nn.ModuleList()
|
||||
attn = nn.ModuleList()
|
||||
block_in = channels * in_ch_mult[i_level]
|
||||
block_out = channels * channels_mult[i_level]
|
||||
for _ in range(self.num_res_blocks):
|
||||
block.append(
|
||||
ResnetBlock(
|
||||
in_channels=block_in,
|
||||
out_channels=block_out,
|
||||
dropout=dropout,
|
||||
)
|
||||
)
|
||||
block_in = block_out
|
||||
if curr_res in attn_resolutions:
|
||||
attn.append(AttnBlock(block_in))
|
||||
down = nn.Module()
|
||||
down.block = block
|
||||
down.attn = attn
|
||||
if i_level < self.num_downsamples:
|
||||
down.downsample = Downsample(block_in)
|
||||
curr_res = curr_res // 2
|
||||
self.down.append(down)
|
||||
|
||||
# middle
|
||||
self.mid = nn.Module()
|
||||
self.mid.block_1 = ResnetBlock(
|
||||
in_channels=block_in, out_channels=block_in, dropout=dropout
|
||||
)
|
||||
self.mid.attn_1 = AttnBlock(block_in)
|
||||
self.mid.block_2 = ResnetBlock(
|
||||
in_channels=block_in, out_channels=block_in, dropout=dropout
|
||||
)
|
||||
|
||||
# end
|
||||
self.norm_out = Normalize(block_in)
|
||||
self.conv_out = torch.nn.Conv2d(
|
||||
block_in, z_channels, kernel_size=3, stride=1, padding=1
|
||||
)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
x = self.patcher(x)
|
||||
|
||||
# downsampling
|
||||
hs = [self.conv_in(x)]
|
||||
for i_level in range(self.num_resolutions):
|
||||
for i_block in range(self.num_res_blocks):
|
||||
h = self.down[i_level].block[i_block](hs[-1])
|
||||
if len(self.down[i_level].attn) > 0:
|
||||
h = self.down[i_level].attn[i_block](h)
|
||||
hs.append(h)
|
||||
if i_level < self.num_downsamples:
|
||||
hs.append(self.down[i_level].downsample(hs[-1]))
|
||||
|
||||
# middle
|
||||
h = hs[-1]
|
||||
h = self.mid.block_1(h)
|
||||
h = self.mid.attn_1(h)
|
||||
h = self.mid.block_2(h)
|
||||
|
||||
# end
|
||||
h = self.norm_out(h)
|
||||
h = nonlinearity(h)
|
||||
h = self.conv_out(h)
|
||||
return h
|
||||
|
||||
|
||||
class Decoder(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
out_channels: int,
|
||||
channels: int,
|
||||
channels_mult: list[int],
|
||||
num_res_blocks: int,
|
||||
attn_resolutions: int,
|
||||
dropout: float,
|
||||
resolution: int,
|
||||
z_channels: int,
|
||||
spatial_compression: int,
|
||||
**ignore_kwargs,
|
||||
):
|
||||
super().__init__()
|
||||
self.num_resolutions = len(channels_mult)
|
||||
self.num_res_blocks = num_res_blocks
|
||||
|
||||
# UnPatcher.
|
||||
patch_size = ignore_kwargs.get("patch_size", 1)
|
||||
self.unpatcher = UnPatcher(
|
||||
patch_size, ignore_kwargs.get("patch_method", "rearrange")
|
||||
)
|
||||
out_ch = out_channels * patch_size * patch_size
|
||||
|
||||
# calculate the number of upsample operations
|
||||
self.num_upsamples = int(math.log2(spatial_compression)) - int(
|
||||
math.log2(patch_size)
|
||||
)
|
||||
assert (
|
||||
self.num_upsamples <= self.num_resolutions
|
||||
), f"we can only upsample {self.num_resolutions} times at most"
|
||||
|
||||
block_in = channels * channels_mult[self.num_resolutions - 1]
|
||||
curr_res = (resolution // patch_size) // 2 ** (self.num_resolutions - 1)
|
||||
self.z_shape = (1, z_channels, curr_res, curr_res)
|
||||
logging.info(
|
||||
"Working with z of shape {} = {} dimensions.".format(
|
||||
self.z_shape, np.prod(self.z_shape)
|
||||
)
|
||||
)
|
||||
|
||||
# z to block_in
|
||||
self.conv_in = torch.nn.Conv2d(
|
||||
z_channels, block_in, kernel_size=3, stride=1, padding=1
|
||||
)
|
||||
|
||||
# middle
|
||||
self.mid = nn.Module()
|
||||
self.mid.block_1 = ResnetBlock(
|
||||
in_channels=block_in, out_channels=block_in, dropout=dropout
|
||||
)
|
||||
self.mid.attn_1 = AttnBlock(block_in)
|
||||
self.mid.block_2 = ResnetBlock(
|
||||
in_channels=block_in, out_channels=block_in, dropout=dropout
|
||||
)
|
||||
|
||||
# upsampling
|
||||
self.up = nn.ModuleList()
|
||||
for i_level in reversed(range(self.num_resolutions)):
|
||||
block = nn.ModuleList()
|
||||
attn = nn.ModuleList()
|
||||
block_out = channels * channels_mult[i_level]
|
||||
for _ in range(self.num_res_blocks + 1):
|
||||
block.append(
|
||||
ResnetBlock(
|
||||
in_channels=block_in,
|
||||
out_channels=block_out,
|
||||
dropout=dropout,
|
||||
)
|
||||
)
|
||||
block_in = block_out
|
||||
if curr_res in attn_resolutions:
|
||||
attn.append(AttnBlock(block_in))
|
||||
up = nn.Module()
|
||||
up.block = block
|
||||
up.attn = attn
|
||||
if i_level >= (self.num_resolutions - self.num_upsamples):
|
||||
up.upsample = Upsample(block_in)
|
||||
curr_res = curr_res * 2
|
||||
self.up.insert(0, up)
|
||||
|
||||
# end
|
||||
self.norm_out = Normalize(block_in)
|
||||
self.conv_out = torch.nn.Conv2d(
|
||||
block_in, out_ch, kernel_size=3, stride=1, padding=1
|
||||
)
|
||||
|
||||
def forward(self, z: torch.Tensor) -> torch.Tensor:
|
||||
h = self.conv_in(z)
|
||||
|
||||
# middle
|
||||
h = self.mid.block_1(h)
|
||||
h = self.mid.attn_1(h)
|
||||
h = self.mid.block_2(h)
|
||||
|
||||
# upsampling
|
||||
for i_level in reversed(range(self.num_resolutions)):
|
||||
for i_block in range(self.num_res_blocks + 1):
|
||||
h = self.up[i_level].block[i_block](h)
|
||||
if len(self.up[i_level].attn) > 0:
|
||||
h = self.up[i_level].attn[i_block](h)
|
||||
if i_level >= (self.num_resolutions - self.num_upsamples):
|
||||
h = self.up[i_level].upsample(h)
|
||||
|
||||
h = self.norm_out(h)
|
||||
h = nonlinearity(h)
|
||||
h = self.conv_out(h)
|
||||
h = self.unpatcher(h)
|
||||
return h
|
||||
1040
tools/decode/vendor/cosmos_tokenizer/modules/layers3d.py
vendored
Normal file
1040
tools/decode/vendor/cosmos_tokenizer/modules/layers3d.py
vendored
Normal file
File diff suppressed because it is too large
Load Diff
356
tools/decode/vendor/cosmos_tokenizer/modules/patching.py
vendored
Normal file
356
tools/decode/vendor/cosmos_tokenizer/modules/patching.py
vendored
Normal file
@@ -0,0 +1,356 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# 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.
|
||||
"""The patcher and unpatcher implementation for 2D and 3D data.
|
||||
|
||||
The idea of Haar wavelet is to compute LL, LH, HL, HH component as two 1D convolutions.
|
||||
One on the rows and one on the columns.
|
||||
For example, in 1D signal, we have [a, b], then the low-freq compoenent is [a + b] / 2 and high-freq is [a - b] / 2.
|
||||
We can use a 1D convolution with kernel [1, 1] and stride 2 to represent the L component.
|
||||
For H component, we can use a 1D convolution with kernel [1, -1] and stride 2.
|
||||
Although in principle, we typically only do additional Haar wavelet over the LL component. But here we do it for all
|
||||
as we need to support downsampling for more than 2x.
|
||||
For example, 4x downsampling can be done by 2x Haar and additional 2x Haar, and the shape would be.
|
||||
[3, 256, 256] -> [12, 128, 128] -> [48, 64, 64]
|
||||
"""
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from einops import rearrange
|
||||
|
||||
_WAVELETS = {
|
||||
"haar": torch.tensor([0.7071067811865476, 0.7071067811865476]),
|
||||
"rearrange": torch.tensor([1.0, 1.0]),
|
||||
}
|
||||
_PERSISTENT = False
|
||||
|
||||
|
||||
class Patcher(torch.nn.Module):
|
||||
"""A module to convert image tensors into patches using torch operations.
|
||||
|
||||
The main difference from `class Patching` is that this module implements
|
||||
all operations using torch, rather than python or numpy, for efficiency purpose.
|
||||
|
||||
It's bit-wise identical to the Patching module outputs, with the added
|
||||
benefit of being torch.jit scriptable.
|
||||
"""
|
||||
|
||||
def __init__(self, patch_size=1, patch_method="haar"):
|
||||
super().__init__()
|
||||
self.patch_size = patch_size
|
||||
self.patch_method = patch_method
|
||||
self.register_buffer(
|
||||
"wavelets", _WAVELETS[patch_method], persistent=_PERSISTENT
|
||||
)
|
||||
self.range = range(int(torch.log2(torch.tensor(self.patch_size)).item()))
|
||||
self.register_buffer(
|
||||
"_arange",
|
||||
torch.arange(_WAVELETS[patch_method].shape[0]),
|
||||
persistent=_PERSISTENT,
|
||||
)
|
||||
for param in self.parameters():
|
||||
param.requires_grad = False
|
||||
|
||||
def forward(self, x):
|
||||
if self.patch_method == "haar":
|
||||
return self._haar(x)
|
||||
elif self.patch_method == "rearrange":
|
||||
return self._arrange(x)
|
||||
else:
|
||||
raise ValueError("Unknown patch method: " + self.patch_method)
|
||||
|
||||
def _dwt(self, x, mode="reflect", rescale=False):
|
||||
dtype = x.dtype
|
||||
h = self.wavelets
|
||||
|
||||
n = h.shape[0]
|
||||
g = x.shape[1]
|
||||
hl = h.flip(0).reshape(1, 1, -1).repeat(g, 1, 1)
|
||||
hh = (h * ((-1) ** self._arange)).reshape(1, 1, -1).repeat(g, 1, 1)
|
||||
hh = hh.to(dtype=dtype)
|
||||
hl = hl.to(dtype=dtype)
|
||||
|
||||
x = F.pad(x, pad=(n - 2, n - 1, n - 2, n - 1), mode=mode).to(dtype)
|
||||
xl = F.conv2d(x, hl.unsqueeze(2), groups=g, stride=(1, 2))
|
||||
xh = F.conv2d(x, hh.unsqueeze(2), groups=g, stride=(1, 2))
|
||||
xll = F.conv2d(xl, hl.unsqueeze(3), groups=g, stride=(2, 1))
|
||||
xlh = F.conv2d(xl, hh.unsqueeze(3), groups=g, stride=(2, 1))
|
||||
xhl = F.conv2d(xh, hl.unsqueeze(3), groups=g, stride=(2, 1))
|
||||
xhh = F.conv2d(xh, hh.unsqueeze(3), groups=g, stride=(2, 1))
|
||||
|
||||
out = torch.cat([xll, xlh, xhl, xhh], dim=1)
|
||||
if rescale:
|
||||
out = out / 2
|
||||
return out
|
||||
|
||||
def _haar(self, x):
|
||||
for _ in self.range:
|
||||
x = self._dwt(x, rescale=True)
|
||||
return x
|
||||
|
||||
def _arrange(self, x):
|
||||
x = rearrange(
|
||||
x,
|
||||
"b c (h p1) (w p2) -> b (c p1 p2) h w",
|
||||
p1=self.patch_size,
|
||||
p2=self.patch_size,
|
||||
).contiguous()
|
||||
return x
|
||||
|
||||
|
||||
class Patcher3D(Patcher):
|
||||
"""A 3D discrete wavelet transform for video data, expects 5D tensor, i.e. a batch of videos."""
|
||||
|
||||
def __init__(self, patch_size=1, patch_method="haar"):
|
||||
super().__init__(patch_method=patch_method, patch_size=patch_size)
|
||||
self.register_buffer(
|
||||
"patch_size_buffer",
|
||||
patch_size * torch.ones([1], dtype=torch.int32),
|
||||
persistent=_PERSISTENT,
|
||||
)
|
||||
|
||||
def _dwt(self, x, wavelet, mode="reflect", rescale=False):
|
||||
dtype = x.dtype
|
||||
h = self.wavelets
|
||||
|
||||
n = h.shape[0]
|
||||
g = x.shape[1]
|
||||
hl = h.flip(0).reshape(1, 1, -1).repeat(g, 1, 1)
|
||||
hh = (h * ((-1) ** self._arange)).reshape(1, 1, -1).repeat(g, 1, 1)
|
||||
hh = hh.to(dtype=dtype)
|
||||
hl = hl.to(dtype=dtype)
|
||||
|
||||
# Handles temporal axis.
|
||||
x = F.pad(
|
||||
x, pad=(max(0, n - 2), n - 1, n - 2, n - 1, n - 2, n - 1), mode=mode
|
||||
).to(dtype)
|
||||
xl = F.conv3d(x, hl.unsqueeze(3).unsqueeze(4), groups=g, stride=(2, 1, 1))
|
||||
xh = F.conv3d(x, hh.unsqueeze(3).unsqueeze(4), groups=g, stride=(2, 1, 1))
|
||||
|
||||
# Handles spatial axes.
|
||||
xll = F.conv3d(xl, hl.unsqueeze(2).unsqueeze(4), groups=g, stride=(1, 2, 1))
|
||||
xlh = F.conv3d(xl, hh.unsqueeze(2).unsqueeze(4), groups=g, stride=(1, 2, 1))
|
||||
xhl = F.conv3d(xh, hl.unsqueeze(2).unsqueeze(4), groups=g, stride=(1, 2, 1))
|
||||
xhh = F.conv3d(xh, hh.unsqueeze(2).unsqueeze(4), groups=g, stride=(1, 2, 1))
|
||||
|
||||
xlll = F.conv3d(xll, hl.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2))
|
||||
xllh = F.conv3d(xll, hh.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2))
|
||||
xlhl = F.conv3d(xlh, hl.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2))
|
||||
xlhh = F.conv3d(xlh, hh.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2))
|
||||
xhll = F.conv3d(xhl, hl.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2))
|
||||
xhlh = F.conv3d(xhl, hh.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2))
|
||||
xhhl = F.conv3d(xhh, hl.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2))
|
||||
xhhh = F.conv3d(xhh, hh.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2))
|
||||
|
||||
out = torch.cat([xlll, xllh, xlhl, xlhh, xhll, xhlh, xhhl, xhhh], dim=1)
|
||||
if rescale:
|
||||
out = out / (2 * torch.sqrt(torch.tensor(2.0)))
|
||||
return out
|
||||
|
||||
def _haar(self, x):
|
||||
xi, xv = torch.split(x, [1, x.shape[2] - 1], dim=2)
|
||||
x = torch.cat([xi.repeat_interleave(self.patch_size, dim=2), xv], dim=2)
|
||||
for _ in self.range:
|
||||
x = self._dwt(x, "haar", rescale=True)
|
||||
return x
|
||||
|
||||
def _arrange(self, x):
|
||||
xi, xv = torch.split(x, [1, x.shape[2] - 1], dim=2)
|
||||
x = torch.cat([xi.repeat_interleave(self.patch_size, dim=2), xv], dim=2)
|
||||
x = rearrange(
|
||||
x,
|
||||
"b c (t p1) (h p2) (w p3) -> b (c p1 p2 p3) t h w",
|
||||
p1=self.patch_size,
|
||||
p2=self.patch_size,
|
||||
p3=self.patch_size,
|
||||
).contiguous()
|
||||
return x
|
||||
|
||||
|
||||
class UnPatcher(torch.nn.Module):
|
||||
"""A module to convert patches into image tensorsusing torch operations.
|
||||
|
||||
The main difference from `class Unpatching` is that this module implements
|
||||
all operations using torch, rather than python or numpy, for efficiency purpose.
|
||||
|
||||
It's bit-wise identical to the Unpatching module outputs, with the added
|
||||
benefit of being torch.jit scriptable.
|
||||
"""
|
||||
|
||||
def __init__(self, patch_size=1, patch_method="haar"):
|
||||
super().__init__()
|
||||
self.patch_size = patch_size
|
||||
self.patch_method = patch_method
|
||||
self.register_buffer(
|
||||
"wavelets", _WAVELETS[patch_method], persistent=_PERSISTENT
|
||||
)
|
||||
self.range = range(int(torch.log2(torch.tensor(self.patch_size)).item()))
|
||||
self.register_buffer(
|
||||
"_arange",
|
||||
torch.arange(_WAVELETS[patch_method].shape[0]),
|
||||
persistent=_PERSISTENT,
|
||||
)
|
||||
for param in self.parameters():
|
||||
param.requires_grad = False
|
||||
|
||||
def forward(self, x):
|
||||
if self.patch_method == "haar":
|
||||
return self._ihaar(x)
|
||||
elif self.patch_method == "rearrange":
|
||||
return self._iarrange(x)
|
||||
else:
|
||||
raise ValueError("Unknown patch method: " + self.patch_method)
|
||||
|
||||
def _idwt(self, x, wavelet="haar", mode="reflect", rescale=False):
|
||||
dtype = x.dtype
|
||||
h = self.wavelets
|
||||
n = h.shape[0]
|
||||
|
||||
g = x.shape[1] // 4
|
||||
hl = h.flip([0]).reshape(1, 1, -1).repeat([g, 1, 1])
|
||||
hh = (h * ((-1) ** self._arange)).reshape(1, 1, -1).repeat(g, 1, 1)
|
||||
hh = hh.to(dtype=dtype)
|
||||
hl = hl.to(dtype=dtype)
|
||||
|
||||
xll, xlh, xhl, xhh = torch.chunk(x.to(dtype), 4, dim=1)
|
||||
|
||||
# Inverse transform.
|
||||
yl = torch.nn.functional.conv_transpose2d(
|
||||
xll, hl.unsqueeze(3), groups=g, stride=(2, 1), padding=(n - 2, 0)
|
||||
)
|
||||
yl += torch.nn.functional.conv_transpose2d(
|
||||
xlh, hh.unsqueeze(3), groups=g, stride=(2, 1), padding=(n - 2, 0)
|
||||
)
|
||||
yh = torch.nn.functional.conv_transpose2d(
|
||||
xhl, hl.unsqueeze(3), groups=g, stride=(2, 1), padding=(n - 2, 0)
|
||||
)
|
||||
yh += torch.nn.functional.conv_transpose2d(
|
||||
xhh, hh.unsqueeze(3), groups=g, stride=(2, 1), padding=(n - 2, 0)
|
||||
)
|
||||
y = torch.nn.functional.conv_transpose2d(
|
||||
yl, hl.unsqueeze(2), groups=g, stride=(1, 2), padding=(0, n - 2)
|
||||
)
|
||||
y += torch.nn.functional.conv_transpose2d(
|
||||
yh, hh.unsqueeze(2), groups=g, stride=(1, 2), padding=(0, n - 2)
|
||||
)
|
||||
|
||||
if rescale:
|
||||
y = y * 2
|
||||
return y
|
||||
|
||||
def _ihaar(self, x):
|
||||
for _ in self.range:
|
||||
x = self._idwt(x, "haar", rescale=True)
|
||||
return x
|
||||
|
||||
def _iarrange(self, x):
|
||||
x = rearrange(
|
||||
x,
|
||||
"b (c p1 p2) h w -> b c (h p1) (w p2)",
|
||||
p1=self.patch_size,
|
||||
p2=self.patch_size,
|
||||
)
|
||||
return x
|
||||
|
||||
|
||||
class UnPatcher3D(UnPatcher):
|
||||
"""A 3D inverse discrete wavelet transform for video wavelet decompositions."""
|
||||
|
||||
def __init__(self, patch_size=1, patch_method="haar"):
|
||||
super().__init__(patch_method=patch_method, patch_size=patch_size)
|
||||
|
||||
def _idwt(self, x, wavelet="haar", mode="reflect", rescale=False):
|
||||
dtype = x.dtype
|
||||
h = self.wavelets
|
||||
n = h.shape[0]
|
||||
|
||||
g = x.shape[1] // 8 # split into 8 spatio-temporal filtered tesnors.
|
||||
hl = h.flip([0]).reshape(1, 1, -1).repeat([g, 1, 1])
|
||||
hh = (h * ((-1) ** self._arange)).reshape(1, 1, -1).repeat(g, 1, 1)
|
||||
hl = hl.to(dtype=dtype)
|
||||
hh = hh.to(dtype=dtype)
|
||||
|
||||
xlll, xllh, xlhl, xlhh, xhll, xhlh, xhhl, xhhh = torch.chunk(x, 8, dim=1)
|
||||
|
||||
# Height height transposed convolutions.
|
||||
xll = F.conv_transpose3d(
|
||||
xlll, hl.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2)
|
||||
)
|
||||
xll += F.conv_transpose3d(
|
||||
xllh, hh.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2)
|
||||
)
|
||||
|
||||
xlh = F.conv_transpose3d(
|
||||
xlhl, hl.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2)
|
||||
)
|
||||
xlh += F.conv_transpose3d(
|
||||
xlhh, hh.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2)
|
||||
)
|
||||
|
||||
xhl = F.conv_transpose3d(
|
||||
xhll, hl.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2)
|
||||
)
|
||||
xhl += F.conv_transpose3d(
|
||||
xhlh, hh.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2)
|
||||
)
|
||||
|
||||
xhh = F.conv_transpose3d(
|
||||
xhhl, hl.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2)
|
||||
)
|
||||
xhh += F.conv_transpose3d(
|
||||
xhhh, hh.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2)
|
||||
)
|
||||
|
||||
# Handles width transposed convolutions.
|
||||
xl = F.conv_transpose3d(
|
||||
xll, hl.unsqueeze(2).unsqueeze(4), groups=g, stride=(1, 2, 1)
|
||||
)
|
||||
xl += F.conv_transpose3d(
|
||||
xlh, hh.unsqueeze(2).unsqueeze(4), groups=g, stride=(1, 2, 1)
|
||||
)
|
||||
xh = F.conv_transpose3d(
|
||||
xhl, hl.unsqueeze(2).unsqueeze(4), groups=g, stride=(1, 2, 1)
|
||||
)
|
||||
xh += F.conv_transpose3d(
|
||||
xhh, hh.unsqueeze(2).unsqueeze(4), groups=g, stride=(1, 2, 1)
|
||||
)
|
||||
|
||||
# Handles time axis transposed convolutions.
|
||||
x = F.conv_transpose3d(
|
||||
xl, hl.unsqueeze(3).unsqueeze(4), groups=g, stride=(2, 1, 1)
|
||||
)
|
||||
x += F.conv_transpose3d(
|
||||
xh, hh.unsqueeze(3).unsqueeze(4), groups=g, stride=(2, 1, 1)
|
||||
)
|
||||
|
||||
if rescale:
|
||||
x = x * (2 * torch.sqrt(torch.tensor(2.0)))
|
||||
return x
|
||||
|
||||
def _ihaar(self, x):
|
||||
for _ in self.range:
|
||||
x = self._idwt(x, "haar", rescale=True)
|
||||
x = x[:, :, self.patch_size - 1 :, ...]
|
||||
return x
|
||||
|
||||
def _iarrange(self, x):
|
||||
x = rearrange(
|
||||
x,
|
||||
"b (c p1 p2 p3) t h w -> b c (t p1) (h p2) (w p3)",
|
||||
p1=self.patch_size,
|
||||
p2=self.patch_size,
|
||||
p3=self.patch_size,
|
||||
)
|
||||
x = x[:, :, self.patch_size - 1 :, ...]
|
||||
return x
|
||||
546
tools/decode/vendor/cosmos_tokenizer/modules/quantizers.py
vendored
Normal file
546
tools/decode/vendor/cosmos_tokenizer/modules/quantizers.py
vendored
Normal file
@@ -0,0 +1,546 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# 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.
|
||||
"""Quantizers for discrete image and video tokenization."""
|
||||
|
||||
from typing import Optional
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from einops import reduce
|
||||
from loguru import logger as logging
|
||||
|
||||
from cosmos_tokenizer.modules.utils import (
|
||||
default,
|
||||
entropy,
|
||||
pack_one,
|
||||
rearrange,
|
||||
round_ste,
|
||||
unpack_one,
|
||||
)
|
||||
|
||||
|
||||
class ResidualFSQuantizer(nn.Module):
|
||||
"""Residual Finite Scalar Quantization
|
||||
|
||||
Follows Algorithm 1. in https://arxiv.org/pdf/2107.03312.pdf
|
||||
"""
|
||||
|
||||
def __init__(self, levels: list[int], num_quantizers: int, **ignore_kwargs):
|
||||
super().__init__()
|
||||
self.dtype = ignore_kwargs.get("dtype", torch.float32)
|
||||
self.layers = nn.ModuleList(
|
||||
[FSQuantizer(levels=levels) for _ in range(num_quantizers)]
|
||||
)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
indices_stack = []
|
||||
residual = x
|
||||
quantized_out = 0
|
||||
loss_out = 0
|
||||
for i, layer in enumerate(self.layers):
|
||||
quant_indices, z, loss = layer(residual)
|
||||
indices_stack.append(quant_indices)
|
||||
residual = residual - z.detach()
|
||||
quantized_out = quantized_out + z
|
||||
loss_out = loss_out + loss
|
||||
self.residual = residual
|
||||
indices = torch.stack(indices_stack, dim=1)
|
||||
return indices, quantized_out.to(self.dtype), loss_out.to(self.dtype)
|
||||
|
||||
def indices_to_codes(self, indices_stack: torch.Tensor) -> torch.Tensor:
|
||||
quantized_out = 0
|
||||
for layer, indices in zip(self.layers, indices_stack.transpose(0, 1)):
|
||||
quantized_out += layer.indices_to_codes(indices)
|
||||
return quantized_out
|
||||
|
||||
|
||||
class FSQuantizer(nn.Module):
|
||||
"""Finite Scalar Quantization: VQ-VAE Made Simple - https://arxiv.org/abs/2309.15505
|
||||
|
||||
Code adapted from Jax version in Appendix A.1.
|
||||
|
||||
Adapted from: https://github.com/lucidrains/vector-quantize-pytorch/blob/9502a1f447876d53fd37685b226bf28f250dc4a3/
|
||||
vector_quantize_pytorch/finite_scalar_quantization.py
|
||||
[Copyright (c) 2020 Phil Wang]
|
||||
https://github.com/lucidrains/vector-quantize-pytorch/blob/9502a1f447876d53fd37685b226bf28f250dc4a3/LICENSE
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
levels: list[int],
|
||||
dim: Optional[int] = None,
|
||||
num_codebooks=1,
|
||||
keep_num_codebooks_dim: Optional[bool] = None,
|
||||
scale: Optional[float] = None,
|
||||
**ignore_kwargs,
|
||||
):
|
||||
super().__init__()
|
||||
self.dtype = ignore_kwargs.get("dtype", torch.bfloat16)
|
||||
_levels = torch.tensor(levels, dtype=torch.int32)
|
||||
self.register_buffer("_levels", _levels, persistent=False)
|
||||
|
||||
_basis = torch.cumprod(
|
||||
torch.tensor([1] + levels[:-1]), dim=0, dtype=torch.int32
|
||||
)
|
||||
self.register_buffer("_basis", _basis, persistent=False)
|
||||
|
||||
self.scale = scale
|
||||
|
||||
codebook_dim = len(levels)
|
||||
self.codebook_dim = codebook_dim
|
||||
|
||||
effective_codebook_dim = codebook_dim * num_codebooks
|
||||
self.num_codebooks = num_codebooks
|
||||
self.effective_codebook_dim = effective_codebook_dim
|
||||
|
||||
keep_num_codebooks_dim = default(keep_num_codebooks_dim, num_codebooks > 1)
|
||||
assert not (num_codebooks > 1 and not keep_num_codebooks_dim)
|
||||
self.keep_num_codebooks_dim = keep_num_codebooks_dim
|
||||
|
||||
self.dim = default(dim, len(_levels) * num_codebooks)
|
||||
|
||||
has_projections = self.dim != effective_codebook_dim
|
||||
self.project_in = (
|
||||
nn.Linear(self.dim, effective_codebook_dim)
|
||||
if has_projections
|
||||
else nn.Identity()
|
||||
)
|
||||
self.project_out = (
|
||||
nn.Linear(effective_codebook_dim, self.dim)
|
||||
if has_projections
|
||||
else nn.Identity()
|
||||
)
|
||||
self.has_projections = has_projections
|
||||
|
||||
self.codebook_size = self._levels.prod().item()
|
||||
|
||||
implicit_codebook = self.indices_to_codes(
|
||||
torch.arange(self.codebook_size), project_out=False
|
||||
)
|
||||
self.register_buffer("implicit_codebook", implicit_codebook, persistent=False)
|
||||
|
||||
def bound(self, z: torch.Tensor, eps: float = 1e-3) -> torch.Tensor:
|
||||
"""Bound `z`, an array of shape (..., d)."""
|
||||
half_l = (self._levels - 1) * (1 + eps) / 2
|
||||
offset = torch.where(self._levels % 2 == 0, 0.5, 0.0)
|
||||
shift = (offset / half_l).atanh()
|
||||
return (z + shift).tanh() * half_l - offset
|
||||
|
||||
def quantize(self, z: torch.Tensor) -> torch.Tensor:
|
||||
"""Quantizes z, returns quantized zhat, same shape as z."""
|
||||
quantized = round_ste(self.bound(z))
|
||||
half_width = self._levels // 2 # Renormalize to [-1, 1].
|
||||
return quantized / half_width
|
||||
|
||||
def _scale_and_shift(self, zhat_normalized: torch.Tensor) -> torch.Tensor:
|
||||
half_width = self._levels // 2
|
||||
return (zhat_normalized * half_width) + half_width
|
||||
|
||||
def _scale_and_shift_inverse(self, zhat: torch.Tensor) -> torch.Tensor:
|
||||
half_width = self._levels // 2
|
||||
return (zhat - half_width) / half_width
|
||||
|
||||
def codes_to_indices(self, zhat: torch.Tensor) -> torch.Tensor:
|
||||
"""Converts a `code` to an index in the codebook."""
|
||||
assert zhat.shape[-1] == self.codebook_dim
|
||||
zhat = self._scale_and_shift(zhat).float()
|
||||
return (zhat * self._basis).sum(dim=-1).to(torch.int32)
|
||||
|
||||
def indices_to_codes(self, indices: torch.Tensor, project_out=True) -> torch.Tensor:
|
||||
"""Inverse of `codes_to_indices`."""
|
||||
is_img_or_video = indices.ndim >= (3 + int(self.keep_num_codebooks_dim))
|
||||
indices = rearrange(indices, "... -> ... 1")
|
||||
codes_non_centered = (indices // self._basis) % self._levels
|
||||
codes = self._scale_and_shift_inverse(codes_non_centered)
|
||||
|
||||
if self.keep_num_codebooks_dim:
|
||||
codes = rearrange(codes, "... c d -> ... (c d)")
|
||||
|
||||
if project_out:
|
||||
codes = self.project_out(codes)
|
||||
|
||||
if is_img_or_video:
|
||||
codes = rearrange(codes, "b ... d -> b d ...")
|
||||
|
||||
return codes.to(self.dtype)
|
||||
|
||||
def forward(self, z: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
einstein notation
|
||||
b - batch
|
||||
n - sequence (or flattened spatial dimensions)
|
||||
d - feature dimension, which is also log2(codebook size)
|
||||
c - number of codebook dim
|
||||
"""
|
||||
is_img_or_video = z.ndim >= 4
|
||||
|
||||
# standardize image or video into (batch, seq, dimension)
|
||||
|
||||
if is_img_or_video:
|
||||
z = rearrange(z, "b d ... -> b ... d")
|
||||
z, ps = pack_one(z, "b * d")
|
||||
|
||||
assert (
|
||||
z.shape[-1] == self.dim
|
||||
), f"expected dimension of {self.dim} but found dimension of {z.shape[-1]}"
|
||||
|
||||
z = self.project_in(z)
|
||||
|
||||
z = rearrange(z, "b n (c d) -> b n c d", c=self.num_codebooks)
|
||||
|
||||
codes = self.quantize(z)
|
||||
indices = self.codes_to_indices(codes)
|
||||
|
||||
codes = rearrange(codes, "b n c d -> b n (c d)")
|
||||
|
||||
out = self.project_out(codes)
|
||||
|
||||
# reconstitute image or video dimensions
|
||||
|
||||
if is_img_or_video:
|
||||
out = unpack_one(out, ps, "b * d")
|
||||
out = rearrange(out, "b ... d -> b d ...")
|
||||
indices = unpack_one(indices, ps, "b * c")
|
||||
dummy_loss = torch.zeros_like(out.mean(dim=[1, 2, 3], keepdim=True))
|
||||
else:
|
||||
dummy_loss = torch.zeros_like(out.mean(dim=[1, 2], keepdim=True)).unsqueeze(
|
||||
1
|
||||
)
|
||||
|
||||
if not self.keep_num_codebooks_dim:
|
||||
indices = rearrange(indices, "... 1 -> ...")
|
||||
|
||||
return (indices, out.to(self.dtype), dummy_loss)
|
||||
|
||||
|
||||
class VectorQuantizer(nn.Module):
|
||||
"""Improved version over VectorQuantizer. Mostly
|
||||
avoids costly matrix multiplications and allows for post-hoc remapping of indices.
|
||||
|
||||
Adapted from: https://github.com/CompVis/taming-transformers/blob/3ba01b241669f5ade541ce990f7650a3b8f65318/
|
||||
taming/modules/vqvae/quantize.py
|
||||
|
||||
[Copyright (c) 2020 Patrick Esser and Robin Rombach and Björn Ommer]
|
||||
https://github.com/CompVis/taming-transformers/blob/3ba01b241669f5ade541ce990f7650a3b8f65318/License.txt
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_embeddings: int,
|
||||
embedding_dim: int,
|
||||
beta: float = 0.25,
|
||||
remap: str = None,
|
||||
unknown_index: str = "random",
|
||||
sane_index_shape: bool = False,
|
||||
legacy: bool = True,
|
||||
use_norm=False,
|
||||
**ignore_kwargs,
|
||||
):
|
||||
super().__init__()
|
||||
self.n_e = num_embeddings
|
||||
self.e_dim = embedding_dim
|
||||
self.beta = beta
|
||||
self.legacy = legacy
|
||||
self.norm = lambda x: F.normalize(x, dim=-1) if use_norm else x
|
||||
|
||||
self.embedding = nn.Embedding(self.n_e, self.e_dim)
|
||||
self.embedding.weight.data.uniform_(-1.0 / self.n_e, 1.0 / self.n_e)
|
||||
|
||||
self.remap = remap
|
||||
if self.remap is not None:
|
||||
self.register_buffer("used", torch.tensor(np.load(self.remap)))
|
||||
self.re_embed = self.used.shape[0]
|
||||
self.unknown_index = unknown_index
|
||||
if self.unknown_index == "extra":
|
||||
self.unknown_index = self.re_embed
|
||||
self.re_embed = self.re_embed + 1
|
||||
print(
|
||||
f"Remapping {self.n_e} indices to {self.re_embed} indices. "
|
||||
f"Using {self.unknown_index} for unknown indices."
|
||||
)
|
||||
else:
|
||||
self.re_embed = num_embeddings
|
||||
|
||||
self.sane_index_shape = sane_index_shape
|
||||
self.dtype = ignore_kwargs.get("dtype", torch.float32)
|
||||
|
||||
def remap_to_used(self, inds):
|
||||
ishape = inds.shape
|
||||
assert len(ishape) > 1
|
||||
inds = inds.reshape(ishape[0], -1)
|
||||
used = self.used.to(inds)
|
||||
match = (inds[:, :, None] == used[None, None, ...]).long()
|
||||
new = match.argmax(-1)
|
||||
unknown = match.sum(2) < 1
|
||||
if self.unknown_index == "random":
|
||||
new[unknown] = torch.randint(0, self.re_embed, size=new[unknown].shape).to(
|
||||
device=new.device
|
||||
)
|
||||
else:
|
||||
new[unknown] = self.unknown_index
|
||||
return new.reshape(ishape)
|
||||
|
||||
def unmap_to_all(self, inds):
|
||||
ishape = inds.shape
|
||||
assert len(ishape) > 1
|
||||
inds = inds.reshape(ishape[0], -1)
|
||||
used = self.used.to(inds)
|
||||
if self.re_embed > self.used.shape[0]: # extra token
|
||||
inds[inds >= self.used.shape[0]] = 0 # simply set to zero
|
||||
back = torch.gather(used[None, :][inds.shape[0] * [0], :], 1, inds)
|
||||
return back.reshape(ishape)
|
||||
|
||||
def forward(self, z, temp=None, rescale_logits=False, return_logits=False):
|
||||
assert temp is None or temp == 1.0, "Only for interface compatible with Gumbel"
|
||||
assert rescale_logits is False, "Only for interface compatible with Gumbel"
|
||||
assert return_logits is False, "Only for interface compatible with Gumbel"
|
||||
z = rearrange(z, "b c h w -> b h w c").contiguous()
|
||||
z_flattened = z.view(-1, self.e_dim)
|
||||
|
||||
d = (
|
||||
torch.sum(z_flattened**2, dim=1, keepdim=True)
|
||||
+ torch.sum(self.embedding.weight**2, dim=1)
|
||||
- 2
|
||||
* torch.einsum(
|
||||
"bd,dn->bn",
|
||||
z_flattened,
|
||||
rearrange(self.embedding.weight, "n d -> d n"),
|
||||
)
|
||||
)
|
||||
|
||||
encoding_indices = torch.argmin(d, dim=1).unsqueeze(1)
|
||||
encodings = torch.zeros(encoding_indices.shape[0], self.n_e, device=z.device)
|
||||
encodings.scatter_(1, encoding_indices, 1)
|
||||
z_q = torch.matmul(encodings, self.embedding.weight).view(z.shape)
|
||||
min_encodings = None
|
||||
|
||||
z_q, z = self.norm(z_q), self.norm(z)
|
||||
|
||||
# compute loss for embedding
|
||||
commit_loss = torch.mean((z_q - z.detach()) ** 2, dim=[1, 2, 3], keepdim=True)
|
||||
emb_loss = torch.mean((z_q.detach() - z) ** 2, dim=[1, 2, 3], keepdim=True)
|
||||
if not self.legacy:
|
||||
loss = self.beta * emb_loss + commit_loss
|
||||
else:
|
||||
loss = emb_loss + self.beta * commit_loss
|
||||
|
||||
# preserve gradients
|
||||
z_q = z + (z_q - z).detach()
|
||||
avg_probs = torch.mean(encodings, dim=0)
|
||||
perplexity = torch.exp(-torch.sum(avg_probs * torch.log(avg_probs + 1e-10)))
|
||||
|
||||
# reshape back to match original input shape
|
||||
z_q = rearrange(z_q, "b h w c -> b c h w").contiguous()
|
||||
|
||||
if self.remap is not None:
|
||||
min_encoding_indices = encoding_indices.squeeze(1).reshape(
|
||||
z.shape[0], -1
|
||||
) # add batch axis
|
||||
min_encoding_indices = self.remap_to_used(encoding_indices.squeeze(1))
|
||||
min_encoding_indices = min_encoding_indices.reshape(-1, 1) # flatten
|
||||
|
||||
if self.sane_index_shape:
|
||||
min_encoding_indices = min_encoding_indices.reshape(
|
||||
z_q.shape[0], z_q.shape[2], z_q.shape[3]
|
||||
)
|
||||
|
||||
# TODO: return (indices, z_q, loss)
|
||||
return (
|
||||
z_q,
|
||||
loss,
|
||||
(
|
||||
encoding_indices.squeeze(1),
|
||||
min_encodings,
|
||||
commit_loss.mean().detach(),
|
||||
self.beta * emb_loss.mean().detach(),
|
||||
perplexity.mean().detach(),
|
||||
),
|
||||
)
|
||||
|
||||
def get_codebook_entry(self, indices, shape):
|
||||
# shape specifying (batch, height, width, channel)
|
||||
if self.remap is not None:
|
||||
indices = indices.reshape(shape[0], -1) # add batch axis
|
||||
indices = self.unmap_to_all(indices)
|
||||
indices = indices.reshape(-1) # flatten again
|
||||
|
||||
# get quantized latent vectors
|
||||
z_q = self.embedding(indices)
|
||||
|
||||
if shape is not None:
|
||||
z_q = z_q.view(shape)
|
||||
# reshape back to match original input shape
|
||||
z_q = z_q.permute(0, 3, 1, 2).contiguous()
|
||||
|
||||
return z_q
|
||||
|
||||
|
||||
class LFQuantizer(nn.Module):
|
||||
"""Lookup-Free Quantization
|
||||
|
||||
Adapted from: https://github.com/lucidrains/vector-quantize-pytorch/blob/9502a1f447876d53fd37685b226bf28f250dc4a3/
|
||||
vector_quantize_pytorch/lookup_free_quantization.py
|
||||
[Copyright (c) 2020 Phil Wang]
|
||||
https://github.com/lucidrains/vector-quantize-pytorch/blob/9502a1f447876d53fd37685b226bf28f250dc4a3/LICENSE
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
codebook_size: int,
|
||||
codebook_dim: int,
|
||||
embed_dim: Optional[int] = None, # if None, use codebook_dim
|
||||
entropy_loss_weight=0.1,
|
||||
commitment_loss_weight=0.25,
|
||||
default_temp: float = 0.01,
|
||||
entropy_loss: bool = False,
|
||||
**ignore_kwargs,
|
||||
):
|
||||
"""Lookup-Free Quantization
|
||||
|
||||
Args:
|
||||
codebook_size (int): The number of entries in the codebook.
|
||||
codebook_dim (int): The number of bits in each code.
|
||||
embed_dim (Optional[int], optional): The dimension of the input embedding. Defaults to None.
|
||||
entropy_loss_weight (float, optional): Whether to use entropy loss. Defaults to 0.1.
|
||||
commitment_loss_weight (float, optional): Weight for commitment loss. Defaults to 0.25.
|
||||
default_temp (float, optional): The temprature to use. Defaults to 0.01.
|
||||
entropy_loss (bool, optional): Flag for entropy loss. Defaults to False.
|
||||
"""
|
||||
super().__init__()
|
||||
self.entropy_loss = entropy_loss
|
||||
self.codebook_dim = codebook_dim
|
||||
self.default_temp = default_temp
|
||||
self.entrop_loss_weight = entropy_loss_weight
|
||||
self.commitment_loss_weight = commitment_loss_weight
|
||||
embed_dim = embed_dim or codebook_dim
|
||||
|
||||
has_projections = embed_dim != codebook_dim
|
||||
self.project_in = (
|
||||
nn.Linear(embed_dim, codebook_dim) if has_projections else nn.Identity()
|
||||
)
|
||||
self.project_out = (
|
||||
nn.Linear(codebook_dim, embed_dim) if has_projections else nn.Identity()
|
||||
)
|
||||
logging.info(
|
||||
f"LFQ: has_projections={has_projections}, dim_in={embed_dim}, codebook_dim={codebook_dim}"
|
||||
)
|
||||
|
||||
self.dtype = ignore_kwargs.get("dtype", torch.float32)
|
||||
|
||||
if entropy_loss:
|
||||
assert (
|
||||
2**codebook_dim == codebook_size
|
||||
), "codebook size must be 2 ** codebook_dim"
|
||||
self.codebook_size = codebook_size
|
||||
|
||||
self.register_buffer(
|
||||
"mask",
|
||||
2 ** torch.arange(codebook_dim - 1, -1, -1),
|
||||
persistent=False,
|
||||
)
|
||||
self.register_buffer("zero", torch.tensor(0.0), persistent=False)
|
||||
|
||||
all_codes = torch.arange(codebook_size)
|
||||
bits = ((all_codes[..., None].int() & self.mask) != 0).float()
|
||||
codebook = 2 * bits - 1.0
|
||||
|
||||
self.register_buffer(
|
||||
"codebook", codebook, persistent=False
|
||||
) # [codebook_size, codebook_dim]
|
||||
|
||||
def forward(self, z: torch.Tensor, temp: float = None) -> torch.Tensor:
|
||||
temp = temp or self.default_temp
|
||||
|
||||
z = rearrange(z, "b d ... -> b ... d")
|
||||
z, ps = pack_one(z, "b * d")
|
||||
z = self.project_in(z)
|
||||
|
||||
# split out number of codebooks
|
||||
z = rearrange(z, "b n (c d) -> b n c d", c=self.num_codebooks)
|
||||
|
||||
# quantization
|
||||
original_input = z
|
||||
|
||||
codebook_value = torch.ones_like(z)
|
||||
z_q = torch.where(z > 0, codebook_value, -codebook_value)
|
||||
|
||||
# preserve gradients
|
||||
z_q = z + (z_q - z).detach()
|
||||
|
||||
# commit loss
|
||||
commit_loss = ((original_input - z_q.detach()) ** 2).mean(dim=[1, 2, 3])
|
||||
|
||||
z_q = rearrange(z_q, "b n c d -> b n (c d)")
|
||||
z_q = self.project_out(z_q)
|
||||
|
||||
# reshape
|
||||
z_q = unpack_one(z_q, ps, "b * d")
|
||||
z_q = rearrange(z_q, "b ... d -> b d ...")
|
||||
|
||||
loss = self.commitment_loss_weight * commit_loss
|
||||
|
||||
# entropy loss (eq-5)
|
||||
if self.entropy_loss:
|
||||
# indices
|
||||
indices = reduce((z > 0).int() * self.mask.int(), "b n c d -> b n c", "sum")
|
||||
indices = unpack_one(indices, ps, "b * c")
|
||||
indices = rearrange(indices, "... 1 -> ...")
|
||||
|
||||
distance = -2 * torch.einsum(
|
||||
"... i d, j d -> ... i j",
|
||||
original_input,
|
||||
self.codebook.to(original_input.dtype),
|
||||
)
|
||||
prob = (-distance / temp).softmax(dim=-1)
|
||||
per_sample_entropy = entropy(prob).mean(dim=[1, 2])
|
||||
avg_prob = reduce(prob, "... c d -> c d", "mean")
|
||||
codebook_entropy = entropy(avg_prob).mean()
|
||||
entropy_aux_loss = per_sample_entropy - codebook_entropy
|
||||
|
||||
loss += self.entrop_loss_weight * entropy_aux_loss
|
||||
|
||||
# TODO: return (indices, z_q, loss)
|
||||
return (
|
||||
z_q,
|
||||
loss.unsqueeze(1).unsqueeze(1).unsqueeze(1),
|
||||
(
|
||||
indices,
|
||||
self.commitment_loss_weight * commit_loss.mean().detach(),
|
||||
self.entrop_loss_weight * entropy_aux_loss.mean().detach(),
|
||||
self.entrop_loss_weight * per_sample_entropy.mean().detach(),
|
||||
self.entrop_loss_weight * codebook_entropy.mean().detach(),
|
||||
),
|
||||
)
|
||||
else:
|
||||
return (
|
||||
z_q,
|
||||
loss.unsqueeze(1).unsqueeze(1).unsqueeze(1),
|
||||
self.commitment_loss_weight * commit_loss.mean().detach(),
|
||||
)
|
||||
|
||||
|
||||
class InvQuantizerJit(nn.Module):
|
||||
"""Use for decoder_jit to trace quantizer in discrete tokenizer"""
|
||||
|
||||
def __init__(self, quantizer):
|
||||
super().__init__()
|
||||
self.quantizer = quantizer
|
||||
|
||||
def forward(self, indices: torch.Tensor):
|
||||
codes = self.quantizer.indices_to_codes(indices)
|
||||
return codes.to(self.quantizer.dtype)
|
||||
117
tools/decode/vendor/cosmos_tokenizer/modules/utils.py
vendored
Normal file
117
tools/decode/vendor/cosmos_tokenizer/modules/utils.py
vendored
Normal file
@@ -0,0 +1,117 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# 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.
|
||||
"""Shared utilities for the networks module."""
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from einops import pack, rearrange, unpack
|
||||
|
||||
|
||||
def time2batch(x: torch.Tensor) -> tuple[torch.Tensor, int]:
|
||||
batch_size = x.shape[0]
|
||||
return rearrange(x, "b c t h w -> (b t) c h w"), batch_size
|
||||
|
||||
|
||||
def batch2time(x: torch.Tensor, batch_size: int) -> torch.Tensor:
|
||||
return rearrange(x, "(b t) c h w -> b c t h w", b=batch_size)
|
||||
|
||||
|
||||
def space2batch(x: torch.Tensor) -> tuple[torch.Tensor, int]:
|
||||
batch_size, height = x.shape[0], x.shape[-2]
|
||||
return rearrange(x, "b c t h w -> (b h w) c t"), batch_size, height
|
||||
|
||||
|
||||
def batch2space(x: torch.Tensor, batch_size: int, height: int) -> torch.Tensor:
|
||||
return rearrange(x, "(b h w) c t -> b c t h w", b=batch_size, h=height)
|
||||
|
||||
|
||||
def cast_tuple(t: Any, length: int = 1) -> Any:
|
||||
return t if isinstance(t, tuple) else ((t,) * length)
|
||||
|
||||
|
||||
def replication_pad(x):
|
||||
return torch.cat([x[:, :, :1, ...], x], dim=2)
|
||||
|
||||
|
||||
def divisible_by(num: int, den: int) -> bool:
|
||||
return (num % den) == 0
|
||||
|
||||
|
||||
def is_odd(n: int) -> bool:
|
||||
return not divisible_by(n, 2)
|
||||
|
||||
|
||||
def nonlinearity(x):
|
||||
return x * torch.sigmoid(x)
|
||||
|
||||
|
||||
def Normalize(in_channels, num_groups=32):
|
||||
return torch.nn.GroupNorm(
|
||||
num_groups=num_groups, num_channels=in_channels, eps=1e-6, affine=True
|
||||
)
|
||||
|
||||
|
||||
class CausalNormalize(torch.nn.Module):
|
||||
def __init__(self, in_channels, num_groups=1):
|
||||
super().__init__()
|
||||
self.norm = torch.nn.GroupNorm(
|
||||
num_groups=num_groups,
|
||||
num_channels=in_channels,
|
||||
eps=1e-6,
|
||||
affine=True,
|
||||
)
|
||||
self.num_groups = num_groups
|
||||
|
||||
def forward(self, x):
|
||||
# if num_groups !=1, we apply a spatio-temporal groupnorm for backward compatibility purpose.
|
||||
# All new models should use num_groups=1, otherwise causality is not guaranteed.
|
||||
if self.num_groups == 1:
|
||||
x, batch_size = time2batch(x)
|
||||
return batch2time(self.norm(x), batch_size)
|
||||
return self.norm(x)
|
||||
|
||||
|
||||
def exists(v):
|
||||
return v is not None
|
||||
|
||||
|
||||
def default(*args):
|
||||
for arg in args:
|
||||
if exists(arg):
|
||||
return arg
|
||||
return None
|
||||
|
||||
|
||||
def pack_one(t, pattern):
|
||||
return pack([t], pattern)
|
||||
|
||||
|
||||
def unpack_one(t, ps, pattern):
|
||||
return unpack(t, ps, pattern)[0]
|
||||
|
||||
|
||||
def round_ste(z: torch.Tensor) -> torch.Tensor:
|
||||
"""Round with straight through gradients."""
|
||||
zhat = z.round()
|
||||
return z + (zhat - z).detach()
|
||||
|
||||
|
||||
def log(t, eps=1e-5):
|
||||
return t.clamp(min=eps).log()
|
||||
|
||||
|
||||
def entropy(prob):
|
||||
return (-prob * log(prob)).sum(dim=-1)
|
||||
Reference in New Issue
Block a user