init
This commit is contained in:
@@ -0,0 +1,33 @@
|
||||
"""
|
||||
-----------------------------------------------------------------------------
|
||||
Copyright (c) 2024, NVIDIA CORPORATION. All rights reserved.
|
||||
|
||||
NVIDIA CORPORATION and its licensors retain all intellectual property
|
||||
and proprietary rights in and to this software, related documentation
|
||||
and any modifications thereto. Any use, reproduction, disclosure or
|
||||
distribution of this software and related documentation without an express
|
||||
license agreement from NVIDIA CORPORATION is strictly prohibited.
|
||||
-----------------------------------------------------------------------------
|
||||
"""
|
||||
|
||||
from .schema import ModelConfig
|
||||
|
||||
|
||||
def make_config():
|
||||
|
||||
model_config = ModelConfig(
|
||||
vae_conf="vae.configs.part_woenc",
|
||||
vae_ckpt_path="pretrained/vae.pt",
|
||||
qknorm=True,
|
||||
qknorm_type="RMSNorm",
|
||||
use_pos_embed=False,
|
||||
dino_model="dinov2_vitg14",
|
||||
hidden_dim=1536,
|
||||
flow_shift=3.0,
|
||||
logitnorm_mean=1.0,
|
||||
logitnorm_std=1.0,
|
||||
latent_size=4096,
|
||||
use_parts=True,
|
||||
)
|
||||
|
||||
return model_config
|
||||
@@ -0,0 +1,57 @@
|
||||
"""
|
||||
-----------------------------------------------------------------------------
|
||||
Copyright (c) 2024, NVIDIA CORPORATION. All rights reserved.
|
||||
|
||||
NVIDIA CORPORATION and its licensors retain all intellectual property
|
||||
and proprietary rights in and to this software, related documentation
|
||||
and any modifications thereto. Any use, reproduction, disclosure or
|
||||
distribution of this software and related documentation without an express
|
||||
license agreement from NVIDIA CORPORATION is strictly prohibited.
|
||||
-----------------------------------------------------------------------------
|
||||
"""
|
||||
|
||||
from typing import Literal, Optional
|
||||
|
||||
import attrs
|
||||
|
||||
|
||||
@attrs.define(slots=False)
|
||||
class ModelConfig:
|
||||
# vae
|
||||
vae_conf: str = "vae.configs.part_woenc"
|
||||
vae_ckpt_path: Optional[str] = None
|
||||
|
||||
# learn & generate parts
|
||||
use_parts: bool = False
|
||||
part_embed_mode: Literal["element", "part", "part2_only"] = "part2_only"
|
||||
shuffle_parts: bool = False
|
||||
use_num_parts_cond: bool = False
|
||||
|
||||
# flow matching hyper-params
|
||||
flow_shift: float = 1.0
|
||||
logitnorm_mean: float = 0.0
|
||||
logitnorm_std: float = 1.0
|
||||
|
||||
# image encoder
|
||||
dino_model: Literal["dinov2_vitl14_reg", "dinov2_vitg14"] = "dinov2_vitg14"
|
||||
|
||||
# backbone DiT
|
||||
hidden_dim: int = 1536
|
||||
num_heads: int = 16
|
||||
num_layers: int = 24
|
||||
qknorm: bool = True
|
||||
qknorm_type: Literal["LayerNorm", "RMSNorm"] = "RMSNorm"
|
||||
use_pos_embed: bool = False
|
||||
|
||||
# latent code
|
||||
latent_size: Optional[int] = None # if None, will load from vae
|
||||
latent_dim: Optional[int] = None
|
||||
|
||||
# preload vae weights
|
||||
preload_vae: bool = True
|
||||
|
||||
# preload dinov2 weights
|
||||
preload_dinov2: bool = True
|
||||
|
||||
# init weights from a pretrained checkpoint
|
||||
pretrain_path: Optional[str] = None
|
||||
@@ -0,0 +1,58 @@
|
||||
"""
|
||||
-----------------------------------------------------------------------------
|
||||
Copyright (c) 2024, NVIDIA CORPORATION. All rights reserved.
|
||||
|
||||
NVIDIA CORPORATION and its licensors retain all intellectual property
|
||||
and proprietary rights in and to this software, related documentation
|
||||
and any modifications thereto. Any use, reproduction, disclosure or
|
||||
distribution of this software and related documentation without an express
|
||||
license agreement from NVIDIA CORPORATION is strictly prohibited.
|
||||
-----------------------------------------------------------------------------
|
||||
"""
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
|
||||
class FlowMatchingScheduler:
|
||||
def __init__(self, num_train_timesteps: int = 1000, shift: float = 1):
|
||||
# set timesteps
|
||||
self.num_train_timesteps = num_train_timesteps
|
||||
self.shift = shift
|
||||
|
||||
timesteps = np.linspace(1, num_train_timesteps, num_train_timesteps, dtype=np.float32)[::-1].copy()
|
||||
timesteps = torch.from_numpy(timesteps).to(dtype=torch.float32)
|
||||
|
||||
sigmas = timesteps / num_train_timesteps
|
||||
sigmas = shift * sigmas / (1 + (shift - 1) * sigmas)
|
||||
|
||||
self.sigmas = sigmas # 1 --> 0
|
||||
self.timesteps = sigmas * num_train_timesteps # num_train_timesteps --> 1
|
||||
|
||||
# set device
|
||||
def to(self, device):
|
||||
self.sigmas = self.sigmas.to(device=device)
|
||||
self.timesteps = self.timesteps.to(device=device)
|
||||
|
||||
# add random noise to latent during training
|
||||
def add_noise(self, latent: torch.Tensor, logit_mean: float = 1.0, logit_std: float = 1.0):
|
||||
# latent: [B, ...]
|
||||
# timesteps: [B]
|
||||
# return: [B, ...] noisy_latent, [B, ...] noise, [B] timesteps
|
||||
|
||||
# logit-normal sampling
|
||||
u = torch.normal(mean=logit_mean, std=logit_std, size=(latent.shape[0],), device=self.sigmas.device)
|
||||
u = torch.nn.functional.sigmoid(u)
|
||||
|
||||
step_indices = (u * self.num_train_timesteps).long()
|
||||
timesteps = self.timesteps[step_indices]
|
||||
|
||||
sigmas = self.sigmas[step_indices].flatten()
|
||||
|
||||
while len(sigmas.shape) < latent.ndim:
|
||||
sigmas = sigmas.unsqueeze(-1)
|
||||
|
||||
noise = torch.randn_like(latent)
|
||||
noisy_latent = (1.0 - sigmas) * latent + sigmas * noise
|
||||
|
||||
return noisy_latent, noise, timesteps
|
||||
@@ -0,0 +1,343 @@
|
||||
"""
|
||||
-----------------------------------------------------------------------------
|
||||
Copyright (c) 2024, NVIDIA CORPORATION. All rights reserved.
|
||||
|
||||
NVIDIA CORPORATION and its licensors retain all intellectual property
|
||||
and proprietary rights in and to this software, related documentation
|
||||
and any modifications thereto. Any use, reproduction, disclosure or
|
||||
distribution of this software and related documentation without an express
|
||||
license agreement from NVIDIA CORPORATION is strictly prohibited.
|
||||
-----------------------------------------------------------------------------
|
||||
"""
|
||||
|
||||
import importlib
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import tqdm
|
||||
from torchvision import transforms
|
||||
from transformers import Dinov2Model
|
||||
|
||||
from .configs.schema import ModelConfig
|
||||
from .flow_matching import FlowMatchingScheduler
|
||||
from .modules.dit import DiT
|
||||
from ..vae.model import Model as VAE
|
||||
from ..vae.utils import sync_timer
|
||||
|
||||
|
||||
class Model(nn.Module):
|
||||
def __init__(self, config: ModelConfig,device,dino_path,cpu_offload=False) -> None:
|
||||
super().__init__()
|
||||
self.config = config
|
||||
self.precision = torch.bfloat16
|
||||
self.cpu_offload = cpu_offload
|
||||
# image condition model (dinov2)
|
||||
# if self.config.dino_model == "dinov2_vitg14":
|
||||
# #self.dino = Dinov2Model.from_pretrained("facebook/dinov2-giant")
|
||||
# elif self.config.dino_model == "dinov2_vitl14_reg":
|
||||
# #self.dino = Dinov2Model.from_pretrained("facebook/dinov2-with-registers-large")
|
||||
self.device = device
|
||||
if dino_path:
|
||||
self.dino = Dinov2Model.from_pretrained(dino_path)
|
||||
else:
|
||||
raise ValueError(f"DINOv2 model {self.config.dino_model} not supported")
|
||||
|
||||
# hack to match our implementation
|
||||
self.dino.layernorm = torch.nn.Identity()
|
||||
|
||||
self.dino.eval().to(dtype=self.precision)
|
||||
self.dino.requires_grad_(False)
|
||||
|
||||
cond_dim = 1024 if self.config.dino_model == "dinov2_vitl14_reg" else 1536
|
||||
assert cond_dim == config.hidden_dim, "DINOv2 dim must match backbone dim"
|
||||
|
||||
self.preprocess_cond_image = transforms.Compose(
|
||||
[
|
||||
transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
|
||||
]
|
||||
)
|
||||
|
||||
# vae encoder
|
||||
vae_config = importlib.import_module(config.vae_conf).make_config()
|
||||
self.vae = VAE(vae_config).eval().to(dtype=self.precision)
|
||||
self.vae.requires_grad_(False)
|
||||
|
||||
# load vae
|
||||
if self.config.preload_vae:
|
||||
try:
|
||||
vae_ckpt = torch.load(self.config.vae_ckpt_path, weights_only=True) # local path
|
||||
if "model" in vae_ckpt:
|
||||
vae_ckpt = vae_ckpt["model"]
|
||||
self.vae.load_state_dict(vae_ckpt, strict=True)
|
||||
del vae_ckpt
|
||||
print(f"Loaded VAE from {self.config.vae_ckpt_path}")
|
||||
except Exception as e:
|
||||
print(
|
||||
f"Failed to load VAE from {self.config.vae_ckpt_path}: {e}, make sure you resumed from a valid checkpoint!"
|
||||
)
|
||||
|
||||
# load info from vae config
|
||||
if config.latent_size is None:
|
||||
config.latent_size = self.vae.config.latent_size
|
||||
if config.latent_dim is None:
|
||||
config.latent_dim = self.vae.config.latent_dim
|
||||
|
||||
# dit
|
||||
self.dit = DiT(
|
||||
hidden_dim=config.hidden_dim,
|
||||
num_heads=config.num_heads,
|
||||
num_layers=config.num_layers,
|
||||
latent_size=config.latent_size,
|
||||
latent_dim=config.latent_dim,
|
||||
qknorm=config.qknorm,
|
||||
qknorm_type=config.qknorm_type,
|
||||
use_pos_embed=config.use_pos_embed,
|
||||
use_parts=config.use_parts,
|
||||
part_embed_mode=config.part_embed_mode,
|
||||
)
|
||||
|
||||
# num_part condition
|
||||
if self.config.use_num_parts_cond:
|
||||
assert self.config.use_parts, "use_num_parts_cond requires use_parts"
|
||||
self.num_part_embed = nn.Embedding(5, config.hidden_dim)
|
||||
|
||||
# preload from a checkpoint (NOTE: this happens BEFORE checkpointer loading latest checkpoint!)
|
||||
if self.config.pretrain_path is not None:
|
||||
try:
|
||||
ckpt = torch.load(self.config.pretrain_path) # local path
|
||||
self.load_state_dict(ckpt["model"], strict=True)
|
||||
del ckpt
|
||||
print(f"Loaded DiT from {self.config.pretrain_path}")
|
||||
except Exception as e:
|
||||
print(
|
||||
f"Failed to load DiT from {self.config.pretrain_path}: {e}, make sure you resumed from a valid checkpoint!"
|
||||
)
|
||||
|
||||
# sampler
|
||||
self.scheduler = FlowMatchingScheduler(shift=config.flow_shift)
|
||||
|
||||
n_params = 0
|
||||
for p in self.dit.parameters():
|
||||
n_params += p.numel()
|
||||
print(f"Number of parameters in DiT: {n_params/1e6:.2f}M")
|
||||
|
||||
# override state_dict to exclude vae and dino, so we only save the trainable params.
|
||||
def state_dict(self, *args, **kwargs):
|
||||
state_dict = super().state_dict(*args, **kwargs)
|
||||
|
||||
keys_to_del = []
|
||||
for k in state_dict.keys():
|
||||
if "vae" in k or "dino" in k:
|
||||
keys_to_del.append(k)
|
||||
|
||||
for k in keys_to_del:
|
||||
del state_dict[k]
|
||||
|
||||
return state_dict
|
||||
|
||||
# override to support tolerant loading (only load matched shape)
|
||||
def load_state_dict(self, state_dict, strict=True, assign=False):
|
||||
local_state_dict = self.state_dict()
|
||||
seen_keys = {k: False for k in local_state_dict.keys()}
|
||||
for k, v in state_dict.items():
|
||||
if k in local_state_dict:
|
||||
seen_keys[k] = True
|
||||
if local_state_dict[k].shape == v.shape:
|
||||
local_state_dict[k].copy_(v)
|
||||
else:
|
||||
print(f"mismatching shape for key {k}: loaded {local_state_dict[k].shape} but model has {v.shape}")
|
||||
else:
|
||||
print(f"unexpected key {k} in loaded state dict")
|
||||
for k in seen_keys:
|
||||
if not seen_keys[k]:
|
||||
print(f"missing key {k} in loaded state dict")
|
||||
|
||||
# this happens before checkpointer loading old models !!!
|
||||
def on_train_start(self, memory_format: torch.memory_format = torch.preserve_format) -> None:
|
||||
super().on_train_start(memory_format=memory_format)
|
||||
device = next(self.dit.parameters()).device
|
||||
|
||||
self.dit.to(dtype=self.precision)
|
||||
|
||||
if self.config.use_num_parts_cond:
|
||||
self.num_part_embed.to(dtype=self.precision)
|
||||
|
||||
# cast scheduler to device
|
||||
self.scheduler.to(device)
|
||||
|
||||
def get_cond(self, cond_image, num_part=None):
|
||||
# image condition
|
||||
cond_image = cond_image.to(dtype=self.precision)
|
||||
with torch.no_grad():
|
||||
cond = self.dino(cond_image).last_hidden_state
|
||||
cond = F.layer_norm(cond.float(), cond.shape[-1:]).to(dtype=self.precision) # [B, L, C]
|
||||
|
||||
# num_part condition
|
||||
if self.config.use_num_parts_cond:
|
||||
if num_part is None:
|
||||
# use a default value (2-10 parts)
|
||||
num_part_coarse = torch.ones(cond.shape[0], dtype=torch.int64, device=cond.device) * 2
|
||||
else:
|
||||
# coarse range
|
||||
num_part_coarse = torch.ones(cond.shape[0], dtype=torch.int64, device=cond.device)
|
||||
num_part_coarse[num_part == 2] = 1
|
||||
num_part_coarse[(num_part > 2) & (num_part <= 10)] = 2
|
||||
num_part_coarse[(num_part > 10) & (num_part <= 100)] = 3
|
||||
num_part_coarse[num_part > 100] = 4
|
||||
num_part_embed = self.num_part_embed(num_part_coarse).unsqueeze(1) # [B, 1, C]
|
||||
cond = torch.cat([cond, num_part_embed], dim=1) # [B, L+1, C]
|
||||
|
||||
return cond
|
||||
|
||||
def training_step(
|
||||
self,
|
||||
data: dict[str, torch.Tensor],
|
||||
iteration: int,
|
||||
) -> tuple[dict[str, torch.Tensor], torch.Tensor]:
|
||||
output = {}
|
||||
loss = 0
|
||||
|
||||
cond_images = self.preprocess_cond_image(
|
||||
data["cond_images"]
|
||||
) # [B, N, 3, 518, 518], we may load multiple (N) cond images for the same shape
|
||||
B, N, C, H, W = cond_images.shape
|
||||
|
||||
if self.config.use_num_parts_cond:
|
||||
cond_num_part = data["num_part"].repeat_interleave(N, dim=0)
|
||||
else:
|
||||
cond_num_part = None
|
||||
|
||||
cond = self.get_cond(cond_images.view(-1, C, H, W), cond_num_part) # [B*N, L, C]
|
||||
|
||||
# random CFG dropout
|
||||
if self.training:
|
||||
mask = torch.rand((B * N, 1, 1), device=cond.device, dtype=cond.dtype) >= 0.1
|
||||
cond = cond * mask
|
||||
|
||||
with torch.no_grad():
|
||||
# encode latent
|
||||
if self.config.use_parts:
|
||||
# encode two parts and concat latent
|
||||
part0_data = {k.replace("_part0", ""): v for k, v in data.items() if "_part0" in k}
|
||||
part1_data = {k.replace("_part1", ""): v for k, v in data.items() if "_part1" in k}
|
||||
posterior0 = self.vae.encode(part0_data)
|
||||
posterior1 = self.vae.encode(part1_data)
|
||||
if self.training and self.config.shuffle_parts:
|
||||
if np.random.rand() < 0.5:
|
||||
posterior0, posterior1 = posterior1, posterior0
|
||||
latent = torch.cat(
|
||||
[
|
||||
posterior0.mode().float().nan_to_num_(0),
|
||||
posterior1.mode().float().nan_to_num_(0),
|
||||
],
|
||||
dim=1,
|
||||
) # [B, 2L, C]
|
||||
else:
|
||||
posterior = self.vae.encode(data)
|
||||
latent = posterior.mode().float().nan_to_num_(0) # use mean as the latent, [B, L, C]
|
||||
|
||||
# repeat latent for each cond image
|
||||
if N != 1:
|
||||
latent = latent.repeat_interleave(N, dim=0)
|
||||
|
||||
# random sample timesteps and add noise
|
||||
noisy_latent, noise, timesteps = self.scheduler.add_noise(
|
||||
latent, self.config.logitnorm_mean, self.config.logitnorm_std
|
||||
)
|
||||
|
||||
noisy_latent = noisy_latent.to(dtype=self.precision)
|
||||
model_pred = self.dit(noisy_latent, cond, timesteps)
|
||||
|
||||
# flow-matching loss
|
||||
target = noise - latent
|
||||
loss = F.mse_loss(model_pred.float(), target.float())
|
||||
|
||||
# metrics
|
||||
with torch.no_grad():
|
||||
output["scalar"] = {} # for wandb logging
|
||||
output["scalar"]["loss_mse"] = loss.detach()
|
||||
|
||||
return output, loss
|
||||
|
||||
@torch.no_grad()
|
||||
def validation_step(
|
||||
self,
|
||||
data: dict[str, torch.Tensor],
|
||||
iteration: int,
|
||||
) -> tuple[dict[str, torch.Tensor], torch.Tensor]:
|
||||
return self.training_step(data, iteration)
|
||||
|
||||
@torch.inference_mode()
|
||||
@sync_timer("flow forward")
|
||||
def forward(
|
||||
self,
|
||||
data: dict[str, torch.Tensor],
|
||||
num_steps: int = 30,
|
||||
cfg_scale: float = 7.0,
|
||||
verbose: bool = True,
|
||||
generator: torch.Generator | None = None,
|
||||
) -> dict[str, torch.Tensor]:
|
||||
# the inference sampling
|
||||
cond_images = self.preprocess_cond_image(data["cond_images"]) # [B, 3, 518, 518]
|
||||
B = cond_images.shape[0]
|
||||
assert B == 1, "Only support batch size 1 for now."
|
||||
|
||||
# num_part condition
|
||||
if self.config.use_num_parts_cond and "num_part" in data:
|
||||
cond_num_part = data["num_part"] # [B,], int
|
||||
else:
|
||||
cond_num_part = None
|
||||
if self.cpu_offload:
|
||||
self.dino.to(device=self.device)
|
||||
cond = self.get_cond(cond_images, cond_num_part)
|
||||
if self.cpu_offload:
|
||||
self.dino.to(device="cpu")
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
if self.config.use_parts:
|
||||
x = torch.randn(
|
||||
B,
|
||||
self.config.latent_size * 2,
|
||||
self.config.latent_dim,
|
||||
device=cond.device,
|
||||
dtype=torch.float32,
|
||||
generator=generator,
|
||||
)
|
||||
else:
|
||||
x = torch.randn(
|
||||
B,
|
||||
self.config.latent_size,
|
||||
self.config.latent_dim,
|
||||
device=cond.device,
|
||||
dtype=torch.float32,
|
||||
generator=generator,
|
||||
)
|
||||
|
||||
cond_input = torch.cat([cond, torch.zeros_like(cond)], dim=0)
|
||||
|
||||
# flow-matching
|
||||
sigmas = np.linspace(1, 0, num_steps + 1)
|
||||
sigmas = self.scheduler.shift * sigmas / (1 + (self.scheduler.shift - 1) * sigmas)
|
||||
sigmas_pair = list((sigmas[i], sigmas[i + 1]) for i in range(num_steps))
|
||||
|
||||
for sigma, sigma_prev in tqdm.tqdm(sigmas_pair, desc="Flow Sampling", disable=not verbose):
|
||||
# classifier-free guidance
|
||||
timesteps = torch.tensor([1000 * sigma] * B * 2, device=x.device, dtype=x.dtype)
|
||||
x_input = torch.cat([x, x], dim=0)
|
||||
|
||||
# predict v
|
||||
x_input = x_input.to(dtype=self.precision)
|
||||
pred = self.dit(x_input, cond_input, timesteps).float()
|
||||
cond_v, uncond_v = pred.chunk(2, dim=0)
|
||||
pred_v = uncond_v + (cond_v - uncond_v) * cfg_scale
|
||||
|
||||
# scheduler step
|
||||
x = x - (sigma - sigma_prev) * pred_v
|
||||
|
||||
output = {}
|
||||
output["latent"] = x
|
||||
|
||||
# leave mesh extraction to vae
|
||||
return output
|
||||
@@ -0,0 +1,235 @@
|
||||
"""
|
||||
-----------------------------------------------------------------------------
|
||||
Copyright (c) 2024, NVIDIA CORPORATION. All rights reserved.
|
||||
|
||||
NVIDIA CORPORATION and its licensors retain all intellectual property
|
||||
and proprietary rights in and to this software, related documentation
|
||||
and any modifications thereto. Any use, reproduction, disclosure or
|
||||
distribution of this software and related documentation without an express
|
||||
license agreement from NVIDIA CORPORATION is strictly prohibited.
|
||||
-----------------------------------------------------------------------------
|
||||
"""
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from torch.utils.checkpoint import checkpoint
|
||||
|
||||
from ...vae.modules.attention import CrossAttention, SelfAttention
|
||||
|
||||
|
||||
class FeedForward(nn.Module):
|
||||
def __init__(self, dim, mult=4):
|
||||
super().__init__()
|
||||
self.net = nn.Sequential(nn.Linear(dim, dim * mult), nn.GELU(), nn.Linear(dim * mult, dim))
|
||||
|
||||
def forward(self, x):
|
||||
return self.net(x)
|
||||
|
||||
|
||||
# Adapted from https://github.com/facebookresearch/DiT/blob/main/models.py#L27
|
||||
class TimestepEmbedder(nn.Module):
|
||||
"""
|
||||
Embeds scalar timesteps into vector representations.
|
||||
"""
|
||||
|
||||
def __init__(self, hidden_size, frequency_embedding_size=256):
|
||||
super().__init__()
|
||||
self.mlp = nn.Sequential(
|
||||
nn.Linear(frequency_embedding_size, hidden_size, bias=True),
|
||||
nn.SiLU(),
|
||||
nn.Linear(hidden_size, hidden_size, bias=True),
|
||||
)
|
||||
self.frequency_embedding_size = frequency_embedding_size
|
||||
|
||||
@staticmethod
|
||||
def timestep_embedding(t, dim, max_period=10000):
|
||||
"""
|
||||
Create sinusoidal timestep embeddings.
|
||||
|
||||
Args:
|
||||
t: a 1-D Tensor of N indices, one per batch element.
|
||||
These may be fractional.
|
||||
dim: the dimension of the output.
|
||||
max_period: controls the minimum frequency of the embeddings.
|
||||
|
||||
Returns:
|
||||
an (N, D) Tensor of positional embeddings.
|
||||
"""
|
||||
# https://github.com/openai/glide-text2im/blob/main/glide_text2im/nn.py
|
||||
half = dim // 2
|
||||
freqs = torch.exp(-np.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32) / half).to(
|
||||
device=t.device
|
||||
)
|
||||
args = t[:, None].float() * freqs[None]
|
||||
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
|
||||
if dim % 2:
|
||||
embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
|
||||
return embedding
|
||||
|
||||
def forward(self, t):
|
||||
dtype = next(self.mlp.parameters()).dtype # need to determine on the fly...
|
||||
t_freq = self.timestep_embedding(t, self.frequency_embedding_size)
|
||||
t_freq = t_freq.to(dtype=dtype)
|
||||
t_emb = self.mlp(t_freq)
|
||||
return t_emb
|
||||
|
||||
|
||||
class DiTLayer(nn.Module):
|
||||
def __init__(self, dim, num_heads, qknorm=False, gradient_checkpointing=True, qknorm_type="LayerNorm"):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.num_heads = num_heads
|
||||
self.gradient_checkpointing = gradient_checkpointing
|
||||
|
||||
self.norm1 = nn.LayerNorm(dim, eps=1e-6, elementwise_affine=False)
|
||||
self.attn1 = SelfAttention(dim, num_heads, qknorm=qknorm, qknorm_type=qknorm_type)
|
||||
self.norm2 = nn.LayerNorm(dim, eps=1e-6, elementwise_affine=False)
|
||||
self.attn2 = CrossAttention(dim, num_heads, context_dim=dim, qknorm=qknorm, qknorm_type=qknorm_type)
|
||||
self.norm3 = nn.LayerNorm(dim, eps=1e-6, elementwise_affine=False)
|
||||
self.ff = FeedForward(dim)
|
||||
self.adaln_linear = nn.Linear(dim, dim * 6, bias=True)
|
||||
|
||||
def forward(self, x, c, t_emb):
|
||||
if self.training and self.gradient_checkpointing:
|
||||
return checkpoint(self._forward, x, c, t_emb, use_reentrant=False)
|
||||
else:
|
||||
return self._forward(x, c, t_emb)
|
||||
|
||||
def _forward(self, x, c, t_emb):
|
||||
# x: [B, N, C], hidden states
|
||||
# c: [B, M, C], condition (assume normed and projected to C)
|
||||
# t_emb: [B, C], timestep embedding of adaln
|
||||
# return: [B, N, C], updated hidden states
|
||||
|
||||
B, N, C = x.shape
|
||||
t_adaln = self.adaln_linear(F.silu(t_emb)).view(B, 6, -1) # [B, 6, C]
|
||||
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = t_adaln.chunk(6, dim=1)
|
||||
|
||||
h = self.norm1(x)
|
||||
h = h * (1 + scale_msa) + shift_msa
|
||||
x = x + gate_msa * self.attn1(h)
|
||||
|
||||
h = self.norm2(x)
|
||||
x = x + self.attn2(h, c)
|
||||
|
||||
h = self.norm3(x)
|
||||
h = h * (1 + scale_mlp) + shift_mlp
|
||||
x = x + gate_mlp * self.ff(h)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class DiT(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
hidden_dim=1024,
|
||||
num_heads=16,
|
||||
latent_size=2048,
|
||||
latent_dim=8,
|
||||
num_layers=24,
|
||||
qknorm=False,
|
||||
gradient_checkpointing=True,
|
||||
qknorm_type="LayerNorm",
|
||||
use_pos_embed=False,
|
||||
use_parts=False,
|
||||
part_embed_mode="part2_only",
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
# project in
|
||||
self.proj_in = nn.Linear(latent_dim, hidden_dim)
|
||||
|
||||
# positional encoding (just use a learnable positional encoding)
|
||||
self.use_pos_embed = use_pos_embed
|
||||
if self.use_pos_embed:
|
||||
self.pos_embed = nn.Parameter(torch.randn(1, latent_size, hidden_dim) / hidden_dim**0.5)
|
||||
|
||||
# part encoding (a must to distinguish parts!)
|
||||
self.use_parts = use_parts
|
||||
self.part_embed_mode = part_embed_mode
|
||||
if self.use_parts:
|
||||
if self.part_embed_mode == "element":
|
||||
self.part_embed = nn.Parameter(torch.randn(latent_size, hidden_dim) / hidden_dim**0.5)
|
||||
elif self.part_embed_mode == "part":
|
||||
self.part_embed = nn.Parameter(torch.randn(2, hidden_dim))
|
||||
elif self.part_embed_mode == "part2_only":
|
||||
# we only add this to the second part to distinguish from the first part
|
||||
self.part_embed = nn.Parameter(torch.randn(1, hidden_dim) / hidden_dim**0.5)
|
||||
|
||||
# timestep encoding
|
||||
self.timestep_embed = TimestepEmbedder(hidden_dim)
|
||||
|
||||
# transformer layers
|
||||
self.layers = nn.ModuleList(
|
||||
[DiTLayer(hidden_dim, num_heads, qknorm, gradient_checkpointing, qknorm_type) for _ in range(num_layers)]
|
||||
)
|
||||
|
||||
# project out
|
||||
self.norm_out = nn.LayerNorm(hidden_dim, eps=1e-6, elementwise_affine=False)
|
||||
self.proj_out = nn.Linear(hidden_dim, latent_dim)
|
||||
|
||||
# init
|
||||
self.init_weight()
|
||||
|
||||
def init_weight(self):
|
||||
# Initialize transformer layers
|
||||
def _basic_init(module):
|
||||
if isinstance(module, nn.Linear):
|
||||
torch.nn.init.xavier_uniform_(module.weight)
|
||||
if module.bias is not None:
|
||||
nn.init.constant_(module.bias, 0)
|
||||
|
||||
self.apply(_basic_init)
|
||||
|
||||
# Initialize timestep embedding MLP:
|
||||
nn.init.normal_(self.timestep_embed.mlp[0].weight, std=0.02)
|
||||
nn.init.normal_(self.timestep_embed.mlp[2].weight, std=0.02)
|
||||
|
||||
# Zero-out adaLN modulation layers in DiT blocks:
|
||||
for layer in self.layers:
|
||||
nn.init.constant_(layer.adaln_linear.weight, 0)
|
||||
nn.init.constant_(layer.adaln_linear.bias, 0)
|
||||
|
||||
# Zero-out output layers:
|
||||
nn.init.constant_(self.proj_out.weight, 0)
|
||||
nn.init.constant_(self.proj_out.bias, 0)
|
||||
|
||||
def forward(self, x, c, t):
|
||||
# x: [B, N, C], hidden states
|
||||
# c: [B, M, C], condition (assume normed and projected to C)
|
||||
# t: [B,], timestep
|
||||
# return: [B, N, C], updated hidden states
|
||||
|
||||
B, N, C = x.shape
|
||||
|
||||
# project in
|
||||
x = self.proj_in(x)
|
||||
|
||||
# positional encoding
|
||||
if self.use_pos_embed:
|
||||
x = x + self.pos_embed
|
||||
|
||||
# part encoding
|
||||
if self.use_parts:
|
||||
if self.part_embed_mode == "element":
|
||||
x += self.part_embed
|
||||
elif self.part_embed_mode == "part":
|
||||
x[:, : x.shape[1] // 2, :] += self.part_embed[0]
|
||||
x[:, x.shape[1] // 2 :, :] += self.part_embed[1]
|
||||
elif self.part_embed_mode == "part2_only":
|
||||
x[:, x.shape[1] // 2 :, :] += self.part_embed[0]
|
||||
|
||||
# timestep encoding
|
||||
t_emb = self.timestep_embed(t) # [B, C]
|
||||
|
||||
# transformer layers
|
||||
for layer in self.layers:
|
||||
x = layer(x, c, t_emb)
|
||||
|
||||
# project out
|
||||
x = self.norm_out(x)
|
||||
x = self.proj_out(x)
|
||||
|
||||
return x
|
||||
@@ -0,0 +1,184 @@
|
||||
"""
|
||||
-----------------------------------------------------------------------------
|
||||
Copyright (c) 2024, NVIDIA CORPORATION. All rights reserved.
|
||||
|
||||
NVIDIA CORPORATION and its licensors retain all intellectual property
|
||||
and proprietary rights in and to this software, related documentation
|
||||
and any modifications thereto. Any use, reproduction, disclosure or
|
||||
distribution of this software and related documentation without an express
|
||||
license agreement from NVIDIA CORPORATION is strictly prohibited.
|
||||
-----------------------------------------------------------------------------
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import glob
|
||||
import importlib
|
||||
import os
|
||||
from datetime import datetime
|
||||
|
||||
import cv2
|
||||
import kiui
|
||||
import numpy as np
|
||||
import rembg
|
||||
import torch
|
||||
import trimesh
|
||||
|
||||
from ..model import Model
|
||||
from ..utils import get_random_color, recenter_foreground
|
||||
from ...vae.utils import postprocess_mesh
|
||||
|
||||
# PYTHONPATH=. python flow/scripts/infer.py
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument(
|
||||
"--config",
|
||||
type=str,
|
||||
help="config file path",
|
||||
default="flow.configs.big_parts_strict_pvae",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ckpt_path",
|
||||
type=str,
|
||||
help="checkpoint path",
|
||||
default="pretrained/flow.pt",
|
||||
)
|
||||
parser.add_argument("--input", type=str, help="input directory", default="assets/images/")
|
||||
parser.add_argument("--limit", type=int, help="limit number of images", default=-1)
|
||||
parser.add_argument("--output_dir", type=str, help="output directory", default="output/")
|
||||
parser.add_argument("--grid_res", type=int, help="grid resolution", default=384)
|
||||
parser.add_argument("--num_steps", type=int, help="number of cfg steps", default=50)
|
||||
parser.add_argument("--cfg_scale", type=float, help="cfg scale", default=7.0)
|
||||
parser.add_argument("--num_repeats", type=int, help="number of repeats per image", default=1)
|
||||
parser.add_argument("--num_faces", type=int, help="target number of faces for decimation", default=-1)
|
||||
parser.add_argument("--seed", type=int, help="seed", default=42)
|
||||
args = parser.parse_args()
|
||||
|
||||
TRIMESH_GLB_EXPORT = np.array([[0, 1, 0], [0, 0, 1], [1, 0, 0]]).astype(np.float32)
|
||||
|
||||
bg_remover = rembg.new_session()
|
||||
|
||||
|
||||
def preprocess_image(path):
|
||||
input_image = kiui.read_image(path, mode="uint8", order="RGBA")
|
||||
|
||||
# bg removal if there is no alpha channel
|
||||
if input_image.shape[-1] == 3:
|
||||
input_image = rembg.remove(input_image, session=bg_remover) # [H, W, 4]
|
||||
|
||||
mask = input_image[..., -1] > 0
|
||||
image = recenter_foreground(input_image, mask, border_ratio=0.1)
|
||||
image = cv2.resize(image, (518, 518), interpolation=cv2.INTER_LINEAR)
|
||||
image = image.astype(np.float32) / 255.0
|
||||
image = image[..., :3] * image[..., 3:4] + (1 - image[..., 3:4]) # white background
|
||||
return image
|
||||
|
||||
|
||||
print(f"Loading checkpoint from {args.ckpt_path}")
|
||||
ckpt_dict = torch.load(args.ckpt_path, weights_only=True)
|
||||
|
||||
# delete all keys other than model
|
||||
if "model" in ckpt_dict:
|
||||
ckpt_dict = ckpt_dict["model"]
|
||||
|
||||
# instantiate model
|
||||
print(f"Instantiating model from {args.config}")
|
||||
model_config = importlib.import_module(args.config).make_config()
|
||||
model = Model(model_config).eval().cuda().bfloat16()
|
||||
|
||||
# load weight
|
||||
print(f"Loading weights from {args.ckpt_path}")
|
||||
model.load_state_dict(ckpt_dict, strict=True)
|
||||
|
||||
# output folder
|
||||
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
workspace = os.path.join(args.output_dir, "flow_" + args.config.split(".")[-1] + "_" + timestamp)
|
||||
if not os.path.exists(workspace):
|
||||
os.makedirs(workspace)
|
||||
else:
|
||||
os.system(f"rm {workspace}/*")
|
||||
print(f"Output directory: {workspace}")
|
||||
|
||||
# load test images
|
||||
if os.path.isdir(args.input):
|
||||
paths = glob.glob(os.path.join(args.input, "*"))
|
||||
paths = sorted(paths)
|
||||
if args.limit > 0:
|
||||
paths = paths[: args.limit]
|
||||
else: # single file
|
||||
paths = [args.input]
|
||||
|
||||
for path in paths:
|
||||
name = os.path.splitext(os.path.basename(path))[0]
|
||||
print(f"Processing {name}")
|
||||
|
||||
image = preprocess_image(path)
|
||||
|
||||
kiui.write_image(os.path.join(workspace, name + ".jpg"), image)
|
||||
image = torch.from_numpy(image).permute(2, 0, 1).contiguous().unsqueeze(0).float().cuda()
|
||||
|
||||
# run model
|
||||
data = {"cond_images": image}
|
||||
|
||||
for i in range(args.num_repeats):
|
||||
|
||||
kiui.seed_everything(args.seed + i)
|
||||
|
||||
with torch.inference_mode():
|
||||
results = model(data, num_steps=args.num_steps, cfg_scale=args.cfg_scale)
|
||||
|
||||
latent = results["latent"]
|
||||
# kiui.lo(latent)
|
||||
|
||||
# query mesh
|
||||
if model.config.use_parts:
|
||||
data_part0 = {"latent": latent[:, : model.config.latent_size, :]}
|
||||
data_part1 = {"latent": latent[:, model.config.latent_size :, :]}
|
||||
|
||||
with torch.inference_mode():
|
||||
results_part0 = model.vae(data_part0, resolution=args.grid_res)
|
||||
results_part1 = model.vae(data_part1, resolution=args.grid_res)
|
||||
|
||||
vertices, faces = results_part0["meshes"][0]
|
||||
mesh_part0 = trimesh.Trimesh(vertices, faces)
|
||||
mesh_part0.vertices = mesh_part0.vertices @ TRIMESH_GLB_EXPORT.T
|
||||
mesh_part0 = postprocess_mesh(mesh_part0, args.num_faces)
|
||||
parts = mesh_part0.split(only_watertight=False)
|
||||
|
||||
vertices, faces = results_part1["meshes"][0]
|
||||
mesh_part1 = trimesh.Trimesh(vertices, faces)
|
||||
mesh_part1.vertices = mesh_part1.vertices @ TRIMESH_GLB_EXPORT.T
|
||||
mesh_part1 = postprocess_mesh(mesh_part1, args.num_faces)
|
||||
parts.extend(mesh_part1.split(only_watertight=False))
|
||||
|
||||
# some parts only have 1 face, seems a problem of trimesh.split.
|
||||
parts = [part for part in parts if len(part.faces) > 10]
|
||||
|
||||
# split connected components and assign different colors
|
||||
for j, part in enumerate(parts):
|
||||
# each component uses a random color
|
||||
part.visual.vertex_colors = get_random_color(j, use_float=True)
|
||||
|
||||
mesh = trimesh.Scene(parts)
|
||||
# export the whole mesh
|
||||
mesh.export(os.path.join(workspace, name + "_" + str(i) + ".glb"))
|
||||
|
||||
# export each part
|
||||
for j, part in enumerate(parts):
|
||||
part.export(os.path.join(workspace, name + "_" + str(i) + "_part" + str(j) + ".glb"))
|
||||
|
||||
# export dual volumes
|
||||
mesh_part0.export(os.path.join(workspace, name + "_" + str(i) + "_vol0.glb"))
|
||||
mesh_part1.export(os.path.join(workspace, name + "_" + str(i) + "_vol1.glb"))
|
||||
|
||||
else:
|
||||
data = {"latent": latent}
|
||||
|
||||
with torch.inference_mode():
|
||||
results = model.vae(data, resolution=args.grid_res)
|
||||
|
||||
vertices, faces = results["meshes"][0]
|
||||
mesh = trimesh.Trimesh(vertices, faces)
|
||||
mesh = postprocess_mesh(mesh, args.num_faces)
|
||||
|
||||
# kiui.lo(mesh.vertices, mesh.faces)
|
||||
mesh.vertices = mesh.vertices @ TRIMESH_GLB_EXPORT.T
|
||||
mesh.export(os.path.join(workspace, name + "_" + str(i) + ".glb"))
|
||||
@@ -0,0 +1,119 @@
|
||||
"""
|
||||
-----------------------------------------------------------------------------
|
||||
Copyright (c) 2024, NVIDIA CORPORATION. All rights reserved.
|
||||
|
||||
NVIDIA CORPORATION and its licensors retain all intellectual property
|
||||
and proprietary rights in and to this software, related documentation
|
||||
and any modifications thereto. Any use, reproduction, disclosure or
|
||||
distribution of this software and related documentation without an express
|
||||
license agreement from NVIDIA CORPORATION is strictly prohibited.
|
||||
-----------------------------------------------------------------------------
|
||||
"""
|
||||
|
||||
from typing import Optional
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
|
||||
|
||||
def recenter_foreground(image, mask, border_ratio: float = 0.1):
|
||||
"""recenter an image to leave some empty space at the image border.
|
||||
|
||||
Args:
|
||||
image (ndarray): input image, float/uint8 [H, W, 3/4]
|
||||
mask (ndarray): alpha mask, bool [H, W]
|
||||
border_ratio (float, optional): border ratio, image will be resized to (1 - border_ratio). Defaults to 0.1.
|
||||
|
||||
Returns:
|
||||
ndarray: output image, float/uint8 [H, W, 3/4]
|
||||
"""
|
||||
|
||||
# empty foreground: just return
|
||||
if mask.sum() == 0:
|
||||
return image
|
||||
|
||||
return_int = False
|
||||
if image.dtype == np.uint8:
|
||||
image = image.astype(np.float32) / 255
|
||||
return_int = True
|
||||
|
||||
H, W, C = image.shape
|
||||
size = max(H, W)
|
||||
|
||||
# default to white bg if rgb, but use 0 if rgba
|
||||
if C == 3:
|
||||
result = np.ones((size, size, C), dtype=np.float32)
|
||||
else:
|
||||
result = np.zeros((size, size, C), dtype=np.float32)
|
||||
|
||||
coords = np.nonzero(mask)
|
||||
x_min, x_max = coords[0].min(), coords[0].max()
|
||||
y_min, y_max = coords[1].min(), coords[1].max()
|
||||
h = x_max - x_min
|
||||
w = y_max - y_min
|
||||
desired_size = int(size * (1 - border_ratio))
|
||||
scale = desired_size / max(h, w)
|
||||
h2 = int(h * scale)
|
||||
w2 = int(w * scale)
|
||||
x2_min = (size - h2) // 2
|
||||
x2_max = x2_min + h2
|
||||
y2_min = (size - w2) // 2
|
||||
y2_max = y2_min + w2
|
||||
result[x2_min:x2_max, y2_min:y2_max] = cv2.resize(
|
||||
image[x_min:x_max, y_min:y_max], (w2, h2), interpolation=cv2.INTER_AREA
|
||||
)
|
||||
|
||||
if return_int:
|
||||
result = (result * 255).astype(np.uint8)
|
||||
|
||||
return result
|
||||
|
||||
|
||||
def get_random_color(index: Optional[int] = None, use_float: bool = False):
|
||||
# some pleasing colors
|
||||
# matplotlib.colormaps['Set3'].colors + matplotlib.colormaps['Set2'].colors + matplotlib.colormaps['Set1'].colors
|
||||
palette = np.array(
|
||||
[
|
||||
[141, 211, 199, 255],
|
||||
[255, 255, 179, 255],
|
||||
[190, 186, 218, 255],
|
||||
[251, 128, 114, 255],
|
||||
[128, 177, 211, 255],
|
||||
[253, 180, 98, 255],
|
||||
[179, 222, 105, 255],
|
||||
[252, 205, 229, 255],
|
||||
[217, 217, 217, 255],
|
||||
[188, 128, 189, 255],
|
||||
[204, 235, 197, 255],
|
||||
[255, 237, 111, 255],
|
||||
[102, 194, 165, 255],
|
||||
[252, 141, 98, 255],
|
||||
[141, 160, 203, 255],
|
||||
[231, 138, 195, 255],
|
||||
[166, 216, 84, 255],
|
||||
[255, 217, 47, 255],
|
||||
[229, 196, 148, 255],
|
||||
[179, 179, 179, 255],
|
||||
[228, 26, 28, 255],
|
||||
[55, 126, 184, 255],
|
||||
[77, 175, 74, 255],
|
||||
[152, 78, 163, 255],
|
||||
[255, 127, 0, 255],
|
||||
[255, 255, 51, 255],
|
||||
[166, 86, 40, 255],
|
||||
[247, 129, 191, 255],
|
||||
[153, 153, 153, 255],
|
||||
],
|
||||
dtype=np.uint8,
|
||||
)
|
||||
|
||||
if index is None:
|
||||
index = np.random.randint(0, len(palette))
|
||||
|
||||
if index >= len(palette):
|
||||
index = index % len(palette)
|
||||
|
||||
if use_float:
|
||||
return palette[index].astype(np.float32) / 255
|
||||
else:
|
||||
return palette[index]
|
||||
Reference in New Issue
Block a user