This commit is contained in:
smthem
2025-06-18 08:54:53 +08:00
parent acefe4b5a0
commit 6b47734cbd
45 changed files with 4035 additions and 141 deletions
View File
View File
@@ -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
+57
View File
@@ -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
+58
View File
@@ -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
+343
View File
@@ -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
View File
+235
View File
@@ -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
+184
View File
@@ -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"))
+119
View File
@@ -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]