初始化项目,由ModelHub XC社区提供模型
Model: EmpathicRobotics/vla-1.7b-qwen3-v2 Source: Original Platform
This commit is contained in:
52
tools/decode/vendor/cosmos_tokenizer/networks/__init__.py
vendored
Normal file
52
tools/decode/vendor/cosmos_tokenizer/networks/__init__.py
vendored
Normal file
@@ -0,0 +1,52 @@
|
||||
# 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.networks.configs import (
|
||||
continuous_image as continuous_image_dict,
|
||||
)
|
||||
from cosmos_tokenizer.networks.configs import (
|
||||
discrete_image as discrete_image_dict,
|
||||
)
|
||||
from cosmos_tokenizer.networks.configs import (
|
||||
continuous_video as continuous_video_dict,
|
||||
)
|
||||
from cosmos_tokenizer.networks.configs import (
|
||||
discrete_video as discrete_video_dict,
|
||||
)
|
||||
|
||||
from cosmos_tokenizer.networks.continuous_image import ContinuousImageTokenizer
|
||||
from cosmos_tokenizer.networks.discrete_image import DiscreteImageTokenizer
|
||||
from cosmos_tokenizer.networks.continuous_video import (
|
||||
CausalContinuousVideoTokenizer,
|
||||
)
|
||||
from cosmos_tokenizer.networks.discrete_video import (
|
||||
CausalDiscreteVideoTokenizer,
|
||||
)
|
||||
|
||||
|
||||
class TokenizerConfigs(Enum):
|
||||
CI = continuous_image_dict
|
||||
DI = discrete_image_dict
|
||||
CV = continuous_video_dict
|
||||
DV = discrete_video_dict
|
||||
|
||||
|
||||
class TokenizerModels(Enum):
|
||||
CI = ContinuousImageTokenizer
|
||||
DI = DiscreteImageTokenizer
|
||||
CV = CausalContinuousVideoTokenizer
|
||||
DV = CausalDiscreteVideoTokenizer
|
||||
146
tools/decode/vendor/cosmos_tokenizer/networks/configs.py
vendored
Normal file
146
tools/decode/vendor/cosmos_tokenizer/networks/configs.py
vendored
Normal file
@@ -0,0 +1,146 @@
|
||||
# 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 default image and video tokenizer configs."""
|
||||
|
||||
from cosmos_tokenizer.modules import (
|
||||
ContinuousFormulation,
|
||||
DiscreteQuantizer,
|
||||
EncoderType,
|
||||
DecoderType,
|
||||
Encoder3DType,
|
||||
Decoder3DType,
|
||||
)
|
||||
|
||||
continuous_image = dict(
|
||||
# The attention resolution for res blocks.
|
||||
attn_resolutions=[32],
|
||||
# The base number of channels.
|
||||
channels=128,
|
||||
# The channel multipler for each resolution.
|
||||
channels_mult=[2, 4, 4],
|
||||
dropout=0.0,
|
||||
in_channels=3,
|
||||
# The spatial compression ratio.
|
||||
spatial_compression=16,
|
||||
# The number of layers in each res block.
|
||||
num_res_blocks=2,
|
||||
out_channels=3,
|
||||
resolution=1024,
|
||||
patch_size=4,
|
||||
patch_method="haar",
|
||||
# The output latent dimension (channels).
|
||||
latent_channels=16,
|
||||
# The encoder output channels just before sampling.
|
||||
# Which is also the decoder's input channels.
|
||||
z_channels=16,
|
||||
# A factor over the z_channels, to get the total channels the encoder should output.
|
||||
# For a VAE for instance, we want to output the mean and variance, so we need 2 * z_channels.
|
||||
z_factor=1,
|
||||
name="CI",
|
||||
# What formulation to use, either "AE" or "VAE".
|
||||
# Chose VAE here, since the pre-trained ckpt were of a VAE formulation.
|
||||
formulation=ContinuousFormulation.AE.name,
|
||||
# Specify type of encoder ["Default", "LiteVAE"]
|
||||
encoder=EncoderType.Default.name,
|
||||
# Specify type of decoder ["Default"]
|
||||
decoder=DecoderType.Default.name,
|
||||
)
|
||||
|
||||
discrete_image = dict(
|
||||
# The attention resolution for res blocks.
|
||||
attn_resolutions=[32],
|
||||
# The base number of channels.
|
||||
channels=128,
|
||||
# The channel multipler for each resolution.
|
||||
channels_mult=[2, 4, 4],
|
||||
dropout=0.0,
|
||||
in_channels=3,
|
||||
# The spatial compression ratio.
|
||||
spatial_compression=16,
|
||||
# The number of layers in each res block.
|
||||
num_res_blocks=2,
|
||||
out_channels=3,
|
||||
resolution=1024,
|
||||
patch_size=4,
|
||||
patch_method="haar",
|
||||
# The encoder output channels just before sampling.
|
||||
z_channels=256,
|
||||
# A factor over the z_channels, to get the total channels the encoder should output.
|
||||
# for discrete tokenization, often we directly use the vector, so z_factor=1.
|
||||
z_factor=1,
|
||||
# The quantizer of choice, VQ, LFQ, FSQ, or ResFSQ.
|
||||
quantizer=DiscreteQuantizer.FSQ.name,
|
||||
# The embedding dimension post-quantization, which is also the input channels of the decoder.
|
||||
# Which is also the output
|
||||
embedding_dim=6,
|
||||
# The number of levels to use for fine-scalar quantization.
|
||||
levels=[8, 8, 8, 5, 5, 5],
|
||||
# The number of quantizers to use for residual fine-scalar quantization.
|
||||
num_quantizers=4,
|
||||
name="DI",
|
||||
# Specify type of encoder ["Default", "LiteVAE"]
|
||||
encoder=EncoderType.Default.name,
|
||||
# Specify type of decoder ["Default"]
|
||||
decoder=DecoderType.Default.name,
|
||||
)
|
||||
|
||||
continuous_video = dict(
|
||||
attn_resolutions=[32],
|
||||
channels=128,
|
||||
channels_mult=[2, 4, 4],
|
||||
dropout=0.0,
|
||||
in_channels=3,
|
||||
num_res_blocks=2,
|
||||
out_channels=3,
|
||||
resolution=1024,
|
||||
patch_size=4,
|
||||
patch_method="haar",
|
||||
latent_channels=16,
|
||||
z_channels=16,
|
||||
z_factor=1,
|
||||
num_groups=1,
|
||||
legacy_mode=False,
|
||||
spatial_compression=8,
|
||||
temporal_compression=8,
|
||||
formulation=ContinuousFormulation.AE.name,
|
||||
encoder=Encoder3DType.FACTORIZED.name,
|
||||
decoder=Decoder3DType.FACTORIZED.name,
|
||||
name="CV",
|
||||
)
|
||||
|
||||
discrete_video = dict(
|
||||
attn_resolutions=[32],
|
||||
channels=128,
|
||||
channels_mult=[2, 4, 4],
|
||||
dropout=0.0,
|
||||
in_channels=3,
|
||||
num_res_blocks=2,
|
||||
out_channels=3,
|
||||
resolution=1024,
|
||||
patch_size=4,
|
||||
patch_method="haar",
|
||||
z_channels=16,
|
||||
z_factor=1,
|
||||
num_groups=1,
|
||||
legacy_mode=False,
|
||||
spatial_compression=16,
|
||||
temporal_compression=8,
|
||||
quantizer=DiscreteQuantizer.FSQ.name,
|
||||
embedding_dim=6,
|
||||
levels=[8, 8, 8, 5, 5, 5],
|
||||
encoder=Encoder3DType.FACTORIZED.name,
|
||||
decoder=Decoder3DType.FACTORIZED.name,
|
||||
name="DV",
|
||||
)
|
||||
104
tools/decode/vendor/cosmos_tokenizer/networks/continuous_image.py
vendored
Normal file
104
tools/decode/vendor/cosmos_tokenizer/networks/continuous_image.py
vendored
Normal file
@@ -0,0 +1,104 @@
|
||||
# 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 continuous image tokenizer with VAE or AE formulation for 2D data."""
|
||||
|
||||
from collections import OrderedDict, namedtuple
|
||||
|
||||
import torch
|
||||
from loguru import logger as logging
|
||||
from torch import nn
|
||||
|
||||
from cosmos_tokenizer.modules import (
|
||||
ContinuousFormulation,
|
||||
DecoderType,
|
||||
EncoderType,
|
||||
)
|
||||
|
||||
NetworkEval = namedtuple("NetworkEval", ["reconstructions", "posteriors", "latent"])
|
||||
|
||||
|
||||
class ContinuousImageTokenizer(nn.Module):
|
||||
def __init__(
|
||||
self, z_channels: int, z_factor: int, latent_channels: int, **kwargs
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.name = kwargs.get("name", "ContinuousImageTokenizer")
|
||||
self.latent_channels = latent_channels
|
||||
|
||||
encoder_name = kwargs.get("encoder", EncoderType.Default.name)
|
||||
self.encoder = EncoderType[encoder_name].value(
|
||||
z_channels=z_factor * z_channels, **kwargs
|
||||
)
|
||||
|
||||
decoder_name = kwargs.get("decoder", DecoderType.Default.name)
|
||||
self.decoder = DecoderType[decoder_name].value(z_channels=z_channels, **kwargs)
|
||||
|
||||
self.quant_conv = torch.nn.Conv2d(
|
||||
z_factor * z_channels, z_factor * latent_channels, 1
|
||||
)
|
||||
self.post_quant_conv = torch.nn.Conv2d(latent_channels, z_channels, 1)
|
||||
|
||||
formulation_name = kwargs.get("formulation", ContinuousFormulation.AE.name)
|
||||
self.distribution = ContinuousFormulation[formulation_name].value()
|
||||
logging.info(
|
||||
f"{self.name} based on {formulation_name} formulation, with {kwargs}."
|
||||
)
|
||||
|
||||
num_parameters = sum(param.numel() for param in self.parameters())
|
||||
logging.info(f"model={self.name}, num_parameters={num_parameters:,}")
|
||||
logging.info(
|
||||
f"z_channels={z_channels}, latent_channels={self.latent_channels}."
|
||||
)
|
||||
|
||||
def encoder_jit(self):
|
||||
return nn.Sequential(
|
||||
OrderedDict(
|
||||
[
|
||||
("encoder", self.encoder),
|
||||
("quant_conv", self.quant_conv),
|
||||
("distribution", self.distribution),
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
def decoder_jit(self):
|
||||
return nn.Sequential(
|
||||
OrderedDict(
|
||||
[
|
||||
("post_quant_conv", self.post_quant_conv),
|
||||
("decoder", self.decoder),
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
def last_decoder_layer(self):
|
||||
return self.decoder.conv_out
|
||||
|
||||
def encode(self, x):
|
||||
h = self.encoder(x)
|
||||
moments = self.quant_conv(h)
|
||||
return self.distribution(moments)
|
||||
|
||||
def decode(self, z):
|
||||
z = self.post_quant_conv(z)
|
||||
dec = self.decoder(z)
|
||||
return dec
|
||||
|
||||
def forward(self, input) -> dict[str, torch.Tensor] | NetworkEval:
|
||||
latent, posteriors = self.encode(input)
|
||||
dec = self.decode(latent)
|
||||
if self.training:
|
||||
return dict(reconstructions=dec, posteriors=posteriors, latent=latent)
|
||||
return NetworkEval(reconstructions=dec, posteriors=posteriors, latent=latent)
|
||||
118
tools/decode/vendor/cosmos_tokenizer/networks/continuous_video.py
vendored
Normal file
118
tools/decode/vendor/cosmos_tokenizer/networks/continuous_video.py
vendored
Normal file
@@ -0,0 +1,118 @@
|
||||
# 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 causal continuous video tokenizer with VAE or AE formulation for 3D data.."""
|
||||
from collections import OrderedDict, namedtuple
|
||||
|
||||
from loguru import logger as logging
|
||||
from torch import nn
|
||||
|
||||
from cosmos_tokenizer.modules import (
|
||||
ContinuousFormulation,
|
||||
Decoder3DType,
|
||||
Encoder3DType,
|
||||
)
|
||||
from cosmos_tokenizer.modules.layers3d import CausalConv3d
|
||||
|
||||
NetworkEval = namedtuple("NetworkEval", ["reconstructions", "posteriors", "latent"])
|
||||
|
||||
|
||||
class CausalContinuousVideoTokenizer(nn.Module):
|
||||
def __init__(
|
||||
self, z_channels: int, z_factor: int, latent_channels: int, **kwargs
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.name = kwargs.get("name", "CausalContinuousVideoTokenizer")
|
||||
self.latent_channels = latent_channels
|
||||
|
||||
encoder_name = kwargs.get("encoder", Encoder3DType.BASE.name)
|
||||
self.encoder = Encoder3DType[encoder_name].value(
|
||||
z_channels=z_factor * z_channels, **kwargs
|
||||
)
|
||||
if kwargs.get("temporal_compression", 4) == 4:
|
||||
kwargs["channels_mult"] = [2, 4]
|
||||
decoder_name = kwargs.get("decoder", Decoder3DType.BASE.name)
|
||||
self.decoder = Decoder3DType[decoder_name].value(
|
||||
z_channels=z_channels, **kwargs
|
||||
)
|
||||
|
||||
self.quant_conv = CausalConv3d(
|
||||
z_factor * z_channels,
|
||||
z_factor * latent_channels,
|
||||
kernel_size=1,
|
||||
padding=0,
|
||||
)
|
||||
self.post_quant_conv = CausalConv3d(
|
||||
latent_channels, z_channels, kernel_size=1, padding=0
|
||||
)
|
||||
|
||||
formulation_name = kwargs.get("formulation", ContinuousFormulation.AE.name)
|
||||
self.distribution = ContinuousFormulation[formulation_name].value()
|
||||
logging.info(
|
||||
f"{self.name} based on {formulation_name} formulation, with {kwargs}."
|
||||
)
|
||||
|
||||
num_parameters = sum(param.numel() for param in self.parameters())
|
||||
logging.info(f"model={self.name}, num_parameters={num_parameters:,}")
|
||||
logging.info(
|
||||
f"z_channels={z_channels}, latent_channels={self.latent_channels}."
|
||||
)
|
||||
|
||||
def encoder_jit(self):
|
||||
return nn.Sequential(
|
||||
OrderedDict(
|
||||
[
|
||||
("encoder", self.encoder),
|
||||
("quant_conv", self.quant_conv),
|
||||
("distribution", self.distribution),
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
def decoder_jit(self):
|
||||
return nn.Sequential(
|
||||
OrderedDict(
|
||||
[
|
||||
("post_quant_conv", self.post_quant_conv),
|
||||
("decoder", self.decoder),
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
def last_decoder_layer(self):
|
||||
return self.decoder.conv_out
|
||||
|
||||
def encode(self, x):
|
||||
h = self.encoder(x)
|
||||
moments = self.quant_conv(h)
|
||||
return self.distribution(moments)
|
||||
|
||||
def decode(self, z):
|
||||
z = self.post_quant_conv(z)
|
||||
return self.decoder(z)
|
||||
|
||||
def forward(self, input):
|
||||
latent, posteriors = self.encode(input)
|
||||
reconstructions = self.decode(latent)
|
||||
if self.training:
|
||||
return dict(
|
||||
reconstructions=reconstructions,
|
||||
posteriors=posteriors,
|
||||
latent=latent,
|
||||
)
|
||||
return NetworkEval(
|
||||
reconstructions=reconstructions,
|
||||
posteriors=posteriors,
|
||||
latent=latent,
|
||||
)
|
||||
129
tools/decode/vendor/cosmos_tokenizer/networks/discrete_image.py
vendored
Normal file
129
tools/decode/vendor/cosmos_tokenizer/networks/discrete_image.py
vendored
Normal file
@@ -0,0 +1,129 @@
|
||||
# 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 network definition for discrete image tokenization with VQ, LFQ, FSQ or ResidualFSQ."""
|
||||
from collections import OrderedDict, namedtuple
|
||||
|
||||
import torch
|
||||
from loguru import logger as logging
|
||||
from torch import nn
|
||||
|
||||
from cosmos_tokenizer.modules import DecoderType, DiscreteQuantizer, EncoderType
|
||||
from cosmos_tokenizer.modules.quantizers import InvQuantizerJit
|
||||
|
||||
NetworkEval = namedtuple("NetworkEval", ["reconstructions", "quant_loss", "quant_info"])
|
||||
|
||||
|
||||
class DiscreteImageTokenizer(nn.Module):
|
||||
def __init__(self, z_channels: int, embedding_dim: int, **kwargs) -> None:
|
||||
super().__init__()
|
||||
self.name = kwargs.get("name", "DiscreteImageTokenizer")
|
||||
self.embedding_dim = embedding_dim
|
||||
|
||||
encoder_name = kwargs.get("encoder", EncoderType.Default.name)
|
||||
self.encoder = EncoderType[encoder_name].value(z_channels=z_channels, **kwargs)
|
||||
|
||||
decoder_name = kwargs.get("decoder", DecoderType.Default.name)
|
||||
self.decoder = DecoderType[decoder_name].value(z_channels=z_channels, **kwargs)
|
||||
self.quant_conv = nn.Conv2d(z_channels, embedding_dim, 1)
|
||||
self.post_quant_conv = nn.Conv2d(embedding_dim, z_channels, 1)
|
||||
|
||||
quantizer_name = kwargs.get("quantizer", DiscreteQuantizer.RESFSQ.name)
|
||||
if quantizer_name == DiscreteQuantizer.VQ.name:
|
||||
assert (
|
||||
"num_embeddings" in kwargs
|
||||
), f"`num_embeddings` must be provided for {quantizer_name}."
|
||||
kwargs.update(dict(embedding_dim=embedding_dim))
|
||||
elif quantizer_name == DiscreteQuantizer.LFQ.name:
|
||||
assert (
|
||||
"codebook_size" in kwargs
|
||||
), f"`codebook_size` must be provided for {quantizer_name}."
|
||||
assert (
|
||||
"codebook_dim" in kwargs
|
||||
), f"`codebook_dim` must be provided for {quantizer_name}."
|
||||
elif quantizer_name == DiscreteQuantizer.FSQ.name:
|
||||
assert (
|
||||
"levels" in kwargs
|
||||
), f"`levels` must be provided for {quantizer_name}."
|
||||
elif quantizer_name == DiscreteQuantizer.RESFSQ.name:
|
||||
assert (
|
||||
"levels" in kwargs
|
||||
), f"`levels` must be provided for {quantizer_name}.name."
|
||||
assert (
|
||||
"num_quantizers" in kwargs
|
||||
), f"`num_quantizers` must be provided for {quantizer_name}."
|
||||
self.quantizer = DiscreteQuantizer[quantizer_name].value(**kwargs)
|
||||
logging.info(f"{self.name} based on {quantizer_name}-VAE, with {kwargs}.")
|
||||
|
||||
num_parameters = sum(param.numel() for param in self.parameters())
|
||||
logging.info(f"model={self.name}, num_parameters={num_parameters:,}")
|
||||
logging.info(f"z_channels={z_channels}, embedding_dim={self.embedding_dim}.")
|
||||
|
||||
def to(self, *args, **kwargs):
|
||||
setattr(self.quantizer, "dtype", kwargs.get("dtype", torch.bfloat16))
|
||||
return super(DiscreteImageTokenizer, self).to(*args, **kwargs)
|
||||
|
||||
def encoder_jit(self):
|
||||
return nn.Sequential(
|
||||
OrderedDict(
|
||||
[
|
||||
("encoder", self.encoder),
|
||||
("quant_conv", self.quant_conv),
|
||||
("quantizer", self.quantizer),
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
def decoder_jit(self):
|
||||
return nn.Sequential(
|
||||
OrderedDict(
|
||||
[
|
||||
("inv_quant", InvQuantizerJit(self.quantizer)),
|
||||
("post_quant_conv", self.post_quant_conv),
|
||||
("decoder", self.decoder),
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
def last_decoder_layer(self):
|
||||
return self.decoder.conv_out
|
||||
|
||||
def encode(self, x):
|
||||
h = self.encoder(x)
|
||||
h = self.quant_conv(h)
|
||||
return self.quantizer(h)
|
||||
|
||||
def decode(self, quant):
|
||||
quant = self.post_quant_conv(quant)
|
||||
return self.decoder(quant)
|
||||
|
||||
def decode_code(self, code_b):
|
||||
quant_b = self.quantizer.indices_to_codes(code_b)
|
||||
quant_b = self.post_quant_conv(quant_b)
|
||||
return self.decoder(quant_b)
|
||||
|
||||
def forward(self, input):
|
||||
quant_info, quant_codes, quant_loss = self.encode(input)
|
||||
reconstructions = self.decode(quant_codes)
|
||||
if self.training:
|
||||
return dict(
|
||||
reconstructions=reconstructions,
|
||||
quant_loss=quant_loss,
|
||||
quant_info=quant_info,
|
||||
)
|
||||
return NetworkEval(
|
||||
reconstructions=reconstructions,
|
||||
quant_loss=quant_loss,
|
||||
quant_info=quant_info,
|
||||
)
|
||||
145
tools/decode/vendor/cosmos_tokenizer/networks/discrete_video.py
vendored
Normal file
145
tools/decode/vendor/cosmos_tokenizer/networks/discrete_video.py
vendored
Normal file
@@ -0,0 +1,145 @@
|
||||
# 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 network definition for discrete video tokenizer with VQ, LFQ, FSQ or ResidualFSQ. """
|
||||
from collections import OrderedDict, namedtuple
|
||||
|
||||
import torch
|
||||
from loguru import logger as logging
|
||||
from torch import nn
|
||||
|
||||
from cosmos_tokenizer.modules import (
|
||||
Decoder3DType,
|
||||
DiscreteQuantizer,
|
||||
Encoder3DType,
|
||||
)
|
||||
from cosmos_tokenizer.modules.layers3d import CausalConv3d
|
||||
from cosmos_tokenizer.modules.quantizers import InvQuantizerJit
|
||||
|
||||
NetworkEval = namedtuple("NetworkEval", ["reconstructions", "quant_loss", "quant_info"])
|
||||
|
||||
|
||||
class CausalDiscreteVideoTokenizer(nn.Module):
|
||||
def __init__(
|
||||
self, z_channels: int, z_factor: int, embedding_dim: int, **kwargs
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.name = kwargs.get("name", "CausalDiscreteVideoTokenizer")
|
||||
self.embedding_dim = embedding_dim
|
||||
|
||||
encoder_name = kwargs.get("encoder", Encoder3DType.BASE.name)
|
||||
self.encoder = Encoder3DType[encoder_name].value(
|
||||
z_channels=z_factor * z_channels, **kwargs
|
||||
)
|
||||
|
||||
decoder_name = kwargs.get("decoder", Decoder3DType.BASE.name)
|
||||
self.decoder = Decoder3DType[decoder_name].value(
|
||||
z_channels=z_channels, **kwargs
|
||||
)
|
||||
|
||||
self.quant_conv = CausalConv3d(
|
||||
z_factor * z_channels, embedding_dim, kernel_size=1, padding=0
|
||||
)
|
||||
self.post_quant_conv = CausalConv3d(
|
||||
embedding_dim, z_channels, kernel_size=1, padding=0
|
||||
)
|
||||
|
||||
quantizer_name = kwargs.get("quantizer", DiscreteQuantizer.RESFSQ.name)
|
||||
if quantizer_name == DiscreteQuantizer.VQ.name:
|
||||
assert (
|
||||
"num_embeddings" in kwargs
|
||||
), f"`num_embeddings` must be provided for {quantizer_name}."
|
||||
kwargs.update(dict(embedding_dim=embedding_dim))
|
||||
elif quantizer_name == DiscreteQuantizer.LFQ.name:
|
||||
assert (
|
||||
"codebook_size" in kwargs
|
||||
), f"`codebook_size` must be provided for {quantizer_name}."
|
||||
assert (
|
||||
"codebook_dim" in kwargs
|
||||
), f"`codebook_dim` must be provided for {quantizer_name}."
|
||||
elif quantizer_name == DiscreteQuantizer.FSQ.name:
|
||||
assert (
|
||||
"levels" in kwargs
|
||||
), f"`levels` must be provided for {quantizer_name}."
|
||||
elif quantizer_name == DiscreteQuantizer.RESFSQ.name:
|
||||
assert (
|
||||
"levels" in kwargs
|
||||
), f"`levels` must be provided for {quantizer_name}."
|
||||
assert (
|
||||
"num_quantizers" in kwargs
|
||||
), f"`num_quantizers` must be provided for {quantizer_name}."
|
||||
self.quantizer = DiscreteQuantizer[quantizer_name].value(**kwargs)
|
||||
logging.info(f"{self.name} based on {quantizer_name}-VAE, with {kwargs}.")
|
||||
|
||||
num_parameters = sum(param.numel() for param in self.parameters())
|
||||
logging.info(f"model={self.name}, num_parameters={num_parameters:,}")
|
||||
logging.info(f"z_channels={z_channels}, embedding_dim={self.embedding_dim}.")
|
||||
|
||||
def to(self, *args, **kwargs):
|
||||
setattr(self.quantizer, "dtype", kwargs.get("dtype", torch.bfloat16))
|
||||
return super(CausalDiscreteVideoTokenizer, self).to(*args, **kwargs)
|
||||
|
||||
def encoder_jit(self):
|
||||
return nn.Sequential(
|
||||
OrderedDict(
|
||||
[
|
||||
("encoder", self.encoder),
|
||||
("quant_conv", self.quant_conv),
|
||||
("quantizer", self.quantizer),
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
def decoder_jit(self):
|
||||
return nn.Sequential(
|
||||
OrderedDict(
|
||||
[
|
||||
("inv_quant", InvQuantizerJit(self.quantizer)),
|
||||
("post_quant_conv", self.post_quant_conv),
|
||||
("decoder", self.decoder),
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
def last_decoder_layer(self):
|
||||
return self.decoder.conv_out
|
||||
|
||||
def encode(self, x):
|
||||
h = self.encoder(x)
|
||||
h = self.quant_conv(h)
|
||||
return self.quantizer(h)
|
||||
|
||||
def decode(self, quant):
|
||||
quant = self.post_quant_conv(quant)
|
||||
return self.decoder(quant)
|
||||
|
||||
def decode_code(self, code_b):
|
||||
quant_b = self.quantizer.indices_to_codes(code_b)
|
||||
quant_b = self.post_quant_conv(quant_b)
|
||||
return self.decoder(quant_b)
|
||||
|
||||
def forward(self, input):
|
||||
quant_info, quant_codes, quant_loss = self.encode(input)
|
||||
reconstructions = self.decode(quant_codes)
|
||||
if self.training:
|
||||
return dict(
|
||||
reconstructions=reconstructions,
|
||||
quant_loss=quant_loss,
|
||||
quant_info=quant_info,
|
||||
)
|
||||
return NetworkEval(
|
||||
reconstructions=reconstructions,
|
||||
quant_loss=quant_loss,
|
||||
quant_info=quant_info,
|
||||
)
|
||||
Reference in New Issue
Block a user