mirror of
https://git.datalinker.icu/comfyanonymous/ComfyUI
synced 2026-08-26 12:17:57 +08:00
file management and removed code redundancy
This commit is contained in:
parent
b184a61046
commit
dff570e364
@ -6,79 +6,47 @@ import torch.nn as nn
|
|||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
|
|
||||||
|
|
||||||
from typing import Union, Tuple, List, Callable, Optional
|
from typing import Optional
|
||||||
|
|
||||||
import numpy as np
|
|
||||||
from einops import repeat, rearrange
|
|
||||||
from tqdm import tqdm
|
|
||||||
import logging
|
import logging
|
||||||
|
|
||||||
import comfy.ops
|
import comfy.ops
|
||||||
ops = comfy.ops.disable_weight_init
|
ops = comfy.ops.disable_weight_init
|
||||||
|
|
||||||
def generate_dense_grid_points(
|
################################################
|
||||||
bbox_min: np.ndarray,
|
# Volume Decoder
|
||||||
bbox_max: np.ndarray,
|
################################################
|
||||||
octree_resolution: int,
|
|
||||||
indexing: str = "ij",
|
|
||||||
):
|
|
||||||
length = bbox_max - bbox_min
|
|
||||||
num_cells = octree_resolution
|
|
||||||
|
|
||||||
x = np.linspace(bbox_min[0], bbox_max[0], int(num_cells) + 1, dtype=np.float32)
|
class VanillaVolumeDecoder():
|
||||||
y = np.linspace(bbox_min[1], bbox_max[1], int(num_cells) + 1, dtype=np.float32)
|
|
||||||
z = np.linspace(bbox_min[2], bbox_max[2], int(num_cells) + 1, dtype=np.float32)
|
|
||||||
[xs, ys, zs] = np.meshgrid(x, y, z, indexing=indexing)
|
|
||||||
xyz = np.stack((xs, ys, zs), axis=-1)
|
|
||||||
grid_size = [int(num_cells) + 1, int(num_cells) + 1, int(num_cells) + 1]
|
|
||||||
|
|
||||||
return xyz, grid_size, length
|
|
||||||
|
|
||||||
|
|
||||||
class VanillaVolumeDecoder:
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def __call__(
|
def __call__(self, latents: torch.Tensor, geo_decoder: callable, octree_resolution: int, bounds = 1.01,
|
||||||
self,
|
num_chunks: int = 10_000):
|
||||||
latents: torch.FloatTensor,
|
|
||||||
geo_decoder: Callable,
|
|
||||||
bounds: Union[Tuple[float], List[float], float] = 1.01,
|
|
||||||
num_chunks: int = 10000,
|
|
||||||
octree_resolution: int = None,
|
|
||||||
enable_pbar: bool = True,
|
|
||||||
**kwargs,
|
|
||||||
):
|
|
||||||
device = latents.device
|
|
||||||
dtype = latents.dtype
|
|
||||||
batch_size = latents.shape[0]
|
|
||||||
|
|
||||||
# 1. generate query points
|
|
||||||
if isinstance(bounds, float):
|
if isinstance(bounds, float):
|
||||||
bounds = [-bounds, -bounds, -bounds, bounds, bounds, bounds]
|
bounds = [-bounds, -bounds, -bounds, bounds, bounds, bounds]
|
||||||
|
|
||||||
bbox_min, bbox_max = np.array(bounds[0:3]), np.array(bounds[3:6])
|
bbox_min, bbox_max = torch.tensor(bounds[:3]), torch.tensor(bounds[3:])
|
||||||
xyz_samples, grid_size, length = generate_dense_grid_points(
|
|
||||||
bbox_min=bbox_min,
|
x = torch.linspace(bbox_min[0], bbox_max[0], int(octree_resolution) + 1, dtype = torch.float32)
|
||||||
bbox_max=bbox_max,
|
y = torch.linspace(bbox_min[1], bbox_max[1], int(octree_resolution) + 1, dtype = torch.float32)
|
||||||
octree_resolution=octree_resolution,
|
z = torch.linspace(bbox_min[2], bbox_max[2], int(octree_resolution) + 1, dtype = torch.float32)
|
||||||
indexing="ij"
|
|
||||||
)
|
[xs, ys, zs] = torch.meshgrid(x, y, z, indexing = "ij")
|
||||||
xyz_samples = torch.from_numpy(xyz_samples).to(device, dtype=dtype).contiguous().reshape(-1, 3)
|
xyz = torch.stack((xs, ys, zs), axis=-1).to(latents.device, dtype = latents.dtype).contiguous().reshape(-1, 3)
|
||||||
|
grid_size = [int(octree_resolution) + 1, int(octree_resolution) + 1, int(octree_resolution) + 1]
|
||||||
|
|
||||||
# 2. latents to 3d volume
|
|
||||||
batch_logits = []
|
batch_logits = []
|
||||||
for start in tqdm(range(0, xyz_samples.shape[0], num_chunks), desc="Volume Decoding",
|
for start in range(0, xyz.shape[0], num_chunks):
|
||||||
disable=not enable_pbar):
|
chunk_queries = xyz[start: start + num_chunks, :]
|
||||||
chunk_queries = xyz_samples[start: start + num_chunks, :]
|
chunk_queries = chunk_queries.unsqueeze(0).repeat(latents.shape[0], 1, 1)
|
||||||
chunk_queries = repeat(chunk_queries, "p c -> b p c", b=batch_size)
|
logits = geo_decoder(queries = chunk_queries, latents = latents)
|
||||||
logits = geo_decoder(queries=chunk_queries, latents=latents)
|
|
||||||
batch_logits.append(logits)
|
batch_logits.append(logits)
|
||||||
|
|
||||||
grid_logits = torch.cat(batch_logits, dim=1)
|
grid_logits = torch.cat(batch_logits, dim = 1)
|
||||||
grid_logits = grid_logits.view((batch_size, *grid_size)).float()
|
grid_logits = grid_logits.view((latents.shape[0], *grid_size)).float()
|
||||||
|
|
||||||
return grid_logits
|
return grid_logits
|
||||||
|
|
||||||
|
|
||||||
class FourierEmbedder(nn.Module):
|
class FourierEmbedder(nn.Module):
|
||||||
"""The sin/cosine positional embedding. Given an input tensor `x` of shape [n_batch, ..., c_dim], it converts
|
"""The sin/cosine positional embedding. Given an input tensor `x` of shape [n_batch, ..., c_dim], it converts
|
||||||
each feature dimension of `x[..., i]` into:
|
each feature dimension of `x[..., i]` into:
|
||||||
@ -175,13 +143,6 @@ class FourierEmbedder(nn.Module):
|
|||||||
else:
|
else:
|
||||||
return x
|
return x
|
||||||
|
|
||||||
|
|
||||||
class CrossAttentionProcessor:
|
|
||||||
def __call__(self, attn, q, k, v):
|
|
||||||
out = F.scaled_dot_product_attention(q, k, v)
|
|
||||||
return out
|
|
||||||
|
|
||||||
|
|
||||||
class DropPath(nn.Module):
|
class DropPath(nn.Module):
|
||||||
"""Drop paths (Stochastic Depth) per sample (when applied in main path of residual blocks).
|
"""Drop paths (Stochastic Depth) per sample (when applied in main path of residual blocks).
|
||||||
"""
|
"""
|
||||||
@ -232,38 +193,41 @@ class MLP(nn.Module):
|
|||||||
def forward(self, x):
|
def forward(self, x):
|
||||||
return self.drop_path(self.c_proj(self.gelu(self.c_fc(x))))
|
return self.drop_path(self.c_proj(self.gelu(self.c_fc(x))))
|
||||||
|
|
||||||
|
|
||||||
class QKVMultiheadCrossAttention(nn.Module):
|
class QKVMultiheadCrossAttention(nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
*,
|
|
||||||
heads: int,
|
heads: int,
|
||||||
|
n_data = None,
|
||||||
width=None,
|
width=None,
|
||||||
qk_norm=False,
|
qk_norm=False,
|
||||||
norm_layer=ops.LayerNorm
|
norm_layer=ops.LayerNorm
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.heads = heads
|
self.heads = heads
|
||||||
|
self.n_data = n_data
|
||||||
self.q_norm = norm_layer(width // heads, elementwise_affine=True, eps=1e-6) if qk_norm else nn.Identity()
|
self.q_norm = norm_layer(width // heads, elementwise_affine=True, eps=1e-6) if qk_norm else nn.Identity()
|
||||||
self.k_norm = norm_layer(width // heads, elementwise_affine=True, eps=1e-6) if qk_norm else nn.Identity()
|
self.k_norm = norm_layer(width // heads, elementwise_affine=True, eps=1e-6) if qk_norm else nn.Identity()
|
||||||
|
|
||||||
self.attn_processor = CrossAttentionProcessor()
|
|
||||||
|
|
||||||
def forward(self, q, kv):
|
def forward(self, q, kv):
|
||||||
|
|
||||||
_, n_ctx, _ = q.shape
|
_, n_ctx, _ = q.shape
|
||||||
bs, n_data, width = kv.shape
|
bs, n_data, width = kv.shape
|
||||||
|
|
||||||
attn_ch = width // self.heads // 2
|
attn_ch = width // self.heads // 2
|
||||||
q = q.view(bs, n_ctx, self.heads, -1)
|
q = q.view(bs, n_ctx, self.heads, -1)
|
||||||
|
|
||||||
kv = kv.view(bs, n_data, self.heads, -1)
|
kv = kv.view(bs, n_data, self.heads, -1)
|
||||||
k, v = torch.split(kv, attn_ch, dim=-1)
|
k, v = torch.split(kv, attn_ch, dim=-1)
|
||||||
|
|
||||||
q = self.q_norm(q)
|
q = self.q_norm(q)
|
||||||
k = self.k_norm(k)
|
k = self.k_norm(k)
|
||||||
q, k, v = map(lambda t: rearrange(t, 'b n h d -> b h n d', h=self.heads), (q, k, v))
|
|
||||||
out = self.attn_processor(self, q, k, v)
|
|
||||||
out = out.transpose(1, 2).reshape(bs, n_ctx, -1)
|
|
||||||
return out
|
|
||||||
|
|
||||||
|
q, k, v = [t.permute(0, 2, 1, 3) for t in (q, k, v)]
|
||||||
|
out = F.scaled_dot_product_attention(q, k, v)
|
||||||
|
|
||||||
|
out = out.transpose(1, 2).reshape(bs, n_ctx, -1)
|
||||||
|
|
||||||
|
return out
|
||||||
|
|
||||||
class MultiheadCrossAttention(nn.Module):
|
class MultiheadCrossAttention(nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
@ -306,7 +270,6 @@ class MultiheadCrossAttention(nn.Module):
|
|||||||
x = self.c_proj(x)
|
x = self.c_proj(x)
|
||||||
return x
|
return x
|
||||||
|
|
||||||
|
|
||||||
class ResidualCrossAttentionBlock(nn.Module):
|
class ResidualCrossAttentionBlock(nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@ -366,7 +329,7 @@ class QKVMultiheadAttention(nn.Module):
|
|||||||
q = self.q_norm(q)
|
q = self.q_norm(q)
|
||||||
k = self.k_norm(k)
|
k = self.k_norm(k)
|
||||||
|
|
||||||
q, k, v = map(lambda t: rearrange(t, 'b n h d -> b h n d', h=self.heads), (q, k, v))
|
q, k, v = [t.permute(0, 2, 1, 3) for t in (q, k, v)]
|
||||||
out = F.scaled_dot_product_attention(q, k, v).transpose(1, 2).reshape(bs, n_ctx, -1)
|
out = F.scaled_dot_product_attention(q, k, v).transpose(1, 2).reshape(bs, n_ctx, -1)
|
||||||
return out
|
return out
|
||||||
|
|
||||||
@ -383,8 +346,7 @@ class MultiheadAttention(nn.Module):
|
|||||||
drop_path_rate: float = 0.0
|
drop_path_rate: float = 0.0
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.width = width
|
|
||||||
self.heads = heads
|
|
||||||
self.c_qkv = ops.Linear(width, width * 3, bias=qkv_bias)
|
self.c_qkv = ops.Linear(width, width * 3, bias=qkv_bias)
|
||||||
self.c_proj = ops.Linear(width, width)
|
self.c_proj = ops.Linear(width, width)
|
||||||
self.attention = QKVMultiheadAttention(
|
self.attention = QKVMultiheadAttention(
|
||||||
|
|||||||
@ -1,189 +0,0 @@
|
|||||||
import torch
|
|
||||||
import torch.nn as nn
|
|
||||||
import torch.nn.functional as F
|
|
||||||
from image_encoders.dino2 import Dinov2Model
|
|
||||||
from dataclasses import dataclass, asdict
|
|
||||||
# avoid using torchvision by recreating image processing functions
|
|
||||||
|
|
||||||
def resize(img: torch.Tensor, size: int) -> torch.Tensor:
|
|
||||||
|
|
||||||
batched = img.ndim == 4
|
|
||||||
|
|
||||||
if not batched:
|
|
||||||
img = img.unsqueeze(0)
|
|
||||||
|
|
||||||
_, _, h, w = img.shape
|
|
||||||
|
|
||||||
# mantain aspect ratio
|
|
||||||
if h < w:
|
|
||||||
new_h = size
|
|
||||||
new_w = int(w * size / h)
|
|
||||||
else:
|
|
||||||
new_w = size
|
|
||||||
new_h = int(h * size / w)
|
|
||||||
|
|
||||||
img = F.interpolate(img, size = (new_h, new_w), mode = 'bilinear', align_corners = False, antialias = True )
|
|
||||||
|
|
||||||
if not batched:
|
|
||||||
img = img.squeeze(0)
|
|
||||||
|
|
||||||
return img
|
|
||||||
|
|
||||||
|
|
||||||
def center_crop(img: torch.Tensor, size: int) -> torch.Tensor:
|
|
||||||
|
|
||||||
batched = img.ndim == 4
|
|
||||||
if not batched:
|
|
||||||
img = img.unsqueeze(0)
|
|
||||||
|
|
||||||
_, _, h, w = img.shape
|
|
||||||
top = (h - size) // 2
|
|
||||||
left = (w - size) // 2
|
|
||||||
|
|
||||||
|
|
||||||
cropped = img[..., top:top + size, left:left + size]
|
|
||||||
|
|
||||||
if not batched:
|
|
||||||
cropped = cropped.squeeze(0)
|
|
||||||
|
|
||||||
return cropped
|
|
||||||
|
|
||||||
def normalize(img: torch.Tensor, mean: list, std: list) -> torch.Tensor:
|
|
||||||
|
|
||||||
mean = torch.tensor(mean, device = img.device).view(-1, 1, 1)
|
|
||||||
std = torch.tensor(std, device = img.device).view(-1, 1, 1)
|
|
||||||
return (img - mean) / std
|
|
||||||
|
|
||||||
def compose(transforms):
|
|
||||||
def apply(img):
|
|
||||||
for t in transforms:
|
|
||||||
img = t(img)
|
|
||||||
return img
|
|
||||||
return apply
|
|
||||||
|
|
||||||
# configuration for Dino Large
|
|
||||||
@dataclass
|
|
||||||
class DinoConfig():
|
|
||||||
|
|
||||||
hidden_size: int = 1024
|
|
||||||
use_mask_token: bool = True
|
|
||||||
patch_size: int = 14
|
|
||||||
image_size: int = 518
|
|
||||||
num_channels: int = 3
|
|
||||||
num_attention_heads: int = 16
|
|
||||||
attention_probs_dropout_prob: float = 0.0
|
|
||||||
hidden_dropout_prob: float = 0.0
|
|
||||||
mlp_ratio: int = 4
|
|
||||||
num_hidden_layers: int = 24
|
|
||||||
layer_norm_eps: float = 1e-6
|
|
||||||
qkv_bias: bool = True
|
|
||||||
layerscale_value: float = 1.0
|
|
||||||
drop_path_rate: float = 0.0
|
|
||||||
device: str = "cuda"
|
|
||||||
dtype = torch.float16
|
|
||||||
|
|
||||||
class ImageEncoder(nn.Module):
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
config: DinoConfig,
|
|
||||||
use_cls_token = True,
|
|
||||||
image_size = 518,
|
|
||||||
**kwargs,
|
|
||||||
):
|
|
||||||
super().__init__()
|
|
||||||
|
|
||||||
|
|
||||||
import comfy.ops
|
|
||||||
ops = comfy.ops.disable_weight_init
|
|
||||||
|
|
||||||
self.model = Dinov2Model(asdict(config), config.dtype, config.device, operations = ops)
|
|
||||||
|
|
||||||
mean = [0.485, 0.456, 0.406]
|
|
||||||
std = [0.229, 0.224, 0.225]
|
|
||||||
|
|
||||||
self.model.eval()
|
|
||||||
self.model.requires_grad_(False)
|
|
||||||
|
|
||||||
self.use_cls_token = use_cls_token
|
|
||||||
self.size = image_size // 14
|
|
||||||
|
|
||||||
self.num_patches = (image_size // 14) ** 2
|
|
||||||
|
|
||||||
if self.use_cls_token:
|
|
||||||
self.num_patches += 1
|
|
||||||
|
|
||||||
self.transform = compose([
|
|
||||||
lambda x: resize(x, image_size),
|
|
||||||
lambda x: center_crop(x, image_size),
|
|
||||||
lambda x: normalize(x, mean, std),
|
|
||||||
])
|
|
||||||
|
|
||||||
def forward(self, image, value_range=(-1, 1), **kwargs):
|
|
||||||
|
|
||||||
if image.ndim == 3:
|
|
||||||
image = image.unsqueeze(0)
|
|
||||||
|
|
||||||
if value_range is not None:
|
|
||||||
low, high = value_range
|
|
||||||
image = (image - low) / (high - low)
|
|
||||||
|
|
||||||
inputs = self.transform(image)
|
|
||||||
inputs = inputs.to(self.model.device, dtype=self.model.dtype)
|
|
||||||
last_hidden_state = self.model(inputs)
|
|
||||||
|
|
||||||
if not self.use_cls_token:
|
|
||||||
last_hidden_state = last_hidden_state[:, 1:, :]
|
|
||||||
|
|
||||||
return last_hidden_state
|
|
||||||
|
|
||||||
def unconditional_embedding(self, batch_size, **kwargs):
|
|
||||||
|
|
||||||
device = next(self.model.parameters()).device
|
|
||||||
dtype = next(self.model.parameters()).dtype
|
|
||||||
|
|
||||||
zero = torch.zeros(
|
|
||||||
batch_size,
|
|
||||||
self.num_patches,
|
|
||||||
self.model.config.hidden_size,
|
|
||||||
device = device,
|
|
||||||
dtype = dtype,
|
|
||||||
)
|
|
||||||
|
|
||||||
return zero
|
|
||||||
|
|
||||||
class SingleImageEncoder(nn.Module):
|
|
||||||
def __init__(self, config):
|
|
||||||
super().__init__()
|
|
||||||
self.main_image_encoder = ImageEncoder(config)
|
|
||||||
|
|
||||||
def forward(self, image, **kwargs):
|
|
||||||
outputs = {
|
|
||||||
'main': self.main_image_encoder(image, **kwargs),
|
|
||||||
}
|
|
||||||
return outputs
|
|
||||||
|
|
||||||
def unconditional_embedding(self, batch_size, **kwargs):
|
|
||||||
outputs = {
|
|
||||||
'main': self.main_image_encoder.unconditional_embedding(batch_size, **kwargs),
|
|
||||||
}
|
|
||||||
return outputs
|
|
||||||
|
|
||||||
def test_image_encoder():
|
|
||||||
|
|
||||||
torch.manual_seed(2025)
|
|
||||||
config = DinoConfig()
|
|
||||||
image_encoder = SingleImageEncoder(config)
|
|
||||||
|
|
||||||
image = torch.rand(3, 224, 224)
|
|
||||||
|
|
||||||
outputs = image_encoder(image)
|
|
||||||
|
|
||||||
print(outputs)
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
#test_image_encoder()
|
|
||||||
conditioner = SingleImageEncoder(DinoConfig())
|
|
||||||
torch.manual_seed(2025)
|
|
||||||
image = torch.rand(1, 3, 224, 224)
|
|
||||||
outputs = conditioner(image)
|
|
||||||
print(outputs["main"].size())
|
|
||||||
@ -2,9 +2,189 @@ import math
|
|||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
from moe import MoEBlock
|
|
||||||
from torch.nn.attention import SDPBackend
|
from torch.nn.attention import SDPBackend
|
||||||
|
|
||||||
|
class GELU(nn.Module):
|
||||||
|
|
||||||
|
def __init__(self, dim_in: int, dim_out: int):
|
||||||
|
super().__init__()
|
||||||
|
self.proj = nn.Linear(dim_in, dim_out)
|
||||||
|
|
||||||
|
def gelu(self, gate: torch.Tensor) -> torch.Tensor:
|
||||||
|
|
||||||
|
if gate.device.type == "mps":
|
||||||
|
return F.gelu(gate.to(dtype = torch.float32)).to(dtype = gate.dtype)
|
||||||
|
|
||||||
|
return F.gelu(gate)
|
||||||
|
|
||||||
|
def forward(self, hidden_states):
|
||||||
|
|
||||||
|
hidden_states = self.proj(hidden_states)
|
||||||
|
hidden_states = self.gelu(hidden_states)
|
||||||
|
|
||||||
|
return hidden_states
|
||||||
|
|
||||||
|
class FeedForward(nn.Module):
|
||||||
|
|
||||||
|
def __init__(self, dim: int, dim_out = None, mult: int = 4,
|
||||||
|
dropout: float = 0.0, inner_dim = None):
|
||||||
|
|
||||||
|
super().__init__()
|
||||||
|
if inner_dim is None:
|
||||||
|
inner_dim = int(dim * mult)
|
||||||
|
|
||||||
|
dim_out = dim_out if dim_out is not None else dim
|
||||||
|
|
||||||
|
act_fn = GELU(dim, inner_dim)
|
||||||
|
|
||||||
|
self.net = nn.ModuleList([])
|
||||||
|
self.net.append(act_fn)
|
||||||
|
|
||||||
|
self.net.append(nn.Dropout(dropout))
|
||||||
|
self.net.append(nn.Linear(inner_dim, dim_out))
|
||||||
|
|
||||||
|
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||||
|
for module in self.net:
|
||||||
|
hidden_states = module(hidden_states)
|
||||||
|
return hidden_states
|
||||||
|
|
||||||
|
class AddAuxLoss(torch.autograd.Function):
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def forward(ctx, x, loss):
|
||||||
|
# do nothing in forward (no computation)
|
||||||
|
ctx.requires_aux_loss = loss.requires_grad
|
||||||
|
ctx.dtype = loss.dtype
|
||||||
|
|
||||||
|
return x
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def backward(ctx, grad_output):
|
||||||
|
# add the aux loss gradients
|
||||||
|
grad_loss = None
|
||||||
|
# put the aux grad the same as the main grad loss
|
||||||
|
# aux grad contributes equally
|
||||||
|
if ctx.requires_aux_loss:
|
||||||
|
grad_loss = torch.ones(1, dtype = ctx.dtype, device = grad_output.device)
|
||||||
|
|
||||||
|
return grad_output, grad_loss
|
||||||
|
|
||||||
|
class MoEGate(nn.Module):
|
||||||
|
|
||||||
|
def __init__(self, embed_dim, num_experts=16, num_experts_per_tok=2, aux_loss_alpha=0.01):
|
||||||
|
|
||||||
|
super().__init__()
|
||||||
|
self.top_k = num_experts_per_tok
|
||||||
|
self.n_routed_experts = num_experts
|
||||||
|
|
||||||
|
self.alpha = aux_loss_alpha
|
||||||
|
|
||||||
|
self.gating_dim = embed_dim
|
||||||
|
self.weight = nn.Parameter(torch.empty((self.n_routed_experts, self.gating_dim)))
|
||||||
|
|
||||||
|
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||||
|
|
||||||
|
# flatten hidden states
|
||||||
|
hidden_states = hidden_states.view(-1, hidden_states.size(-1))
|
||||||
|
|
||||||
|
# get logits and pass it to softmax
|
||||||
|
logits = F.linear(hidden_states, self.weight, bias = None)
|
||||||
|
scores = logits.softmax(dim = -1)
|
||||||
|
|
||||||
|
topk_weight, topk_idx = torch.topk(scores, k = self.top_k, dim = -1, sorted = False)
|
||||||
|
|
||||||
|
if self.training and self.alpha > 0.0:
|
||||||
|
scores_for_aux = scores
|
||||||
|
|
||||||
|
# used bincount instead of one hot encoding
|
||||||
|
counts = torch.bincount(topk_idx.view(-1), minlength = self.n_routed_experts).float()
|
||||||
|
ce = counts / topk_idx.numel() # normalized expert usage
|
||||||
|
|
||||||
|
# mean expert score
|
||||||
|
Pi = scores_for_aux.mean(0)
|
||||||
|
|
||||||
|
# expert balance loss
|
||||||
|
aux_loss = (Pi * ce * self.n_routed_experts).sum() * self.alpha
|
||||||
|
else:
|
||||||
|
aux_loss = None
|
||||||
|
|
||||||
|
return topk_idx, topk_weight, aux_loss
|
||||||
|
|
||||||
|
class MoEBlock(nn.Module):
|
||||||
|
def __init__(self, dim, num_experts: int = 6, moe_top_k: int = 2, dropout: float = 0.0, ff_inner_dim: int = None):
|
||||||
|
super().__init__()
|
||||||
|
|
||||||
|
self.moe_top_k = moe_top_k
|
||||||
|
self.num_experts = num_experts
|
||||||
|
|
||||||
|
self.experts = nn.ModuleList([
|
||||||
|
FeedForward(dim, dropout = dropout, inner_dim = ff_inner_dim)
|
||||||
|
for _ in range(num_experts)
|
||||||
|
])
|
||||||
|
|
||||||
|
self.gate = MoEGate(dim, num_experts = num_experts, num_experts_per_tok = moe_top_k)
|
||||||
|
self.shared_experts = FeedForward(dim, dropout = dropout, inner_dim = ff_inner_dim)
|
||||||
|
|
||||||
|
def forward(self, hidden_states) -> torch.Tensor:
|
||||||
|
|
||||||
|
identity = hidden_states
|
||||||
|
orig_shape = hidden_states.shape
|
||||||
|
topk_idx, topk_weight, aux_loss = self.gate(hidden_states)
|
||||||
|
|
||||||
|
hidden_states = hidden_states.view(-1, hidden_states.shape[-1])
|
||||||
|
flat_topk_idx = topk_idx.view(-1)
|
||||||
|
|
||||||
|
if self.training:
|
||||||
|
|
||||||
|
hidden_states = hidden_states.repeat_interleave(self.moe_top_k, dim = 0)
|
||||||
|
y = torch.empty_like(hidden_states, dtype = hidden_states.dtype)
|
||||||
|
|
||||||
|
for i, expert in enumerate(self.experts):
|
||||||
|
tmp = expert(hidden_states[flat_topk_idx == i])
|
||||||
|
y[flat_topk_idx == i] = tmp.to(hidden_states.dtype)
|
||||||
|
|
||||||
|
y = (y.view(*topk_weight.shape, -1) * topk_weight.unsqueeze(-1)).sum(dim = 1)
|
||||||
|
y = y.view(*orig_shape)
|
||||||
|
|
||||||
|
y = AddAuxLoss.apply(y, aux_loss)
|
||||||
|
else:
|
||||||
|
y = self.moe_infer(hidden_states, flat_expert_indices = flat_topk_idx,flat_expert_weights = topk_weight.view(-1, 1)).view(*orig_shape)
|
||||||
|
|
||||||
|
y = y + self.shared_experts(identity)
|
||||||
|
|
||||||
|
return y
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def moe_infer(self, x, flat_expert_indices, flat_expert_weights):
|
||||||
|
|
||||||
|
expert_cache = torch.zeros_like(x)
|
||||||
|
idxs = flat_expert_indices.argsort()
|
||||||
|
|
||||||
|
# no need for .numpy().cpu() here
|
||||||
|
tokens_per_expert = flat_expert_indices.bincount().cumsum(0)
|
||||||
|
token_idxs = idxs // self.moe_top_k
|
||||||
|
|
||||||
|
for i, end_idx in enumerate(tokens_per_expert):
|
||||||
|
|
||||||
|
start_idx = 0 if i == 0 else tokens_per_expert[i-1]
|
||||||
|
|
||||||
|
if start_idx == end_idx:
|
||||||
|
continue
|
||||||
|
|
||||||
|
expert = self.experts[i]
|
||||||
|
exp_token_idx = token_idxs[start_idx:end_idx]
|
||||||
|
|
||||||
|
expert_tokens = x[exp_token_idx]
|
||||||
|
expert_out = expert(expert_tokens)
|
||||||
|
|
||||||
|
expert_out.mul_(flat_expert_weights[idxs[start_idx:end_idx]])
|
||||||
|
|
||||||
|
# use index_add_ with a 1-D index tensor directly avoids building a large [N, D] index map and extra memcopy required by scatter_reduce_
|
||||||
|
# + avoid dtype conversion
|
||||||
|
expert_cache.index_add_(0, exp_token_idx, expert_out)
|
||||||
|
|
||||||
|
return expert_cache
|
||||||
|
|
||||||
class Timesteps(nn.Module):
|
class Timesteps(nn.Module):
|
||||||
def __init__(self, num_channels: int, downscale_freq_shift: float = 0.0,
|
def __init__(self, num_channels: int, downscale_freq_shift: float = 0.0,
|
||||||
scale: float = 1.0, max_period: int = 10000):
|
scale: float = 1.0, max_period: int = 10000):
|
||||||
|
|||||||
@ -1,203 +0,0 @@
|
|||||||
import torch
|
|
||||||
import torch.nn as nn
|
|
||||||
import torch.nn.functional as F
|
|
||||||
|
|
||||||
class GELU(nn.Module):
|
|
||||||
|
|
||||||
def __init__(self, dim_in: int, dim_out: int):
|
|
||||||
super().__init__()
|
|
||||||
self.proj = nn.Linear(dim_in, dim_out)
|
|
||||||
|
|
||||||
def gelu(self, gate: torch.Tensor) -> torch.Tensor:
|
|
||||||
|
|
||||||
if gate.device.type == "mps":
|
|
||||||
return F.gelu(gate.to(dtype = torch.float32)).to(dtype = gate.dtype)
|
|
||||||
|
|
||||||
return F.gelu(gate)
|
|
||||||
|
|
||||||
def forward(self, hidden_states):
|
|
||||||
|
|
||||||
hidden_states = self.proj(hidden_states)
|
|
||||||
hidden_states = self.gelu(hidden_states)
|
|
||||||
|
|
||||||
return hidden_states
|
|
||||||
|
|
||||||
class FeedForward(nn.Module):
|
|
||||||
|
|
||||||
def __init__(self, dim: int, dim_out = None, mult: int = 4,
|
|
||||||
dropout: float = 0.0, inner_dim = None):
|
|
||||||
|
|
||||||
super().__init__()
|
|
||||||
if inner_dim is None:
|
|
||||||
inner_dim = int(dim * mult)
|
|
||||||
|
|
||||||
dim_out = dim_out if dim_out is not None else dim
|
|
||||||
|
|
||||||
act_fn = GELU(dim, inner_dim)
|
|
||||||
|
|
||||||
self.net = nn.ModuleList([])
|
|
||||||
self.net.append(act_fn)
|
|
||||||
|
|
||||||
self.net.append(nn.Dropout(dropout))
|
|
||||||
self.net.append(nn.Linear(inner_dim, dim_out))
|
|
||||||
|
|
||||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
|
||||||
for module in self.net:
|
|
||||||
hidden_states = module(hidden_states)
|
|
||||||
return hidden_states
|
|
||||||
|
|
||||||
class AddAuxLoss(torch.autograd.Function):
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def forward(ctx, x, loss):
|
|
||||||
# do nothing in forward (no computation)
|
|
||||||
ctx.requires_aux_loss = loss.requires_grad
|
|
||||||
ctx.dtype = loss.dtype
|
|
||||||
|
|
||||||
return x
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def backward(ctx, grad_output):
|
|
||||||
# add the aux loss gradients
|
|
||||||
grad_loss = None
|
|
||||||
# put the aux grad the same as the main grad loss
|
|
||||||
# aux grad contributes equally
|
|
||||||
if ctx.requires_aux_loss:
|
|
||||||
grad_loss = torch.ones(1, dtype = ctx.dtype, device = grad_output.device)
|
|
||||||
|
|
||||||
return grad_output, grad_loss
|
|
||||||
|
|
||||||
class MoEGate(nn.Module):
|
|
||||||
|
|
||||||
def __init__(self, embed_dim, num_experts=16, num_experts_per_tok=2, aux_loss_alpha=0.01):
|
|
||||||
|
|
||||||
super().__init__()
|
|
||||||
self.top_k = num_experts_per_tok
|
|
||||||
self.n_routed_experts = num_experts
|
|
||||||
|
|
||||||
self.alpha = aux_loss_alpha
|
|
||||||
|
|
||||||
self.gating_dim = embed_dim
|
|
||||||
self.weight = nn.Parameter(torch.empty((self.n_routed_experts, self.gating_dim)))
|
|
||||||
|
|
||||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
|
||||||
|
|
||||||
# flatten hidden states
|
|
||||||
hidden_states = hidden_states.view(-1, hidden_states.size(-1))
|
|
||||||
|
|
||||||
# get logits and pass it to softmax
|
|
||||||
logits = F.linear(hidden_states, self.weight, bias = None)
|
|
||||||
scores = logits.softmax(dim = -1)
|
|
||||||
|
|
||||||
topk_weight, topk_idx = torch.topk(scores, k = self.top_k, dim = -1, sorted = False)
|
|
||||||
|
|
||||||
if self.training and self.alpha > 0.0:
|
|
||||||
scores_for_aux = scores
|
|
||||||
|
|
||||||
# used bincount instead of one hot encoding
|
|
||||||
counts = torch.bincount(topk_idx.view(-1), minlength = self.n_routed_experts).float()
|
|
||||||
ce = counts / topk_idx.numel() # normalized expert usage
|
|
||||||
|
|
||||||
# mean expert score
|
|
||||||
Pi = scores_for_aux.mean(0)
|
|
||||||
|
|
||||||
# expert balance loss
|
|
||||||
aux_loss = (Pi * ce * self.n_routed_experts).sum() * self.alpha
|
|
||||||
else:
|
|
||||||
aux_loss = None
|
|
||||||
|
|
||||||
return topk_idx, topk_weight, aux_loss
|
|
||||||
|
|
||||||
class MoEBlock(nn.Module):
|
|
||||||
def __init__(self, dim, num_experts: int = 6, moe_top_k: int = 2, dropout: float = 0.0, ff_inner_dim: int = None):
|
|
||||||
super().__init__()
|
|
||||||
|
|
||||||
self.moe_top_k = moe_top_k
|
|
||||||
self.num_experts = num_experts
|
|
||||||
|
|
||||||
self.experts = nn.ModuleList([
|
|
||||||
FeedForward(dim, dropout = dropout, inner_dim = ff_inner_dim)
|
|
||||||
for _ in range(num_experts)
|
|
||||||
])
|
|
||||||
|
|
||||||
self.gate = MoEGate(dim, num_experts = num_experts, num_experts_per_tok = moe_top_k)
|
|
||||||
self.shared_experts = FeedForward(dim, dropout = dropout, inner_dim = ff_inner_dim)
|
|
||||||
|
|
||||||
def forward(self, hidden_states) -> torch.Tensor:
|
|
||||||
|
|
||||||
identity = hidden_states
|
|
||||||
orig_shape = hidden_states.shape
|
|
||||||
topk_idx, topk_weight, aux_loss = self.gate(hidden_states)
|
|
||||||
|
|
||||||
hidden_states = hidden_states.view(-1, hidden_states.shape[-1])
|
|
||||||
flat_topk_idx = topk_idx.view(-1)
|
|
||||||
|
|
||||||
if self.training:
|
|
||||||
|
|
||||||
hidden_states = hidden_states.repeat_interleave(self.moe_top_k, dim = 0)
|
|
||||||
y = torch.empty_like(hidden_states, dtype = hidden_states.dtype)
|
|
||||||
|
|
||||||
for i, expert in enumerate(self.experts):
|
|
||||||
tmp = expert(hidden_states[flat_topk_idx == i])
|
|
||||||
y[flat_topk_idx == i] = tmp.to(hidden_states.dtype)
|
|
||||||
|
|
||||||
y = (y.view(*topk_weight.shape, -1) * topk_weight.unsqueeze(-1)).sum(dim = 1)
|
|
||||||
y = y.view(*orig_shape)
|
|
||||||
|
|
||||||
y = AddAuxLoss.apply(y, aux_loss)
|
|
||||||
else:
|
|
||||||
y = self.moe_infer(hidden_states, flat_expert_indices = flat_topk_idx,flat_expert_weights = topk_weight.view(-1, 1)).view(*orig_shape)
|
|
||||||
|
|
||||||
y = y + self.shared_experts(identity)
|
|
||||||
|
|
||||||
return y
|
|
||||||
|
|
||||||
@torch.no_grad()
|
|
||||||
def moe_infer(self, x, flat_expert_indices, flat_expert_weights):
|
|
||||||
|
|
||||||
expert_cache = torch.zeros_like(x)
|
|
||||||
idxs = flat_expert_indices.argsort()
|
|
||||||
|
|
||||||
# no need for .numpy().cpu() here
|
|
||||||
tokens_per_expert = flat_expert_indices.bincount().cumsum(0)
|
|
||||||
token_idxs = idxs // self.moe_top_k
|
|
||||||
|
|
||||||
for i, end_idx in enumerate(tokens_per_expert):
|
|
||||||
|
|
||||||
start_idx = 0 if i == 0 else tokens_per_expert[i-1]
|
|
||||||
|
|
||||||
if start_idx == end_idx:
|
|
||||||
continue
|
|
||||||
|
|
||||||
expert = self.experts[i]
|
|
||||||
exp_token_idx = token_idxs[start_idx:end_idx]
|
|
||||||
|
|
||||||
expert_tokens = x[exp_token_idx]
|
|
||||||
expert_out = expert(expert_tokens)
|
|
||||||
|
|
||||||
expert_out.mul_(flat_expert_weights[idxs[start_idx:end_idx]])
|
|
||||||
|
|
||||||
# use index_add_ with a 1-D index tensor directly avoids building a large [N, D] index map and extra memcopy required by scatter_reduce_
|
|
||||||
# + avoid dtype conversion
|
|
||||||
expert_cache.index_add_(0, exp_token_idx, expert_out)
|
|
||||||
|
|
||||||
return expert_cache
|
|
||||||
|
|
||||||
def test_moe():
|
|
||||||
|
|
||||||
torch.manual_seed(2025)
|
|
||||||
import time
|
|
||||||
|
|
||||||
start = time.time()
|
|
||||||
|
|
||||||
moe_gate = MoEGate(512)
|
|
||||||
print(moe_gate(torch.rand(1, 71, 512)))
|
|
||||||
|
|
||||||
#moe_block = MoEBlock(512)
|
|
||||||
#moe_block(torch.rand(1, 77, 512))
|
|
||||||
|
|
||||||
timing = time.time() - start
|
|
||||||
print(timing)
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
test_moe()
|
|
||||||
@ -1,192 +0,0 @@
|
|||||||
import torch
|
|
||||||
import torch.nn as nn
|
|
||||||
from PIL import Image
|
|
||||||
from typing import List, Union
|
|
||||||
from torch.utils._pytree import tree_map
|
|
||||||
from torch.utils.data._utils.collate import default_collate
|
|
||||||
|
|
||||||
import sys
|
|
||||||
import os
|
|
||||||
|
|
||||||
def find_project_root(target_folder_name="ComfyUI"):
|
|
||||||
""" Walks directory until it finds ComfyUI base directroy """
|
|
||||||
current = os.path.abspath(os.path.dirname(__file__))
|
|
||||||
while True:
|
|
||||||
if os.path.basename(current) == target_folder_name:
|
|
||||||
return current
|
|
||||||
parent = os.path.dirname(current)
|
|
||||||
if parent == current:
|
|
||||||
raise RuntimeError(f"Could not find folder named '{target_folder_name}' in parent paths.")
|
|
||||||
current = parent
|
|
||||||
|
|
||||||
comfyui_root = find_project_root()
|
|
||||||
sys.path.append(comfyui_root)
|
|
||||||
|
|
||||||
from comfy_extras.nodes_hunyuan3d import save_glb
|
|
||||||
|
|
||||||
class Hunyuan3DDiTFlowMatchingPipeline(nn.Module):
|
|
||||||
def __init__(self, model, vae, conditioner, image_processor, scheduler, device, dtype):
|
|
||||||
super().__init__()
|
|
||||||
|
|
||||||
self.vae = vae
|
|
||||||
self.model = model
|
|
||||||
self.conditioner = conditioner
|
|
||||||
self.image_processor = image_processor
|
|
||||||
self.scheduler = scheduler
|
|
||||||
self.device = device
|
|
||||||
self.dtype = dtype
|
|
||||||
|
|
||||||
def compile(self):
|
|
||||||
self.vae = torch.compile(self.vae)
|
|
||||||
self.model = torch.compile(self.model)
|
|
||||||
self.conditioner = torch.compile(self.conditioner)
|
|
||||||
|
|
||||||
def load_ckpt(self, checkpoint_path: str):
|
|
||||||
|
|
||||||
checkpoint = torch.load(checkpoint_path, weights_only = True)
|
|
||||||
self.model.load_state_dict(checkpoint["model"])
|
|
||||||
self.vae.load_state_dict(checkpoint["vae"])
|
|
||||||
self.conditioner.load_state_dict(checkpoint["conditioner"])
|
|
||||||
|
|
||||||
|
|
||||||
def encode_cond(self, image, additional_cond_inputs, do_classifier_free_guidance):
|
|
||||||
|
|
||||||
bsz = image.shape[0]
|
|
||||||
cond = self.conditioner(image=image, **additional_cond_inputs)
|
|
||||||
|
|
||||||
if do_classifier_free_guidance:
|
|
||||||
un_cond = self.conditioner.unconditional_embedding(bsz, **additional_cond_inputs)
|
|
||||||
|
|
||||||
# avoid python recursion by using tree_map
|
|
||||||
_fn = lambda x, y: torch.cat([x, y], dim=0).to(self.dtype)
|
|
||||||
|
|
||||||
cond = tree_map(_fn, cond, un_cond)
|
|
||||||
|
|
||||||
return cond
|
|
||||||
|
|
||||||
def to(self, device=None, dtype=None):
|
|
||||||
if dtype is not None:
|
|
||||||
self.dtype = dtype
|
|
||||||
self.vae.to(dtype=dtype)
|
|
||||||
self.model.to(dtype=dtype)
|
|
||||||
self.conditioner.to(dtype=dtype)
|
|
||||||
if device is not None:
|
|
||||||
self.device = torch.device(device)
|
|
||||||
self.vae.to(device)
|
|
||||||
self.model.to(device)
|
|
||||||
self.conditioner.to(device)
|
|
||||||
|
|
||||||
def prepare_images(self, images):
|
|
||||||
|
|
||||||
if isinstance(images, (str, Image.Image)):
|
|
||||||
return self.image_processor(images)
|
|
||||||
|
|
||||||
outputs = []
|
|
||||||
for image in images:
|
|
||||||
output = self.image_processor(image)
|
|
||||||
outputs.append(output)
|
|
||||||
|
|
||||||
return default_collate(outputs)
|
|
||||||
|
|
||||||
def get_guidance_scale_embedding(self, w, embedding_dim=512, dtype=torch.float32):
|
|
||||||
|
|
||||||
assert len(w.shape) == 1
|
|
||||||
w = w * 1000.0
|
|
||||||
|
|
||||||
half_dim = embedding_dim // 2
|
|
||||||
emb = torch.log(torch.tensor(10000.0)) / (half_dim - 1)
|
|
||||||
|
|
||||||
emb = torch.exp(torch.arange(half_dim, dtype=dtype) * -emb)
|
|
||||||
emb = w.to(dtype)[:, None] * emb[None, :]
|
|
||||||
|
|
||||||
emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=1)
|
|
||||||
if embedding_dim % 2 == 1: # zero pad
|
|
||||||
emb = torch.nn.functional.pad(emb, (0, 1))
|
|
||||||
|
|
||||||
assert emb.shape == (w.shape[0], embedding_dim)
|
|
||||||
|
|
||||||
return emb
|
|
||||||
|
|
||||||
def prepare_latents(self, batch_size, dtype, device):
|
|
||||||
|
|
||||||
shape = (batch_size, *self.vae.latent_shape)
|
|
||||||
latents = torch.randn(shape, dtype = dtype, device = device)
|
|
||||||
|
|
||||||
return latents
|
|
||||||
|
|
||||||
@torch.inference_mode()
|
|
||||||
def __call__(
|
|
||||||
self,
|
|
||||||
image: Union[str, List[str], Image.Image, dict, List[dict], torch.Tensor] = None,
|
|
||||||
guidance_scale: float = 5.0,
|
|
||||||
bounds = 1.01,
|
|
||||||
octree_res = 384,
|
|
||||||
num_chunks = 8000,
|
|
||||||
save_file = None,
|
|
||||||
**kwargs,
|
|
||||||
):
|
|
||||||
|
|
||||||
callback = kwargs.pop("callback", None)
|
|
||||||
callback_steps = kwargs.pop("callback_steps", None)
|
|
||||||
|
|
||||||
device = self.device
|
|
||||||
dtype = self.dtype
|
|
||||||
do_classifier_free_guidance = guidance_scale >= 0 and not (
|
|
||||||
hasattr(self.model, 'guidance_embed') and
|
|
||||||
self.model.guidance_embed is True
|
|
||||||
)
|
|
||||||
|
|
||||||
cond_inputs = self.prepare_images(image)
|
|
||||||
image = cond_inputs.pop('image')
|
|
||||||
|
|
||||||
cond = self.encode_cond(
|
|
||||||
image = image,
|
|
||||||
additional_cond_inputs = cond_inputs,
|
|
||||||
do_classifier_free_guidance = do_classifier_free_guidance,
|
|
||||||
)
|
|
||||||
|
|
||||||
guidance = None
|
|
||||||
batch_size = image.shape[0]
|
|
||||||
|
|
||||||
latents = self.prepare_latents(batch_size, dtype, device)
|
|
||||||
|
|
||||||
if hasattr(self.model, 'guidance_embed') and \
|
|
||||||
self.model.guidance_embed is True:
|
|
||||||
guidance = torch.tensor([guidance_scale] * batch_size, device=device, dtype=dtype)
|
|
||||||
|
|
||||||
timesteps = self.scheduler.timesteps
|
|
||||||
|
|
||||||
for i, t in enumerate(timesteps):
|
|
||||||
# expand the latents if we are doing classifier free guidance
|
|
||||||
if do_classifier_free_guidance:
|
|
||||||
latent_model_input = torch.cat([latents] * 2)
|
|
||||||
else:
|
|
||||||
latent_model_input = latents
|
|
||||||
|
|
||||||
timestep = t.expand(latent_model_input.shape[0]).to(latents.dtype)
|
|
||||||
timestep = timestep / self.scheduler.num_training_timesteps
|
|
||||||
noise_pred = self.model(latent_model_input, timestep, cond, guidance=guidance)
|
|
||||||
|
|
||||||
if do_classifier_free_guidance:
|
|
||||||
noise_pred_cond, noise_pred_uncond = noise_pred.chunk(2)
|
|
||||||
noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_cond - noise_pred_uncond)
|
|
||||||
|
|
||||||
# compute the previous noisy sample x_t -> x_t-1
|
|
||||||
latents = self.scheduler.step(noise_pred, latents)
|
|
||||||
|
|
||||||
if callback is not None and i % callback_steps == 0:
|
|
||||||
step_idx = i // getattr(self.scheduler, "order", 1)
|
|
||||||
callback(step_idx, t, latents)
|
|
||||||
|
|
||||||
latents = 1. / self.vae.scale_factor * latents
|
|
||||||
mesh = self.vae.decode(latents, bounds = bounds, octree_res = octree_res, num_chunks = num_chunks)
|
|
||||||
|
|
||||||
try:
|
|
||||||
if save_file is not None:
|
|
||||||
for i, output in enumerate(mesh):
|
|
||||||
output_file = f"{save_file}_{i}" if len(mesh) > 1 else save_file
|
|
||||||
save_glb(output.mesh_v, output.mesh_f, output_file, numpy_ready = True)
|
|
||||||
except Exception as e:
|
|
||||||
print(e)
|
|
||||||
|
|
||||||
return mesh
|
|
||||||
@ -1,100 +0,0 @@
|
|||||||
import torch
|
|
||||||
|
|
||||||
class EulerScheduler(torch.nn.Module):
|
|
||||||
def __init__(self, num_training_timesteps: int = 1_000, shift: float = 1,
|
|
||||||
num_inference_timesteps: int = 50, inference: bool = True, device: str = "cuda"):
|
|
||||||
super(EulerScheduler, self).__init__()
|
|
||||||
|
|
||||||
# compute timestep values so we can index into them later
|
|
||||||
timesteps = torch.linspace(1, num_training_timesteps, num_training_timesteps).to(torch.float32)
|
|
||||||
|
|
||||||
# normalize between 0 and 1
|
|
||||||
sigmas = timesteps / num_training_timesteps
|
|
||||||
|
|
||||||
# staticaly shift (fixed image size assumed)
|
|
||||||
self.sigmas = sigmas * shift / (1 + (shift - 1) * sigmas)
|
|
||||||
|
|
||||||
# get timesteps after shifting
|
|
||||||
self.timesteps = self.sigmas * num_training_timesteps
|
|
||||||
|
|
||||||
self.num_training_timesteps = num_training_timesteps
|
|
||||||
self.num_inference_timesteps = num_inference_timesteps
|
|
||||||
|
|
||||||
if inference:
|
|
||||||
|
|
||||||
sigmas = torch.linspace(0, 1, num_inference_timesteps, dtype = torch.float32, device = device)
|
|
||||||
timesteps = sigmas * self.num_training_timesteps
|
|
||||||
|
|
||||||
self.timesteps = timesteps.to(device = device)
|
|
||||||
self.sigmas = torch.cat([sigmas, torch.ones(1, device=sigmas.device)])
|
|
||||||
|
|
||||||
self._step_index = 0
|
|
||||||
|
|
||||||
def sigma_to_timestep(self, sigma):
|
|
||||||
return sigma * self.num_training_timesteps
|
|
||||||
|
|
||||||
def index_for_timestep(self, timestep, schedule_timesteps = None):
|
|
||||||
|
|
||||||
indices = (schedule_timesteps == timestep).nonzero()
|
|
||||||
|
|
||||||
return indices[0].item()
|
|
||||||
|
|
||||||
|
|
||||||
def add_noise(self, image: torch.FloatTensor, timestep: float):
|
|
||||||
|
|
||||||
noise = torch.randn_like(image)
|
|
||||||
|
|
||||||
if image.device.type == "mps" and torch.is_floating_point(timestep):
|
|
||||||
# mps does not support float64
|
|
||||||
schedule_timesteps = self.timesteps.to(image.device, dtype=torch.float32)
|
|
||||||
timestep = timestep.to(image.device, dtype=torch.float32)
|
|
||||||
else:
|
|
||||||
schedule_timesteps = self.timesteps.to(image.device)
|
|
||||||
timestep = timestep.to(image.device)
|
|
||||||
|
|
||||||
# supports a list and a float
|
|
||||||
if not isinstance(timestep, torch.Tensor) or timestep.ndim == 0:
|
|
||||||
step_indices = [self.index_for_timestep(timestep, schedule_timesteps)]
|
|
||||||
else:
|
|
||||||
step_indices = [self.index_for_timestep(t, schedule_timesteps) for t in timestep]
|
|
||||||
|
|
||||||
sigma = self.sigmas[step_indices].flatten().to(dtype = image.dtype, device = image.device)
|
|
||||||
|
|
||||||
while len(sigma.shape) < len(image.shape):
|
|
||||||
sigma = sigma.unsqueeze(-1)
|
|
||||||
|
|
||||||
noised_image = (1.0 - sigma) * image + noise * sigma
|
|
||||||
|
|
||||||
return noised_image
|
|
||||||
|
|
||||||
@torch.no_grad()
|
|
||||||
def step(self, model_output: torch.FloatTensor, sample: torch.FloatTensor,):
|
|
||||||
|
|
||||||
sample = sample.to(torch.float32)
|
|
||||||
|
|
||||||
sigma = self.sigmas[self._step_index]
|
|
||||||
sigma_next = self.sigmas[self._step_index + 1]
|
|
||||||
|
|
||||||
prev_sample = sample + (sigma_next - sigma) * model_output
|
|
||||||
|
|
||||||
prev_sample = prev_sample.to(model_output.dtype)
|
|
||||||
|
|
||||||
self._step_index += 1
|
|
||||||
|
|
||||||
return prev_sample
|
|
||||||
|
|
||||||
def test_scheduler():
|
|
||||||
|
|
||||||
scheduler = EulerScheduler()
|
|
||||||
|
|
||||||
torch.manual_seed(2025)
|
|
||||||
image = torch.rand(1, 224, 224, dtype = torch.float32)
|
|
||||||
latent = torch.rand(1, 224, 224, dtype = torch.float32)
|
|
||||||
|
|
||||||
output = scheduler.add_noise(image, timestep = torch.tensor([scheduler.timesteps[15]]))
|
|
||||||
output = scheduler.reverse_flow(model_output = image, current_sample = latent)
|
|
||||||
|
|
||||||
print(output)
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
test_scheduler()
|
|
||||||
@ -1,7 +1,3 @@
|
|||||||
# replaced torch.ops.torch_cluster.fps with a manual implementation
|
|
||||||
# to avoid having torch_cluster downloaded as dependency
|
|
||||||
# also the dependency takes a long time to install
|
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
import math
|
import math
|
||||||
@ -9,7 +5,9 @@ import numpy as np
|
|||||||
from skimage import measure
|
from skimage import measure
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
import torch.nn.functional as F
|
from hunyuan3d.vae import (
|
||||||
|
CrossAttentionDecoder, Transformer, ResidualCrossAttentionBlock, FourierEmbedder, VanillaVolumeDecoder
|
||||||
|
)
|
||||||
|
|
||||||
def fps(src: Tensor, batch: Tensor, sampling_ratio: float, start_random: bool = True):
|
def fps(src: Tensor, batch: Tensor, sampling_ratio: float, start_random: bool = True):
|
||||||
|
|
||||||
@ -54,332 +52,6 @@ def fps(src: Tensor, batch: Tensor, sampling_ratio: float, start_random: bool =
|
|||||||
sampled_indicies.append(torch.arange(start, end)[selected])
|
sampled_indicies.append(torch.arange(start, end)[selected])
|
||||||
|
|
||||||
return torch.cat(sampled_indicies, dim = 0)
|
return torch.cat(sampled_indicies, dim = 0)
|
||||||
|
|
||||||
class DropPath(nn.Module):
|
|
||||||
"""Drop paths (Stochastic Depth) per sample (when applied in main path of residual blocks).
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(self, drop_prob: float = 0., scale_by_keep: bool = True):
|
|
||||||
super(DropPath, self).__init__()
|
|
||||||
self.drop_prob = drop_prob
|
|
||||||
self.scale_by_keep = scale_by_keep
|
|
||||||
|
|
||||||
def forward(self, x):
|
|
||||||
|
|
||||||
keep_prob = 1 - self.drop_prob
|
|
||||||
shape = (x.shape[0],) + (1,) * (x.ndim - 1) # work with diff dim tensors, not just 2D ConvNets
|
|
||||||
|
|
||||||
random_tensor = x.new_empty(shape).bernoulli_(keep_prob)
|
|
||||||
|
|
||||||
if keep_prob > 0.0 and self.scale_by_keep:
|
|
||||||
random_tensor.div_(keep_prob)
|
|
||||||
|
|
||||||
return x * random_tensor
|
|
||||||
|
|
||||||
class MLP(nn.Module):
|
|
||||||
def __init__(self, width: int, ratio: int = 4, drop_path_rate: float = 0):
|
|
||||||
super().__init__()
|
|
||||||
self.gelu = nn.GELU()
|
|
||||||
self.c_fc = nn.Linear(width, width * ratio)
|
|
||||||
self.c_proj = nn.Linear(width * ratio, width)
|
|
||||||
self.drop_path = DropPath(drop_path_rate) if drop_path_rate > 0. else nn.Identity()
|
|
||||||
|
|
||||||
def forward(self, x):
|
|
||||||
return self.drop_path(self.c_proj(self.gelu(self.c_fc(x))))
|
|
||||||
|
|
||||||
|
|
||||||
class QKVMultiheadAttention(nn.Module):
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
heads: int,
|
|
||||||
n_ctx: int,
|
|
||||||
width=None,
|
|
||||||
qk_norm=False,
|
|
||||||
norm_layer=nn.LayerNorm
|
|
||||||
):
|
|
||||||
super().__init__()
|
|
||||||
self.heads = heads
|
|
||||||
self.n_ctx = n_ctx
|
|
||||||
self.q_norm = norm_layer(width // heads, elementwise_affine=True, eps=1e-6) if qk_norm else nn.Identity()
|
|
||||||
self.k_norm = norm_layer(width // heads, elementwise_affine=True, eps=1e-6) if qk_norm else nn.Identity()
|
|
||||||
|
|
||||||
def forward(self, qkv):
|
|
||||||
bs, n_ctx, width = qkv.shape
|
|
||||||
attn_ch = width // self.heads // 3
|
|
||||||
qkv = qkv.view(bs, n_ctx, self.heads, -1)
|
|
||||||
q, k, v = torch.split(qkv, attn_ch, dim=-1)
|
|
||||||
|
|
||||||
q = self.q_norm(q)
|
|
||||||
k = self.k_norm(k)
|
|
||||||
|
|
||||||
q, k, v = [t.permute(0, 2, 1, 3) for t in (q, k, v)]
|
|
||||||
out = F.scaled_dot_product_attention(q, k, v).transpose(1, 2).reshape(bs, n_ctx, -1)
|
|
||||||
return out
|
|
||||||
|
|
||||||
|
|
||||||
class MultiheadAttention(nn.Module):
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
n_ctx: int,
|
|
||||||
width: int,
|
|
||||||
heads: int,
|
|
||||||
qkv_bias: bool,
|
|
||||||
norm_layer = nn.LayerNorm,
|
|
||||||
qk_norm: bool = False,
|
|
||||||
drop_path_rate: float = 0.0
|
|
||||||
):
|
|
||||||
super().__init__()
|
|
||||||
|
|
||||||
self.c_qkv = nn.Linear(width, width * 3, bias=qkv_bias)
|
|
||||||
self.c_proj = nn.Linear(width, width)
|
|
||||||
|
|
||||||
self.attention = QKVMultiheadAttention(
|
|
||||||
heads = heads,
|
|
||||||
n_ctx = n_ctx,
|
|
||||||
width = width,
|
|
||||||
norm_layer = norm_layer,
|
|
||||||
qk_norm = qk_norm
|
|
||||||
)
|
|
||||||
self.drop_path = DropPath(drop_path_rate) if drop_path_rate > 0. else nn.Identity()
|
|
||||||
|
|
||||||
def forward(self, x):
|
|
||||||
x = self.c_qkv(x)
|
|
||||||
x = self.attention(x)
|
|
||||||
x = self.drop_path(self.c_proj(x))
|
|
||||||
return x
|
|
||||||
|
|
||||||
|
|
||||||
class ResAttnBlock(nn.Module):
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
*,
|
|
||||||
n_ctx: int,
|
|
||||||
width: int,
|
|
||||||
heads: int,
|
|
||||||
qkv_bias: bool = True,
|
|
||||||
norm_layer=nn.LayerNorm,
|
|
||||||
qk_norm: bool = False,
|
|
||||||
drop_path_rate: float = 0.0,
|
|
||||||
):
|
|
||||||
super().__init__()
|
|
||||||
self.attn = MultiheadAttention(
|
|
||||||
n_ctx=n_ctx,
|
|
||||||
width=width,
|
|
||||||
heads=heads,
|
|
||||||
qkv_bias=qkv_bias,
|
|
||||||
norm_layer=norm_layer,
|
|
||||||
qk_norm=qk_norm,
|
|
||||||
drop_path_rate=drop_path_rate
|
|
||||||
)
|
|
||||||
self.ln_1 = norm_layer(width, elementwise_affine=True, eps=1e-6)
|
|
||||||
self.mlp = MLP(width=width, drop_path_rate=drop_path_rate)
|
|
||||||
self.ln_2 = norm_layer(width, elementwise_affine=True, eps=1e-6)
|
|
||||||
|
|
||||||
def forward(self, x: torch.Tensor):
|
|
||||||
x = x + self.attn(self.ln_1(x))
|
|
||||||
x = x + self.mlp(self.ln_2(x))
|
|
||||||
return x
|
|
||||||
|
|
||||||
class Transformer(nn.Module):
|
|
||||||
def __init__(self, n_ctx: int, heads: int, width: int, depth: int,
|
|
||||||
qkv_bias: bool = True, qk_norm: bool = False, drop_path_rate: float = 0.0):
|
|
||||||
super().__init__()
|
|
||||||
|
|
||||||
self.resblocks = nn.ModuleList([
|
|
||||||
ResAttnBlock(n_ctx = n_ctx,
|
|
||||||
heads = heads,
|
|
||||||
width = width,
|
|
||||||
qkv_bias = qkv_bias,
|
|
||||||
qk_norm = qk_norm,
|
|
||||||
drop_path_rate = drop_path_rate)
|
|
||||||
for _ in range(depth)
|
|
||||||
])
|
|
||||||
|
|
||||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
|
||||||
|
|
||||||
for resnet in self.resblocks:
|
|
||||||
x = resnet(x)
|
|
||||||
|
|
||||||
return x
|
|
||||||
|
|
||||||
class QKVMultiheadCrossAttention(nn.Module):
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
heads: int,
|
|
||||||
n_data = None,
|
|
||||||
width=None,
|
|
||||||
qk_norm=False,
|
|
||||||
norm_layer=nn.LayerNorm
|
|
||||||
):
|
|
||||||
super().__init__()
|
|
||||||
self.heads = heads
|
|
||||||
self.n_data = n_data
|
|
||||||
self.q_norm = norm_layer(width // heads, elementwise_affine=True, eps=1e-6) if qk_norm else nn.Identity()
|
|
||||||
self.k_norm = norm_layer(width // heads, elementwise_affine=True, eps=1e-6) if qk_norm else nn.Identity()
|
|
||||||
|
|
||||||
def forward(self, q, kv):
|
|
||||||
|
|
||||||
_, n_ctx, _ = q.shape
|
|
||||||
bs, n_data, width = kv.shape
|
|
||||||
|
|
||||||
attn_ch = width // self.heads // 2
|
|
||||||
q = q.view(bs, n_ctx, self.heads, -1)
|
|
||||||
|
|
||||||
kv = kv.view(bs, n_data, self.heads, -1)
|
|
||||||
k, v = torch.split(kv, attn_ch, dim=-1)
|
|
||||||
|
|
||||||
q = self.q_norm(q)
|
|
||||||
k = self.k_norm(k)
|
|
||||||
|
|
||||||
q, k, v = [t.permute(0, 2, 1, 3) for t in (q, k, v)]
|
|
||||||
out = F.scaled_dot_product_attention(q, k, v)
|
|
||||||
|
|
||||||
out = out.transpose(1, 2).reshape(bs, n_ctx, -1)
|
|
||||||
|
|
||||||
return out
|
|
||||||
|
|
||||||
|
|
||||||
class MultiheadCrossAttention(nn.Module):
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
width: int,
|
|
||||||
heads: int,
|
|
||||||
qkv_bias: bool = False,
|
|
||||||
n_data = None,
|
|
||||||
norm_layer = nn.LayerNorm,
|
|
||||||
qk_norm: bool = False,
|
|
||||||
kv_cache: bool = False,
|
|
||||||
):
|
|
||||||
super().__init__()
|
|
||||||
|
|
||||||
self.c_q = nn.Linear(width, width, bias=qkv_bias)
|
|
||||||
self.c_kv = nn.Linear(width, width * 2, bias=qkv_bias)
|
|
||||||
self.c_proj = nn.Linear(width, width)
|
|
||||||
|
|
||||||
self.attention = QKVMultiheadCrossAttention(
|
|
||||||
heads = heads,
|
|
||||||
n_data = n_data,
|
|
||||||
width = width,
|
|
||||||
norm_layer = norm_layer,
|
|
||||||
qk_norm = qk_norm
|
|
||||||
)
|
|
||||||
|
|
||||||
self.kv_cache = kv_cache
|
|
||||||
self.data = None
|
|
||||||
|
|
||||||
def forward(self, x, data):
|
|
||||||
x = self.c_q(x)
|
|
||||||
|
|
||||||
if self.kv_cache:
|
|
||||||
if self.data is None:
|
|
||||||
self.data = self.c_kv(data)
|
|
||||||
|
|
||||||
data = self.data
|
|
||||||
else:
|
|
||||||
data = self.c_kv(data)
|
|
||||||
|
|
||||||
x = self.attention(x, data)
|
|
||||||
x = self.c_proj(x)
|
|
||||||
|
|
||||||
return x
|
|
||||||
|
|
||||||
|
|
||||||
class ResidualCrossAttentionBlock(nn.Module):
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
width: int,
|
|
||||||
heads: int,
|
|
||||||
n_data: int = None,
|
|
||||||
mlp_expand_ratio: int = 4,
|
|
||||||
qkv_bias: bool = False,
|
|
||||||
norm_layer=nn.LayerNorm,
|
|
||||||
qk_norm: bool = False
|
|
||||||
):
|
|
||||||
super().__init__()
|
|
||||||
|
|
||||||
self.attn = MultiheadCrossAttention(
|
|
||||||
n_data=n_data,
|
|
||||||
width = width,
|
|
||||||
heads=heads,
|
|
||||||
qkv_bias=qkv_bias,
|
|
||||||
norm_layer=norm_layer,
|
|
||||||
qk_norm=qk_norm
|
|
||||||
)
|
|
||||||
|
|
||||||
self.ln_1 = norm_layer(width, elementwise_affine = True, eps = 1e-6)
|
|
||||||
self.ln_2 = norm_layer(width, elementwise_affine = True, eps = 1e-6)
|
|
||||||
self.ln_3 = norm_layer(width, elementwise_affine = True, eps = 1e-6)
|
|
||||||
|
|
||||||
self.mlp = MLP(width=width, ratio = mlp_expand_ratio)
|
|
||||||
|
|
||||||
def forward(self, x: torch.Tensor, data: torch.Tensor):
|
|
||||||
x = x + self.attn(self.ln_1(x), self.ln_2(data))
|
|
||||||
x = x + self.mlp(self.ln_3(x))
|
|
||||||
return x
|
|
||||||
|
|
||||||
class CrossAttentionDecoder(nn.Module):
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
num_latents: int,
|
|
||||||
out_channels: int,
|
|
||||||
fourier_embedder,
|
|
||||||
width: int,
|
|
||||||
heads: int,
|
|
||||||
mlp_expand_ratio: int = 4,
|
|
||||||
downsample_ratio: int = 1,
|
|
||||||
enable_ln_post: bool = True,
|
|
||||||
qkv_bias: bool = False,
|
|
||||||
qk_norm: bool = False):
|
|
||||||
|
|
||||||
super().__init__()
|
|
||||||
|
|
||||||
self.enable_ln_post = enable_ln_post
|
|
||||||
self.fourier_embedder = fourier_embedder
|
|
||||||
self.downsample_ratio = downsample_ratio
|
|
||||||
|
|
||||||
self.query_proj = nn.Linear(self.fourier_embedder.out_dim, width)
|
|
||||||
|
|
||||||
if self.downsample_ratio != 1:
|
|
||||||
self.latents_proj = nn.Linear(width * downsample_ratio, width)
|
|
||||||
|
|
||||||
if self.enable_ln_post == False:
|
|
||||||
qk_norm = False
|
|
||||||
|
|
||||||
self.cross_attn_decoder = ResidualCrossAttentionBlock(
|
|
||||||
n_data=num_latents,
|
|
||||||
width=width,
|
|
||||||
mlp_expand_ratio=mlp_expand_ratio,
|
|
||||||
heads=heads,
|
|
||||||
qkv_bias=qkv_bias,
|
|
||||||
qk_norm=qk_norm
|
|
||||||
)
|
|
||||||
|
|
||||||
if self.enable_ln_post:
|
|
||||||
self.ln_post = nn.LayerNorm(width)
|
|
||||||
|
|
||||||
self.output_proj = nn.Linear(width, out_channels)
|
|
||||||
self.count = 0
|
|
||||||
|
|
||||||
def forward(self, queries = None, query_embeddings = None, latents = None):
|
|
||||||
|
|
||||||
if query_embeddings is None:
|
|
||||||
query_embeddings = self.query_proj(self.fourier_embedder(queries).to(latents.dtype))
|
|
||||||
|
|
||||||
self.count += query_embeddings.shape[1]
|
|
||||||
|
|
||||||
if self.downsample_ratio != 1:
|
|
||||||
latents = self.latents_proj(latents)
|
|
||||||
|
|
||||||
x = self.cross_attn_decoder(query_embeddings, latents)
|
|
||||||
|
|
||||||
if self.enable_ln_post:
|
|
||||||
x = self.ln_post(x)
|
|
||||||
|
|
||||||
out = self.output_proj(x)
|
|
||||||
|
|
||||||
return out
|
|
||||||
|
|
||||||
|
|
||||||
class PointCrossAttention(nn.Module):
|
class PointCrossAttention(nn.Module):
|
||||||
def __init__(self,
|
def __init__(self,
|
||||||
num_latents: int,
|
num_latents: int,
|
||||||
@ -652,40 +324,6 @@ class SufraceExtractor():
|
|||||||
|
|
||||||
return outputs
|
return outputs
|
||||||
|
|
||||||
################################################
|
|
||||||
# Volume Decoder
|
|
||||||
################################################
|
|
||||||
|
|
||||||
class VanillaVolumeDecoder():
|
|
||||||
@torch.no_grad()
|
|
||||||
def __call__(self, latents: torch.Tensor, geo_decoder: callable, octree_res: int, bounds = 1.01,
|
|
||||||
num_chunks: int = 10_000):
|
|
||||||
|
|
||||||
if isinstance(bounds, float):
|
|
||||||
bounds = [-bounds, -bounds, -bounds, bounds, bounds, bounds]
|
|
||||||
|
|
||||||
bbox_min, bbox_max = torch.tensor(bounds[:3]), torch.tensor(bounds[3:])
|
|
||||||
|
|
||||||
x = torch.linspace(bbox_min[0], bbox_max[0], int(octree_res) + 1, dtype = torch.float32)
|
|
||||||
y = torch.linspace(bbox_min[1], bbox_max[1], int(octree_res) + 1, dtype = torch.float32)
|
|
||||||
z = torch.linspace(bbox_min[2], bbox_max[2], int(octree_res) + 1, dtype = torch.float32)
|
|
||||||
|
|
||||||
[xs, ys, zs] = torch.meshgrid(x, y, z, indexing = "ij")
|
|
||||||
xyz = torch.stack((xs, ys, zs), axis=-1).to(latents.device, dtype = latents.dtype).contiguous().reshape(-1, 3)
|
|
||||||
grid_size = [int(octree_res) + 1, int(octree_res) + 1, int(octree_res) + 1]
|
|
||||||
|
|
||||||
batch_logits = []
|
|
||||||
for start in range(0, xyz.shape[0], num_chunks):
|
|
||||||
chunk_queries = xyz[start: start + num_chunks, :]
|
|
||||||
chunk_queries = chunk_queries.unsqueeze(0).repeat(latents.shape[0], 1, 1)
|
|
||||||
logits = geo_decoder(queries = chunk_queries, latents = latents)
|
|
||||||
batch_logits.append(logits)
|
|
||||||
|
|
||||||
grid_logits = torch.cat(batch_logits, dim = 1)
|
|
||||||
grid_logits = grid_logits.view((latents.shape[0], *grid_size)).float()
|
|
||||||
|
|
||||||
return grid_logits
|
|
||||||
|
|
||||||
def normalize_mesh(mesh, scale = 0.9999):
|
def normalize_mesh(mesh, scale = 0.9999):
|
||||||
"""Normalize mesh to fit in [-scale, scale]. Translate mesh so its center is [0,0,0]"""
|
"""Normalize mesh to fit in [-scale, scale]. Translate mesh so its center is [0,0,0]"""
|
||||||
|
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user