mirror of
https://git.datalinker.icu/comfyanonymous/ComfyUI
synced 2026-08-24 11:51:20 +08:00
merged vaes and improved surface net
This commit is contained in:
parent
ee65d6ea41
commit
491a49c828
@ -4,7 +4,9 @@
|
|||||||
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
|
||||||
|
import numpy as np
|
||||||
|
import math
|
||||||
|
from tqdm import tqdm
|
||||||
|
|
||||||
from typing import Optional
|
from typing import Optional
|
||||||
|
|
||||||
@ -13,6 +15,457 @@ import logging
|
|||||||
import comfy.ops
|
import comfy.ops
|
||||||
ops = comfy.ops.disable_weight_init
|
ops = comfy.ops.disable_weight_init
|
||||||
|
|
||||||
|
def fps(src: torch.Tensor, batch: torch.Tensor, sampling_ratio: float, start_random: bool = True):
|
||||||
|
|
||||||
|
# manually create the pointer vector
|
||||||
|
assert src.size(0) == batch.numel()
|
||||||
|
|
||||||
|
batch_size = int(batch.max()) + 1
|
||||||
|
deg = src.new_zeros(batch_size, dtype = torch.long)
|
||||||
|
|
||||||
|
deg.scatter_add_(0, batch, torch.ones_like(batch))
|
||||||
|
|
||||||
|
ptr_vec = deg.new_zeros(batch_size + 1)
|
||||||
|
torch.cumsum(deg, 0, out=ptr_vec[1:])
|
||||||
|
|
||||||
|
#return fps_sampling(src, ptr_vec, ratio)
|
||||||
|
sampled_indicies = []
|
||||||
|
|
||||||
|
for b in range(batch_size):
|
||||||
|
# start and the end of each batch
|
||||||
|
start, end = ptr_vec[b].item(), ptr_vec[b + 1].item()
|
||||||
|
# points from the point cloud
|
||||||
|
points = src[start:end]
|
||||||
|
|
||||||
|
num_points = points.size(0)
|
||||||
|
num_samples = max(1, math.ceil(num_points * sampling_ratio))
|
||||||
|
|
||||||
|
selected = torch.zeros(num_samples, device = src.device, dtype = torch.long)
|
||||||
|
distances = torch.full((num_points,), float("inf"), device = src.device)
|
||||||
|
|
||||||
|
# select a random start point
|
||||||
|
if start_random:
|
||||||
|
farthest = torch.randint(0, num_points, (1,), device = src.device)
|
||||||
|
else: farthest = torch.tensor([0], device = src.device, dtype = torch.long)
|
||||||
|
|
||||||
|
for i in range(num_samples):
|
||||||
|
selected[i] = farthest
|
||||||
|
centroid = points[farthest].squeeze(0)
|
||||||
|
dist = torch.norm(points - centroid, dim = 1) # compute euclidean distance
|
||||||
|
distances = torch.minimum(distances, dist)
|
||||||
|
farthest = torch.argmax(distances)
|
||||||
|
|
||||||
|
sampled_indicies.append(torch.arange(start, end)[selected])
|
||||||
|
|
||||||
|
return torch.cat(sampled_indicies, dim = 0)
|
||||||
|
class PointCrossAttention(nn.Module):
|
||||||
|
def __init__(self,
|
||||||
|
num_latents: int,
|
||||||
|
downsample_ratio: float,
|
||||||
|
pc_size: int,
|
||||||
|
pc_sharpedge_size: int,
|
||||||
|
point_feats: int,
|
||||||
|
width: int,
|
||||||
|
heads: int,
|
||||||
|
layers: int,
|
||||||
|
fourier_embedder,
|
||||||
|
normal_pe: bool = False,
|
||||||
|
qkv_bias: bool = False,
|
||||||
|
use_ln_post: bool = True,
|
||||||
|
qk_norm: bool = True):
|
||||||
|
|
||||||
|
super().__init__()
|
||||||
|
|
||||||
|
self.fourier_embedder = fourier_embedder
|
||||||
|
|
||||||
|
self.pc_size = pc_size
|
||||||
|
self.normal_pe = normal_pe
|
||||||
|
self.downsample_ratio = downsample_ratio
|
||||||
|
self.pc_sharpedge_size = pc_sharpedge_size
|
||||||
|
self.num_latents = num_latents
|
||||||
|
self.point_feats = point_feats
|
||||||
|
|
||||||
|
self.input_proj = nn.Linear(self.fourier_embedder.out_dim + point_feats, width)
|
||||||
|
|
||||||
|
self.cross_attn = ResidualCrossAttentionBlock(
|
||||||
|
width = width,
|
||||||
|
heads = heads,
|
||||||
|
qkv_bias = qkv_bias,
|
||||||
|
qk_norm = qk_norm
|
||||||
|
)
|
||||||
|
|
||||||
|
self.self_attn = None
|
||||||
|
if layers > 0:
|
||||||
|
self.self_attn = Transformer(
|
||||||
|
width = width,
|
||||||
|
heads = heads,
|
||||||
|
qkv_bias = qkv_bias,
|
||||||
|
qk_norm = qk_norm,
|
||||||
|
layers = layers
|
||||||
|
)
|
||||||
|
|
||||||
|
if use_ln_post:
|
||||||
|
self.ln_post = nn.LayerNorm(width)
|
||||||
|
else:
|
||||||
|
self.ln_post = None
|
||||||
|
|
||||||
|
def sample_points_and_latents(self, point_cloud: torch.Tensor, features: torch.Tensor):
|
||||||
|
|
||||||
|
"""
|
||||||
|
Subsample points randomly from the point cloud (input_pc)
|
||||||
|
Further sample the subsampled points to get query_pc
|
||||||
|
take the fourier embeddings for both input and query pc
|
||||||
|
|
||||||
|
Mental Note: FPS-sampled points (query_pc) act as latent tokens that attend to and learn from the broader context in input_pc.
|
||||||
|
Goal: get a smaller represenation (query_pc) to represent the entire scence structure by learning from a broader subset (input_pc).
|
||||||
|
More computationally efficient.
|
||||||
|
|
||||||
|
Features are additional information for each point in the cloud
|
||||||
|
"""
|
||||||
|
|
||||||
|
B, _, D = point_cloud.shape
|
||||||
|
|
||||||
|
num_latents = int(self.num_latents)
|
||||||
|
|
||||||
|
num_random_query = self.pc_size / (self.pc_size + self.pc_sharpedge_size) * num_latents
|
||||||
|
num_sharpedge_query = num_latents - num_random_query
|
||||||
|
|
||||||
|
# Split random and sharpedge surface points
|
||||||
|
random_pc, sharpedge_pc = torch.split(point_cloud, [self.pc_size, self.pc_sharpedge_size], dim=1)
|
||||||
|
|
||||||
|
# assert statements
|
||||||
|
assert random_pc.shape[1] <= self.pc_size, "Random surface points size must be less than or equal to pc_size"
|
||||||
|
assert sharpedge_pc.shape[1] <= self.pc_sharpedge_size, "Sharpedge surface points size must be less than or equal to pc_sharpedge_size"
|
||||||
|
|
||||||
|
input_random_pc_size = int(num_random_query * self.downsample_ratio)
|
||||||
|
random_query_pc, random_input_pc, random_idx_pc, random_idx_query = \
|
||||||
|
self.subsample(pc = random_pc, num_query = num_random_query, input_pc_size = input_random_pc_size)
|
||||||
|
|
||||||
|
input_sharpedge_pc_size = int(num_sharpedge_query * self.downsample_ratio)
|
||||||
|
|
||||||
|
if input_sharpedge_pc_size == 0:
|
||||||
|
sharpedge_input_pc = torch.zeros(B, 0, D, dtype = random_input_pc.dtype).to(point_cloud.device)
|
||||||
|
sharpedge_query_pc = torch.zeros(B, 0, D, dtype= random_query_pc.dtype).to(point_cloud.device)
|
||||||
|
|
||||||
|
else: sharpedge_query_pc, sharpedge_input_pc, sharpedge_idx_pc, sharpedge_idx_query = \
|
||||||
|
self.subsample(pc = sharpedge_pc, num_query = num_sharpedge_query, input_pc_size = input_sharpedge_pc_size)
|
||||||
|
|
||||||
|
# concat the random and sharpedges
|
||||||
|
query_pc = torch.cat([random_query_pc, sharpedge_query_pc], dim = 1)
|
||||||
|
input_pc = torch.cat([random_input_pc, sharpedge_input_pc], dim = 1)
|
||||||
|
|
||||||
|
query = self.fourier_embedder(query_pc)
|
||||||
|
data = self.fourier_embedder(input_pc)
|
||||||
|
|
||||||
|
if self.point_feats > 0:
|
||||||
|
random_surface_features, sharpedge_surface_features = torch.split(features, [self.pc_size, self.pc_sharpedge_size], dim = 1)
|
||||||
|
|
||||||
|
input_random_surface_features, query_random_features = \
|
||||||
|
self.handle_features(features = random_surface_features, idx_pc = random_idx_pc, batch_size = B,
|
||||||
|
input_pc_size = input_random_pc_size, idx_query = random_idx_query)
|
||||||
|
|
||||||
|
if input_sharpedge_pc_size == 0:
|
||||||
|
input_sharpedge_surface_features = torch.zeros(B, 0, self.point_feats,
|
||||||
|
dtype = input_random_surface_features.dtype, device = point_cloud.device)
|
||||||
|
|
||||||
|
query_sharpedge_features = torch.zeros(B, 0, self.point_feats,
|
||||||
|
dtype = query_random_features.dtype, device = point_cloud.device)
|
||||||
|
else:
|
||||||
|
|
||||||
|
input_sharpedge_surface_features, query_sharpedge_features = \
|
||||||
|
self.handle_features(idx_pc = sharpedge_idx_pc, features = sharpedge_surface_features,
|
||||||
|
batch_size = B, idx_query = sharpedge_idx_query, input_pc_size = input_sharpedge_pc_size)
|
||||||
|
|
||||||
|
query_features = torch.cat([query_random_features, query_sharpedge_features], dim = 1)
|
||||||
|
input_features = torch.cat([input_random_surface_features, input_sharpedge_surface_features], dim = 1)
|
||||||
|
|
||||||
|
if self.normal_pe:
|
||||||
|
# apply the fourier embeddings on the first 3 dims (xyz)
|
||||||
|
input_features_pe = self.fourier_embedder(input_features[..., :3])
|
||||||
|
query_features_pe = self.fourier_embedder(query_features[..., :3])
|
||||||
|
# replace the first 3 dims with the new PE ones
|
||||||
|
input_features = torch.cat([input_features_pe, input_features[..., :3]], dim = -1)
|
||||||
|
query_features = torch.cat([query_features_pe, query_features[..., :3]], dim = -1)
|
||||||
|
|
||||||
|
# concat at the channels dim
|
||||||
|
query = torch.cat([query, query_features], dim = -1)
|
||||||
|
data = torch.cat([data, input_features], dim = -1)
|
||||||
|
|
||||||
|
# don't return pc_info to avoid unnecessary memory usuage
|
||||||
|
return query.view(B, -1, query.shape[-1]), data.view(B, -1, data.shape[-1])
|
||||||
|
|
||||||
|
def forward(self, point_cloud: torch.Tensor, features: torch.Tensor):
|
||||||
|
|
||||||
|
query, data = self.sample_points_and_latents(point_cloud = point_cloud, features = features)
|
||||||
|
|
||||||
|
# apply projections
|
||||||
|
query = self.input_proj(query)
|
||||||
|
data = self.input_proj(data)
|
||||||
|
|
||||||
|
# apply cross attention between query and data
|
||||||
|
latents = self.cross_attn(query, data)
|
||||||
|
|
||||||
|
if self.self_attn is not None:
|
||||||
|
latents = self.self_attn(latents)
|
||||||
|
|
||||||
|
if self.ln_post is not None:
|
||||||
|
latents = self.ln_post(latents)
|
||||||
|
|
||||||
|
return latents
|
||||||
|
|
||||||
|
|
||||||
|
def subsample(self, pc, num_query, input_pc_size: int):
|
||||||
|
|
||||||
|
"""
|
||||||
|
num_query: number of points to keep after FPS
|
||||||
|
input_pc_size: number of points to select before FPS
|
||||||
|
"""
|
||||||
|
|
||||||
|
B, _, D = pc.shape
|
||||||
|
query_ratio = num_query / input_pc_size
|
||||||
|
|
||||||
|
# random subsampling of points inside the point cloud
|
||||||
|
idx_pc = torch.randperm(pc.shape[1], device = pc.device)[:input_pc_size]
|
||||||
|
input_pc = pc[:, idx_pc, :]
|
||||||
|
|
||||||
|
# flatten to allow applying fps across the whole batch
|
||||||
|
flattent_input_pc = input_pc.view(B * input_pc_size, D)
|
||||||
|
|
||||||
|
# construct a batch_down tensor to tell fps
|
||||||
|
# which points belong to which batch
|
||||||
|
N_down = int(flattent_input_pc.shape[0] / B)
|
||||||
|
batch_down = torch.arange(B).to(pc.device)
|
||||||
|
batch_down = torch.repeat_interleave(batch_down, N_down)
|
||||||
|
|
||||||
|
idx_query = fps(flattent_input_pc, batch_down, sampling_ratio = query_ratio)
|
||||||
|
query_pc = flattent_input_pc[idx_query].view(B, -1, D)
|
||||||
|
|
||||||
|
return query_pc, input_pc, idx_pc, idx_query
|
||||||
|
|
||||||
|
def handle_features(self, features, idx_pc, input_pc_size, batch_size: int, idx_query):
|
||||||
|
|
||||||
|
B = batch_size
|
||||||
|
|
||||||
|
input_surface_features = features[:, idx_pc, :]
|
||||||
|
flattent_input_features = input_surface_features.view(B * input_pc_size, -1)
|
||||||
|
query_features = flattent_input_features[idx_query].view(B, -1,
|
||||||
|
flattent_input_features.shape[-1])
|
||||||
|
|
||||||
|
return input_surface_features, query_features
|
||||||
|
|
||||||
|
def forward(self, pc, feats):
|
||||||
|
"""
|
||||||
|
|
||||||
|
Args:
|
||||||
|
pc (torch.FloatTensor): [B, N, 3]
|
||||||
|
feats (torch.FloatTensor or None): [B, N, C]
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
|
||||||
|
"""
|
||||||
|
|
||||||
|
query, data = self.sample_points_and_latents(pc, feats)
|
||||||
|
|
||||||
|
query = self.input_proj(query)
|
||||||
|
query = query
|
||||||
|
data = self.input_proj(data)
|
||||||
|
data = data
|
||||||
|
|
||||||
|
latents = self.cross_attn(query, data)
|
||||||
|
if self.self_attn is not None:
|
||||||
|
latents = self.self_attn(latents)
|
||||||
|
|
||||||
|
if self.ln_post is not None:
|
||||||
|
latents = self.ln_post(latents)
|
||||||
|
|
||||||
|
return latents
|
||||||
|
|
||||||
|
def normalize_mesh(mesh, scale = 0.9999):
|
||||||
|
"""Normalize mesh to fit in [-scale, scale]. Translate mesh so its center is [0,0,0]"""
|
||||||
|
|
||||||
|
bbox = mesh.bounds
|
||||||
|
center = (bbox[1] + bbox[0]) / 2
|
||||||
|
|
||||||
|
max_extent = (bbox[1] - bbox[0]).max()
|
||||||
|
mesh.apply_translation(-center)
|
||||||
|
mesh.apply_scale((2 * scale) / max_extent)
|
||||||
|
|
||||||
|
return mesh
|
||||||
|
|
||||||
|
def sample_pointcloud(mesh, num = 200000):
|
||||||
|
""" Uniformly sample points from the surface of the mesh """
|
||||||
|
|
||||||
|
points, face_idx = mesh.sample(num, return_index = True)
|
||||||
|
normals = mesh.face_normals[face_idx]
|
||||||
|
return torch.from_numpy(points.astype(np.float32)), torch.from_numpy(normals.astype(np.float32))
|
||||||
|
|
||||||
|
def detect_sharp_edges(mesh, threshold=0.985):
|
||||||
|
"""Return edge indices (a, b) that lie on sharp boundaries of the mesh."""
|
||||||
|
|
||||||
|
V, F = mesh.vertices, mesh.faces
|
||||||
|
VN, FN = mesh.vertex_normals, mesh.face_normals
|
||||||
|
|
||||||
|
sharp_mask = np.ones(V.shape[0])
|
||||||
|
for i in range(3):
|
||||||
|
indices = F[:, i]
|
||||||
|
alignment = np.einsum('ij,ij->i', VN[indices], FN)
|
||||||
|
dot_stack = np.stack((sharp_mask[indices], alignment), axis=-1)
|
||||||
|
sharp_mask[indices] = np.min(dot_stack, axis=-1)
|
||||||
|
|
||||||
|
edge_a = np.concatenate([F[:, 0], F[:, 1], F[:, 2]])
|
||||||
|
edge_b = np.concatenate([F[:, 1], F[:, 2], F[:, 0]])
|
||||||
|
sharp_edges = (sharp_mask[edge_a] < threshold) & (sharp_mask[edge_b] < threshold)
|
||||||
|
|
||||||
|
return edge_a[sharp_edges], edge_b[sharp_edges]
|
||||||
|
|
||||||
|
|
||||||
|
def sharp_sample_pointcloud(mesh, num = 16384):
|
||||||
|
""" Sample points preferentially from sharp edges in the mesh. """
|
||||||
|
|
||||||
|
edge_a, edge_b = detect_sharp_edges(mesh)
|
||||||
|
V, VN = mesh.vertices, mesh.vertex_normals
|
||||||
|
|
||||||
|
va, vb = V[edge_a], V[edge_b]
|
||||||
|
na, nb = VN[edge_a], VN[edge_b]
|
||||||
|
|
||||||
|
edge_lengths = np.linalg.norm(vb - va, axis=-1)
|
||||||
|
weights = edge_lengths / edge_lengths.sum()
|
||||||
|
|
||||||
|
indices = np.searchsorted(np.cumsum(weights), np.random.rand(num))
|
||||||
|
t = np.random.rand(num, 1)
|
||||||
|
|
||||||
|
samples = t * va[indices] + (1 - t) * vb[indices]
|
||||||
|
normals = t * na[indices] + (1 - t) * nb[indices]
|
||||||
|
|
||||||
|
return samples.astype(np.float32), normals.astype(np.float32)
|
||||||
|
|
||||||
|
def load_surface_sharpedge(mesh, num_points=4096, num_sharp_points=4096, sharpedge_flag = True, device = "cuda"):
|
||||||
|
"""Load a surface with optional sharp-edge annotations from a trimesh mesh."""
|
||||||
|
|
||||||
|
import trimesh
|
||||||
|
|
||||||
|
try:
|
||||||
|
mesh_full = trimesh.util.concatenate(mesh.dump())
|
||||||
|
except Exception:
|
||||||
|
mesh_full = trimesh.util.concatenate(mesh)
|
||||||
|
|
||||||
|
mesh_full = normalize_mesh(mesh_full)
|
||||||
|
|
||||||
|
faces = mesh_full.faces
|
||||||
|
vertices = mesh_full.vertices
|
||||||
|
origin_face_count = faces.shape[0]
|
||||||
|
|
||||||
|
mesh_surface = trimesh.Trimesh(vertices=vertices, faces=faces[:origin_face_count])
|
||||||
|
mesh_fill = trimesh.Trimesh(vertices=vertices, faces=faces[origin_face_count:])
|
||||||
|
|
||||||
|
area_surface = mesh_surface.area
|
||||||
|
area_fill = mesh_fill.area
|
||||||
|
total_area = area_surface + area_fill
|
||||||
|
|
||||||
|
sample_num = 499712 // 2
|
||||||
|
fill_ratio = area_fill / total_area if total_area > 0 else 0
|
||||||
|
|
||||||
|
num_fill = int(sample_num * fill_ratio)
|
||||||
|
num_surface = sample_num - num_fill
|
||||||
|
|
||||||
|
surf_pts, surf_normals = sample_pointcloud(mesh_surface, num_surface)
|
||||||
|
fill_pts, fill_normals = (torch.zeros(0, 3), torch.zeros(0, 3)) if num_fill == 0 else sample_pointcloud(mesh_fill, num_fill)
|
||||||
|
|
||||||
|
sharp_pts, sharp_normals = sharp_sample_pointcloud(mesh_surface, sample_num)
|
||||||
|
|
||||||
|
def assemble_tensor(points, normals, label=None):
|
||||||
|
|
||||||
|
data = torch.cat([points, normals], dim=1).half().to(device)
|
||||||
|
|
||||||
|
if label is not None:
|
||||||
|
label_tensor = torch.full((data.shape[0], 1), float(label), dtype=torch.float16).to(device)
|
||||||
|
data = torch.cat([data, label_tensor], dim=1)
|
||||||
|
|
||||||
|
return data
|
||||||
|
|
||||||
|
surface = assemble_tensor(torch.cat([surf_pts.to(device), fill_pts.to(device)], dim=0),
|
||||||
|
torch.cat([surf_normals.to(device), fill_normals.to(device)], dim=0),
|
||||||
|
label = 0 if sharpedge_flag else None)
|
||||||
|
|
||||||
|
sharp_surface = assemble_tensor(torch.from_numpy(sharp_pts), torch.from_numpy(sharp_normals),
|
||||||
|
label = 1 if sharpedge_flag else None)
|
||||||
|
|
||||||
|
rng = np.random.default_rng()
|
||||||
|
|
||||||
|
surface = surface[rng.choice(surface.shape[0], num_points, replace = False)]
|
||||||
|
sharp_surface = sharp_surface[rng.choice(sharp_surface.shape[0], num_sharp_points, replace = False)]
|
||||||
|
|
||||||
|
full = torch.cat([surface, sharp_surface], dim = 0).unsqueeze(0)
|
||||||
|
|
||||||
|
return full
|
||||||
|
|
||||||
|
class SharpEdgeSurfaceLoader:
|
||||||
|
""" Load mesh surface and sharp edge samples. """
|
||||||
|
|
||||||
|
def __init__(self, num_uniform_points = 8192, num_sharp_points = 8192):
|
||||||
|
|
||||||
|
self.num_uniform_points = num_uniform_points
|
||||||
|
self.num_sharp_points = num_sharp_points
|
||||||
|
self.total_points = num_uniform_points + num_sharp_points
|
||||||
|
|
||||||
|
def __call__(self, mesh_input, device = "cuda"):
|
||||||
|
mesh = self._load_mesh(mesh_input)
|
||||||
|
return load_surface_sharpedge(mesh, self.num_uniform_points, self.num_sharp_points, device = device)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _load_mesh(mesh_input):
|
||||||
|
import trimesh
|
||||||
|
|
||||||
|
if isinstance(mesh_input, str):
|
||||||
|
mesh = trimesh.load(mesh_input, force="mesh", merge_primitives = True)
|
||||||
|
else:
|
||||||
|
mesh = mesh_input
|
||||||
|
|
||||||
|
if isinstance(mesh, trimesh.Scene):
|
||||||
|
combined = None
|
||||||
|
for obj in mesh.geometry.values():
|
||||||
|
combined = obj if combined is None else combined + obj
|
||||||
|
return combined
|
||||||
|
|
||||||
|
return mesh
|
||||||
|
|
||||||
|
class FourierEmbedder(nn.Module):
|
||||||
|
def __init__(self, num_freq: int = 8, input_dim: int = 3, include_pi: bool = False):
|
||||||
|
super().__init__()
|
||||||
|
|
||||||
|
frequencies = 2.0 ** torch.arange(
|
||||||
|
num_freq,
|
||||||
|
dtype = torch.float32
|
||||||
|
)
|
||||||
|
|
||||||
|
if include_pi:
|
||||||
|
frequencies *= torch.pi
|
||||||
|
|
||||||
|
self.register_buffer("frequencies", frequencies, persistent = False)
|
||||||
|
|
||||||
|
self.out_dim = input_dim * (num_freq * 2 + 1)
|
||||||
|
|
||||||
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
|
||||||
|
embed = (x[..., None].contiguous() * self.frequencies).view(*x.shape[:-1], -1)
|
||||||
|
return torch.cat((x, embed.sin(), embed.cos()), dim = -1)
|
||||||
|
|
||||||
|
class DiagonalGaussianDistribution:
|
||||||
|
def __init__(self, params: torch.Tensor, feature_dim: int = -1):
|
||||||
|
|
||||||
|
# divide quant channels (8) into mean and log variance
|
||||||
|
self.mean, self.logvar = torch.chunk(params, 2, dim = feature_dim)
|
||||||
|
|
||||||
|
self.logvar = torch.clamp(self.logvar, -30.0, 20.0)
|
||||||
|
self.std = torch.exp(0.5 * self.logvar)
|
||||||
|
|
||||||
|
def sample(self):
|
||||||
|
|
||||||
|
eps = torch.randn_like(self.std)
|
||||||
|
z = self.mean + eps * self.std
|
||||||
|
|
||||||
|
return z
|
||||||
|
|
||||||
################################################
|
################################################
|
||||||
# Volume Decoder
|
# Volume Decoder
|
||||||
################################################
|
################################################
|
||||||
@ -20,7 +473,7 @@ ops = comfy.ops.disable_weight_init
|
|||||||
class VanillaVolumeDecoder():
|
class VanillaVolumeDecoder():
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def __call__(self, latents: torch.Tensor, geo_decoder: callable, octree_resolution: int, bounds = 1.01,
|
def __call__(self, latents: torch.Tensor, geo_decoder: callable, octree_resolution: int, bounds = 1.01,
|
||||||
num_chunks: int = 10_000):
|
num_chunks: int = 10_000, enable_pbar: bool = True, **kwargs):
|
||||||
|
|
||||||
if isinstance(bounds, float):
|
if isinstance(bounds, float):
|
||||||
bounds = [-bounds, -bounds, -bounds, bounds, bounds, bounds]
|
bounds = [-bounds, -bounds, -bounds, bounds, bounds, bounds]
|
||||||
@ -36,7 +489,9 @@ class VanillaVolumeDecoder():
|
|||||||
grid_size = [int(octree_resolution) + 1, int(octree_resolution) + 1, int(octree_resolution) + 1]
|
grid_size = [int(octree_resolution) + 1, int(octree_resolution) + 1, int(octree_resolution) + 1]
|
||||||
|
|
||||||
batch_logits = []
|
batch_logits = []
|
||||||
for start in range(0, xyz.shape[0], num_chunks):
|
for start in tqdm(range(0, xyz.shape[0], num_chunks), desc=f"Volume Decoding",
|
||||||
|
disable=not enable_pbar):
|
||||||
|
|
||||||
chunk_queries = xyz[start: start + num_chunks, :]
|
chunk_queries = xyz[start: start + num_chunks, :]
|
||||||
chunk_queries = chunk_queries.unsqueeze(0).repeat(latents.shape[0], 1, 1)
|
chunk_queries = chunk_queries.unsqueeze(0).repeat(latents.shape[0], 1, 1)
|
||||||
logits = geo_decoder(queries = chunk_queries, latents = latents)
|
logits = geo_decoder(queries = chunk_queries, latents = latents)
|
||||||
@ -484,28 +939,44 @@ class CrossAttentionDecoder(nn.Module):
|
|||||||
|
|
||||||
class ShapeVAE(nn.Module):
|
class ShapeVAE(nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
embed_dim: int,
|
num_latents: int = 4096,
|
||||||
width: int,
|
embed_dim: int = 64,
|
||||||
heads: int,
|
width: int = 1024,
|
||||||
num_decoder_layers: int,
|
heads: int = 16,
|
||||||
geo_decoder_downsample_ratio: int = 1,
|
num_decoder_layers: int = 16,
|
||||||
geo_decoder_mlp_expand_ratio: int = 4,
|
num_encoder_layers: int = 8,
|
||||||
geo_decoder_ln_post: bool = True,
|
pc_size: int = 81920,
|
||||||
num_freqs: int = 8,
|
pc_sharpedge_size: int = 0,
|
||||||
include_pi: bool = True,
|
point_feats: int = 4,
|
||||||
qkv_bias: bool = True,
|
downsample_ratio: int = 20,
|
||||||
qk_norm: bool = False,
|
geo_decoder_downsample_ratio: int = 1,
|
||||||
label_type: str = "binary",
|
geo_decoder_mlp_expand_ratio: int = 4,
|
||||||
drop_path_rate: float = 0.0,
|
geo_decoder_ln_post: bool = True,
|
||||||
scale_factor: float = 1.0,
|
num_freqs: int = 8,
|
||||||
|
qkv_bias: bool = False,
|
||||||
|
qk_norm: bool = True,
|
||||||
|
drop_path_rate: float = 0.0,
|
||||||
|
include_pi: bool = False,
|
||||||
|
scale_factor: float = 1.0039506158752403,
|
||||||
|
label_type: str = "binary",
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.geo_decoder_ln_post = geo_decoder_ln_post
|
self.geo_decoder_ln_post = geo_decoder_ln_post
|
||||||
|
|
||||||
self.fourier_embedder = FourierEmbedder(num_freqs=num_freqs, include_pi=include_pi)
|
self.fourier_embedder = FourierEmbedder(num_freqs=num_freqs, include_pi=include_pi)
|
||||||
|
|
||||||
|
self.encoder = PointCrossAttention(layers = num_encoder_layers,
|
||||||
|
num_latents = num_latents,
|
||||||
|
downsample_ratio = downsample_ratio,
|
||||||
|
heads = heads,
|
||||||
|
pc_size = pc_size,
|
||||||
|
width = width,
|
||||||
|
point_feats = point_feats,
|
||||||
|
fourier_embedder = self.fourier_embedder,
|
||||||
|
pc_sharpedge_size = pc_sharpedge_size)
|
||||||
|
|
||||||
self.post_kl = ops.Linear(embed_dim, width)
|
self.post_kl = ops.Linear(embed_dim, width)
|
||||||
|
|
||||||
self.transformer = Transformer(
|
self.transformer = Transformer(
|
||||||
@ -534,7 +1005,7 @@ class ShapeVAE(nn.Module):
|
|||||||
self.scale_factor = scale_factor
|
self.scale_factor = scale_factor
|
||||||
|
|
||||||
def decode(self, latents, **kwargs):
|
def decode(self, latents, **kwargs):
|
||||||
latents = self.post_kl(latents.movedim(-2, -1))
|
latents = self.post_kl(latents)
|
||||||
latents = self.transformer(latents)
|
latents = self.transformer(latents)
|
||||||
|
|
||||||
bounds = kwargs.get("bounds", 1.01)
|
bounds = kwargs.get("bounds", 1.01)
|
||||||
@ -543,7 +1014,16 @@ class ShapeVAE(nn.Module):
|
|||||||
enable_pbar = kwargs.get("enable_pbar", True)
|
enable_pbar = kwargs.get("enable_pbar", True)
|
||||||
|
|
||||||
grid_logits = self.volume_decoder(latents, self.geo_decoder, bounds=bounds, num_chunks=num_chunks, octree_resolution=octree_resolution, enable_pbar=enable_pbar)
|
grid_logits = self.volume_decoder(latents, self.geo_decoder, bounds=bounds, num_chunks=num_chunks, octree_resolution=octree_resolution, enable_pbar=enable_pbar)
|
||||||
return grid_logits.movedim(-2, -1)
|
return grid_logits
|
||||||
|
|
||||||
def encode(self, x):
|
def encode(self, surface):
|
||||||
return None
|
|
||||||
|
pc, feats = surface[:, :, :3], surface[:, :, 3:]
|
||||||
|
latents = self.encoder(pc, feats)
|
||||||
|
|
||||||
|
moments = self.pre_kl(latents)
|
||||||
|
posterior = DiagonalGaussianDistribution(moments, feature_dim = -1)
|
||||||
|
|
||||||
|
latents = posterior.sample()
|
||||||
|
|
||||||
|
return latents
|
||||||
@ -1,629 +0,0 @@
|
|||||||
import torch
|
|
||||||
from torch import Tensor
|
|
||||||
import math
|
|
||||||
import numpy as np
|
|
||||||
from skimage import measure
|
|
||||||
from dataclasses import dataclass
|
|
||||||
import torch.nn as nn
|
|
||||||
|
|
||||||
import sys, os;
|
|
||||||
sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), "../../..")))
|
|
||||||
|
|
||||||
from comfy.ldm.hunyuan3d.vae import (
|
|
||||||
CrossAttentionDecoder, Transformer, ResidualCrossAttentionBlock, FourierEmbedder, VanillaVolumeDecoder
|
|
||||||
)
|
|
||||||
def fps(src: Tensor, batch: Tensor, sampling_ratio: float, start_random: bool = True):
|
|
||||||
|
|
||||||
# manually create the pointer vector
|
|
||||||
assert src.size(0) == batch.numel()
|
|
||||||
|
|
||||||
batch_size = int(batch.max()) + 1
|
|
||||||
deg = src.new_zeros(batch_size, dtype = torch.long)
|
|
||||||
|
|
||||||
deg.scatter_add_(0, batch, torch.ones_like(batch))
|
|
||||||
|
|
||||||
ptr_vec = deg.new_zeros(batch_size + 1)
|
|
||||||
torch.cumsum(deg, 0, out=ptr_vec[1:])
|
|
||||||
|
|
||||||
#return fps_sampling(src, ptr_vec, ratio)
|
|
||||||
sampled_indicies = []
|
|
||||||
|
|
||||||
for b in range(batch_size):
|
|
||||||
# start and the end of each batch
|
|
||||||
start, end = ptr_vec[b].item(), ptr_vec[b + 1].item()
|
|
||||||
# points from the point cloud
|
|
||||||
points = src[start:end]
|
|
||||||
|
|
||||||
num_points = points.size(0)
|
|
||||||
num_samples = max(1, math.ceil(num_points * sampling_ratio))
|
|
||||||
|
|
||||||
selected = torch.zeros(num_samples, device = src.device, dtype = torch.long)
|
|
||||||
distances = torch.full((num_points,), float("inf"), device = src.device)
|
|
||||||
|
|
||||||
# select a random start point
|
|
||||||
if start_random:
|
|
||||||
farthest = torch.randint(0, num_points, (1,), device = src.device)
|
|
||||||
else: farthest = torch.tensor([0], device = src.device, dtype = torch.long)
|
|
||||||
|
|
||||||
for i in range(num_samples):
|
|
||||||
selected[i] = farthest
|
|
||||||
centroid = points[farthest].squeeze(0)
|
|
||||||
dist = torch.norm(points - centroid, dim = 1) # compute euclidean distance
|
|
||||||
distances = torch.minimum(distances, dist)
|
|
||||||
farthest = torch.argmax(distances)
|
|
||||||
|
|
||||||
sampled_indicies.append(torch.arange(start, end)[selected])
|
|
||||||
|
|
||||||
return torch.cat(sampled_indicies, dim = 0)
|
|
||||||
class PointCrossAttention(nn.Module):
|
|
||||||
def __init__(self,
|
|
||||||
num_latents: int,
|
|
||||||
downsample_ratio: float,
|
|
||||||
pc_size: int,
|
|
||||||
pc_sharpedge_size: int,
|
|
||||||
point_feats: int,
|
|
||||||
width: int,
|
|
||||||
heads: int,
|
|
||||||
layers: int,
|
|
||||||
fourier_embedder,
|
|
||||||
normal_pe: bool = False,
|
|
||||||
qkv_bias: bool = False,
|
|
||||||
use_ln_post: bool = True,
|
|
||||||
qk_norm: bool = True):
|
|
||||||
|
|
||||||
super().__init__()
|
|
||||||
|
|
||||||
self.fourier_embedder = fourier_embedder
|
|
||||||
|
|
||||||
self.pc_size = pc_size
|
|
||||||
self.normal_pe = normal_pe
|
|
||||||
self.downsample_ratio = downsample_ratio
|
|
||||||
self.pc_sharpedge_size = pc_sharpedge_size
|
|
||||||
self.num_latents = num_latents
|
|
||||||
self.point_feats = point_feats
|
|
||||||
|
|
||||||
self.input_proj = nn.Linear(self.fourier_embedder.out_dim + point_feats, width)
|
|
||||||
|
|
||||||
self.cross_attn = ResidualCrossAttentionBlock(
|
|
||||||
width = width,
|
|
||||||
heads = heads,
|
|
||||||
qkv_bias = qkv_bias,
|
|
||||||
qk_norm = qk_norm
|
|
||||||
)
|
|
||||||
|
|
||||||
self.self_attn = None
|
|
||||||
if layers > 0:
|
|
||||||
self.self_attn = Transformer(
|
|
||||||
width = width,
|
|
||||||
heads = heads,
|
|
||||||
qkv_bias = qkv_bias,
|
|
||||||
qk_norm = qk_norm,
|
|
||||||
layers = layers
|
|
||||||
)
|
|
||||||
|
|
||||||
if use_ln_post:
|
|
||||||
self.ln_post = nn.LayerNorm(width)
|
|
||||||
else:
|
|
||||||
self.ln_post = None
|
|
||||||
|
|
||||||
def sample_points_and_latents(self, point_cloud: torch.Tensor, features: torch.Tensor):
|
|
||||||
|
|
||||||
"""
|
|
||||||
Subsample points randomly from the point cloud (input_pc)
|
|
||||||
Further sample the subsampled points to get query_pc
|
|
||||||
take the fourier embeddings for both input and query pc
|
|
||||||
|
|
||||||
Mental Note: FPS-sampled points (query_pc) act as latent tokens that attend to and learn from the broader context in input_pc.
|
|
||||||
Goal: get a smaller represenation (query_pc) to represent the entire scence structure by learning from a broader subset (input_pc).
|
|
||||||
More computationally efficient.
|
|
||||||
|
|
||||||
Features are additional information for each point in the cloud
|
|
||||||
"""
|
|
||||||
|
|
||||||
B, _, D = point_cloud.shape
|
|
||||||
|
|
||||||
num_latents = int(self.num_latents)
|
|
||||||
|
|
||||||
num_random_query = self.pc_size / (self.pc_size + self.pc_sharpedge_size) * num_latents
|
|
||||||
num_sharpedge_query = num_latents - num_random_query
|
|
||||||
|
|
||||||
# Split random and sharpedge surface points
|
|
||||||
random_pc, sharpedge_pc = torch.split(point_cloud, [self.pc_size, self.pc_sharpedge_size], dim=1)
|
|
||||||
|
|
||||||
# assert statements
|
|
||||||
assert random_pc.shape[1] <= self.pc_size, "Random surface points size must be less than or equal to pc_size"
|
|
||||||
assert sharpedge_pc.shape[1] <= self.pc_sharpedge_size, "Sharpedge surface points size must be less than or equal to pc_sharpedge_size"
|
|
||||||
|
|
||||||
input_random_pc_size = int(num_random_query * self.downsample_ratio)
|
|
||||||
random_query_pc, random_input_pc, random_idx_pc, random_idx_query = \
|
|
||||||
self.subsample(pc = random_pc, num_query = num_random_query, input_pc_size = input_random_pc_size)
|
|
||||||
|
|
||||||
input_sharpedge_pc_size = int(num_sharpedge_query * self.downsample_ratio)
|
|
||||||
|
|
||||||
if input_sharpedge_pc_size == 0:
|
|
||||||
sharpedge_input_pc = torch.zeros(B, 0, D, dtype = random_input_pc.dtype).to(point_cloud.device)
|
|
||||||
sharpedge_query_pc = torch.zeros(B, 0, D, dtype= random_query_pc.dtype).to(point_cloud.device)
|
|
||||||
|
|
||||||
else: sharpedge_query_pc, sharpedge_input_pc, sharpedge_idx_pc, sharpedge_idx_query = \
|
|
||||||
self.subsample(pc = sharpedge_pc, num_query = num_sharpedge_query, input_pc_size = input_sharpedge_pc_size)
|
|
||||||
|
|
||||||
# concat the random and sharpedges
|
|
||||||
query_pc = torch.cat([random_query_pc, sharpedge_query_pc], dim = 1)
|
|
||||||
input_pc = torch.cat([random_input_pc, sharpedge_input_pc], dim = 1)
|
|
||||||
|
|
||||||
query = self.fourier_embedder(query_pc)
|
|
||||||
data = self.fourier_embedder(input_pc)
|
|
||||||
|
|
||||||
if self.point_feats > 0:
|
|
||||||
random_surface_features, sharpedge_surface_features = torch.split(features, [self.pc_size, self.pc_sharpedge_size], dim = 1)
|
|
||||||
|
|
||||||
input_random_surface_features, query_random_features = \
|
|
||||||
self.handle_features(features = random_surface_features, idx_pc = random_idx_pc, batch_size = B,
|
|
||||||
input_pc_size = input_random_pc_size, idx_query = random_idx_query)
|
|
||||||
|
|
||||||
if input_sharpedge_pc_size == 0:
|
|
||||||
input_sharpedge_surface_features = torch.zeros(B, 0, self.point_feats,
|
|
||||||
dtype = input_random_surface_features.dtype, device = point_cloud.device)
|
|
||||||
|
|
||||||
query_sharpedge_features = torch.zeros(B, 0, self.point_feats,
|
|
||||||
dtype = query_random_features.dtype, device = point_cloud.device)
|
|
||||||
else:
|
|
||||||
|
|
||||||
input_sharpedge_surface_features, query_sharpedge_features = \
|
|
||||||
self.handle_features(idx_pc = sharpedge_idx_pc, features = sharpedge_surface_features,
|
|
||||||
batch_size = B, idx_query = sharpedge_idx_query, input_pc_size = input_sharpedge_pc_size)
|
|
||||||
|
|
||||||
query_features = torch.cat([query_random_features, query_sharpedge_features], dim = 1)
|
|
||||||
input_features = torch.cat([input_random_surface_features, input_sharpedge_surface_features], dim = 1)
|
|
||||||
|
|
||||||
if self.normal_pe:
|
|
||||||
# apply the fourier embeddings on the first 3 dims (xyz)
|
|
||||||
input_features_pe = self.fourier_embedder(input_features[..., :3])
|
|
||||||
query_features_pe = self.fourier_embedder(query_features[..., :3])
|
|
||||||
# replace the first 3 dims with the new PE ones
|
|
||||||
input_features = torch.cat([input_features_pe, input_features[..., :3]], dim = -1)
|
|
||||||
query_features = torch.cat([query_features_pe, query_features[..., :3]], dim = -1)
|
|
||||||
|
|
||||||
# concat at the channels dim
|
|
||||||
query = torch.cat([query, query_features], dim = -1)
|
|
||||||
data = torch.cat([data, input_features], dim = -1)
|
|
||||||
|
|
||||||
# don't return pc_info to avoid unnecessary memory usuage
|
|
||||||
return query.view(B, -1, query.shape[-1]), data.view(B, -1, data.shape[-1])
|
|
||||||
|
|
||||||
def forward(self, point_cloud: torch.Tensor, features: torch.Tensor):
|
|
||||||
|
|
||||||
query, data = self.sample_points_and_latents(point_cloud = point_cloud, features = features)
|
|
||||||
|
|
||||||
# apply projections
|
|
||||||
query = self.input_proj(query)
|
|
||||||
data = self.input_proj(data)
|
|
||||||
|
|
||||||
# apply cross attention between query and data
|
|
||||||
latents = self.cross_attn(query, data)
|
|
||||||
|
|
||||||
if self.self_attn is not None:
|
|
||||||
latents = self.self_attn(latents)
|
|
||||||
|
|
||||||
if self.ln_post is not None:
|
|
||||||
latents = self.ln_post(latents)
|
|
||||||
|
|
||||||
return latents
|
|
||||||
|
|
||||||
|
|
||||||
def subsample(self, pc, num_query, input_pc_size: int):
|
|
||||||
|
|
||||||
"""
|
|
||||||
num_query: number of points to keep after FPS
|
|
||||||
input_pc_size: number of points to select before FPS
|
|
||||||
"""
|
|
||||||
|
|
||||||
B, _, D = pc.shape
|
|
||||||
query_ratio = num_query / input_pc_size
|
|
||||||
|
|
||||||
# random subsampling of points inside the point cloud
|
|
||||||
idx_pc = torch.randperm(pc.shape[1], device = pc.device)[:input_pc_size]
|
|
||||||
input_pc = pc[:, idx_pc, :]
|
|
||||||
|
|
||||||
# flatten to allow applying fps across the whole batch
|
|
||||||
flattent_input_pc = input_pc.view(B * input_pc_size, D)
|
|
||||||
|
|
||||||
# construct a batch_down tensor to tell fps
|
|
||||||
# which points belong to which batch
|
|
||||||
N_down = int(flattent_input_pc.shape[0] / B)
|
|
||||||
batch_down = torch.arange(B).to(pc.device)
|
|
||||||
batch_down = torch.repeat_interleave(batch_down, N_down)
|
|
||||||
|
|
||||||
idx_query = fps(flattent_input_pc, batch_down, sampling_ratio = query_ratio)
|
|
||||||
query_pc = flattent_input_pc[idx_query].view(B, -1, D)
|
|
||||||
|
|
||||||
return query_pc, input_pc, idx_pc, idx_query
|
|
||||||
|
|
||||||
def handle_features(self, features, idx_pc, input_pc_size, batch_size: int, idx_query):
|
|
||||||
|
|
||||||
B = batch_size
|
|
||||||
|
|
||||||
input_surface_features = features[:, idx_pc, :]
|
|
||||||
flattent_input_features = input_surface_features.view(B * input_pc_size, -1)
|
|
||||||
query_features = flattent_input_features[idx_query].view(B, -1,
|
|
||||||
flattent_input_features.shape[-1])
|
|
||||||
|
|
||||||
return input_surface_features, query_features
|
|
||||||
|
|
||||||
def forward(self, pc, feats):
|
|
||||||
"""
|
|
||||||
|
|
||||||
Args:
|
|
||||||
pc (torch.FloatTensor): [B, N, 3]
|
|
||||||
feats (torch.FloatTensor or None): [B, N, C]
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
|
|
||||||
"""
|
|
||||||
|
|
||||||
query, data = self.sample_points_and_latents(pc, feats)
|
|
||||||
|
|
||||||
query = self.input_proj(query)
|
|
||||||
query = query
|
|
||||||
data = self.input_proj(data)
|
|
||||||
data = data
|
|
||||||
|
|
||||||
latents = self.cross_attn(query, data)
|
|
||||||
if self.self_attn is not None:
|
|
||||||
latents = self.self_attn(latents)
|
|
||||||
|
|
||||||
if self.ln_post is not None:
|
|
||||||
latents = self.ln_post(latents)
|
|
||||||
|
|
||||||
return latents
|
|
||||||
|
|
||||||
@dataclass
|
|
||||||
class Latent2MeshOutput():
|
|
||||||
# mesh for vertices and faces
|
|
||||||
vertices: None
|
|
||||||
faces: None
|
|
||||||
|
|
||||||
class SurfaceExtractor():
|
|
||||||
def compute_box_stat(self, bounds, octree_resolution: int):
|
|
||||||
|
|
||||||
# if float, turn it into a cube
|
|
||||||
if isinstance(bounds, float):
|
|
||||||
bounds = [-bounds, -bounds, -bounds, bounds, bounds, bounds]
|
|
||||||
|
|
||||||
bbox_min, bbox_max = np.array(bounds[0:3]), np.array(bounds[3:6])
|
|
||||||
bbox_size = bbox_max - bbox_min
|
|
||||||
grid_size = [int(octree_resolution) + 1, int(octree_resolution) + 1, int(octree_resolution) + 1]
|
|
||||||
return grid_size, bbox_min, bbox_size
|
|
||||||
|
|
||||||
def run(self, grid_logit, *, bounds, octree_resolution, level: float = 0.0, **kwargs):
|
|
||||||
# grid_logit from volume decoder
|
|
||||||
# use marching cube algo to turn an sdf to a mesh
|
|
||||||
vertices, faces, _, _ = measure.marching_cubes(grid_logit.cpu().numpy(),
|
|
||||||
level,
|
|
||||||
method = "lewiner")
|
|
||||||
|
|
||||||
grid_size, bbox_min, bbox_size = self.compute_box_stat(bounds = bounds, octree_resolution = octree_resolution)
|
|
||||||
vertices = vertices / grid_size * bbox_size + bbox_min
|
|
||||||
|
|
||||||
return vertices, faces
|
|
||||||
|
|
||||||
def __call__(self, grid_logits, **kwds):
|
|
||||||
|
|
||||||
outputs = []
|
|
||||||
veritces_list = []
|
|
||||||
faces_list = []
|
|
||||||
# loop over the batches
|
|
||||||
for i in range(grid_logits.shape[0]):
|
|
||||||
try:
|
|
||||||
# process each batch
|
|
||||||
vertices, faces = self.run(grid_logits[i], **kwds)
|
|
||||||
vertices = vertices.astype(np.float32)
|
|
||||||
faces = np.ascontiguousarray(faces)
|
|
||||||
#outputs.append(Latent2MeshOutput(vertices = vertices, faces = faces))
|
|
||||||
veritces_list.append(vertices)
|
|
||||||
faces_list.append(faces)
|
|
||||||
|
|
||||||
except Exception:
|
|
||||||
import traceback
|
|
||||||
traceback.print_exc()
|
|
||||||
outputs.append(None)
|
|
||||||
|
|
||||||
return outputs
|
|
||||||
|
|
||||||
def normalize_mesh(mesh, scale = 0.9999):
|
|
||||||
"""Normalize mesh to fit in [-scale, scale]. Translate mesh so its center is [0,0,0]"""
|
|
||||||
|
|
||||||
bbox = mesh.bounds
|
|
||||||
center = (bbox[1] + bbox[0]) / 2
|
|
||||||
|
|
||||||
max_extent = (bbox[1] - bbox[0]).max()
|
|
||||||
mesh.apply_translation(-center)
|
|
||||||
mesh.apply_scale((2 * scale) / max_extent)
|
|
||||||
|
|
||||||
return mesh
|
|
||||||
|
|
||||||
def sample_pointcloud(mesh, num = 200000):
|
|
||||||
""" Uniformly sample points from the surface of the mesh """
|
|
||||||
|
|
||||||
points, face_idx = mesh.sample(num, return_index = True)
|
|
||||||
normals = mesh.face_normals[face_idx]
|
|
||||||
return torch.from_numpy(points.astype(np.float32)), torch.from_numpy(normals.astype(np.float32))
|
|
||||||
|
|
||||||
def detect_sharp_edges(mesh, threshold=0.985):
|
|
||||||
"""Return edge indices (a, b) that lie on sharp boundaries of the mesh."""
|
|
||||||
|
|
||||||
V, F = mesh.vertices, mesh.faces
|
|
||||||
VN, FN = mesh.vertex_normals, mesh.face_normals
|
|
||||||
|
|
||||||
sharp_mask = np.ones(V.shape[0])
|
|
||||||
for i in range(3):
|
|
||||||
indices = F[:, i]
|
|
||||||
alignment = np.einsum('ij,ij->i', VN[indices], FN)
|
|
||||||
dot_stack = np.stack((sharp_mask[indices], alignment), axis=-1)
|
|
||||||
sharp_mask[indices] = np.min(dot_stack, axis=-1)
|
|
||||||
|
|
||||||
edge_a = np.concatenate([F[:, 0], F[:, 1], F[:, 2]])
|
|
||||||
edge_b = np.concatenate([F[:, 1], F[:, 2], F[:, 0]])
|
|
||||||
sharp_edges = (sharp_mask[edge_a] < threshold) & (sharp_mask[edge_b] < threshold)
|
|
||||||
|
|
||||||
return edge_a[sharp_edges], edge_b[sharp_edges]
|
|
||||||
|
|
||||||
|
|
||||||
def sharp_sample_pointcloud(mesh, num = 16384):
|
|
||||||
""" Sample points preferentially from sharp edges in the mesh. """
|
|
||||||
|
|
||||||
edge_a, edge_b = detect_sharp_edges(mesh)
|
|
||||||
V, VN = mesh.vertices, mesh.vertex_normals
|
|
||||||
|
|
||||||
va, vb = V[edge_a], V[edge_b]
|
|
||||||
na, nb = VN[edge_a], VN[edge_b]
|
|
||||||
|
|
||||||
edge_lengths = np.linalg.norm(vb - va, axis=-1)
|
|
||||||
weights = edge_lengths / edge_lengths.sum()
|
|
||||||
|
|
||||||
indices = np.searchsorted(np.cumsum(weights), np.random.rand(num))
|
|
||||||
t = np.random.rand(num, 1)
|
|
||||||
|
|
||||||
samples = t * va[indices] + (1 - t) * vb[indices]
|
|
||||||
normals = t * na[indices] + (1 - t) * nb[indices]
|
|
||||||
|
|
||||||
return samples.astype(np.float32), normals.astype(np.float32)
|
|
||||||
|
|
||||||
def load_surface_sharpedge(mesh, num_points=4096, num_sharp_points=4096, sharpedge_flag = True, device = "cuda"):
|
|
||||||
"""Load a surface with optional sharp-edge annotations from a trimesh mesh."""
|
|
||||||
|
|
||||||
import trimesh
|
|
||||||
|
|
||||||
try:
|
|
||||||
mesh_full = trimesh.util.concatenate(mesh.dump())
|
|
||||||
except Exception:
|
|
||||||
mesh_full = trimesh.util.concatenate(mesh)
|
|
||||||
|
|
||||||
mesh_full = normalize_mesh(mesh_full)
|
|
||||||
|
|
||||||
faces = mesh_full.faces
|
|
||||||
vertices = mesh_full.vertices
|
|
||||||
origin_face_count = faces.shape[0]
|
|
||||||
|
|
||||||
mesh_surface = trimesh.Trimesh(vertices=vertices, faces=faces[:origin_face_count])
|
|
||||||
mesh_fill = trimesh.Trimesh(vertices=vertices, faces=faces[origin_face_count:])
|
|
||||||
|
|
||||||
area_surface = mesh_surface.area
|
|
||||||
area_fill = mesh_fill.area
|
|
||||||
total_area = area_surface + area_fill
|
|
||||||
|
|
||||||
sample_num = 499712 // 2
|
|
||||||
fill_ratio = area_fill / total_area if total_area > 0 else 0
|
|
||||||
|
|
||||||
num_fill = int(sample_num * fill_ratio)
|
|
||||||
num_surface = sample_num - num_fill
|
|
||||||
|
|
||||||
surf_pts, surf_normals = sample_pointcloud(mesh_surface, num_surface)
|
|
||||||
fill_pts, fill_normals = (torch.zeros(0, 3), torch.zeros(0, 3)) if num_fill == 0 else sample_pointcloud(mesh_fill, num_fill)
|
|
||||||
|
|
||||||
sharp_pts, sharp_normals = sharp_sample_pointcloud(mesh_surface, sample_num)
|
|
||||||
|
|
||||||
def assemble_tensor(points, normals, label=None):
|
|
||||||
|
|
||||||
data = torch.cat([points, normals], dim=1).half().to(device)
|
|
||||||
|
|
||||||
if label is not None:
|
|
||||||
label_tensor = torch.full((data.shape[0], 1), float(label), dtype=torch.float16).to(device)
|
|
||||||
data = torch.cat([data, label_tensor], dim=1)
|
|
||||||
|
|
||||||
return data
|
|
||||||
|
|
||||||
surface = assemble_tensor(torch.cat([surf_pts.to(device), fill_pts.to(device)], dim=0),
|
|
||||||
torch.cat([surf_normals.to(device), fill_normals.to(device)], dim=0),
|
|
||||||
label = 0 if sharpedge_flag else None)
|
|
||||||
|
|
||||||
sharp_surface = assemble_tensor(torch.from_numpy(sharp_pts), torch.from_numpy(sharp_normals),
|
|
||||||
label = 1 if sharpedge_flag else None)
|
|
||||||
|
|
||||||
rng = np.random.default_rng()
|
|
||||||
|
|
||||||
surface = surface[rng.choice(surface.shape[0], num_points, replace = False)]
|
|
||||||
sharp_surface = sharp_surface[rng.choice(sharp_surface.shape[0], num_sharp_points, replace = False)]
|
|
||||||
|
|
||||||
full = torch.cat([surface, sharp_surface], dim = 0).unsqueeze(0)
|
|
||||||
|
|
||||||
return full
|
|
||||||
|
|
||||||
class SharpEdgeSurfaceLoader:
|
|
||||||
""" Load mesh surface and sharp edge samples. """
|
|
||||||
|
|
||||||
def __init__(self, num_uniform_points = 8192, num_sharp_points = 8192):
|
|
||||||
|
|
||||||
self.num_uniform_points = num_uniform_points
|
|
||||||
self.num_sharp_points = num_sharp_points
|
|
||||||
self.total_points = num_uniform_points + num_sharp_points
|
|
||||||
|
|
||||||
def __call__(self, mesh_input, device = "cuda"):
|
|
||||||
mesh = self._load_mesh(mesh_input)
|
|
||||||
return load_surface_sharpedge(mesh, self.num_uniform_points, self.num_sharp_points, device = device)
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _load_mesh(mesh_input):
|
|
||||||
import trimesh
|
|
||||||
|
|
||||||
if isinstance(mesh_input, str):
|
|
||||||
mesh = trimesh.load(mesh_input, force="mesh", merge_primitives = True)
|
|
||||||
else:
|
|
||||||
mesh = mesh_input
|
|
||||||
|
|
||||||
if isinstance(mesh, trimesh.Scene):
|
|
||||||
combined = None
|
|
||||||
for obj in mesh.geometry.values():
|
|
||||||
combined = obj if combined is None else combined + obj
|
|
||||||
return combined
|
|
||||||
|
|
||||||
return mesh
|
|
||||||
|
|
||||||
class FourierEmbedder(nn.Module):
|
|
||||||
def __init__(self, num_freq: int = 8, input_dim: int = 3, include_pi: bool = False):
|
|
||||||
super().__init__()
|
|
||||||
|
|
||||||
frequencies = 2.0 ** torch.arange(
|
|
||||||
num_freq,
|
|
||||||
dtype = torch.float32
|
|
||||||
)
|
|
||||||
|
|
||||||
if include_pi:
|
|
||||||
frequencies *= torch.pi
|
|
||||||
|
|
||||||
self.register_buffer("frequencies", frequencies, persistent = False)
|
|
||||||
|
|
||||||
self.out_dim = input_dim * (num_freq * 2 + 1)
|
|
||||||
|
|
||||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
|
||||||
|
|
||||||
embed = (x[..., None].contiguous() * self.frequencies).view(*x.shape[:-1], -1)
|
|
||||||
return torch.cat((x, embed.sin(), embed.cos()), dim = -1)
|
|
||||||
|
|
||||||
class DiagonalGaussianDistribution:
|
|
||||||
def __init__(self, params: torch.Tensor, feature_dim: int = -1):
|
|
||||||
|
|
||||||
# divide quant channels (8) into mean and log variance
|
|
||||||
self.mean, self.logvar = torch.chunk(params, 2, dim = feature_dim)
|
|
||||||
|
|
||||||
self.logvar = torch.clamp(self.logvar, -30.0, 20.0)
|
|
||||||
self.std = torch.exp(0.5 * self.logvar)
|
|
||||||
|
|
||||||
def sample(self):
|
|
||||||
|
|
||||||
eps = torch.randn_like(self.std)
|
|
||||||
z = self.mean + eps * self.std
|
|
||||||
|
|
||||||
return z
|
|
||||||
|
|
||||||
class VAE(nn.Module):
|
|
||||||
def __init__(self,
|
|
||||||
*,
|
|
||||||
num_latents: int = 4096,
|
|
||||||
embed_dim: int = 64,
|
|
||||||
width: int = 1024,
|
|
||||||
heads: int = 16,
|
|
||||||
num_decoder_layers: int = 16,
|
|
||||||
num_encoder_layers: int = 8,
|
|
||||||
pc_size: int = 81920,
|
|
||||||
pc_sharpedge_size: int = 0,
|
|
||||||
point_feats: int = 4,
|
|
||||||
downsample_ratio: int = 20,
|
|
||||||
geo_decoder_downsample_ratio: int = 1,
|
|
||||||
geo_decoder_mlp_expand_ratio: int = 4,
|
|
||||||
geo_decoder_ln_post: bool = True,
|
|
||||||
num_frequencies: int = 8,
|
|
||||||
qkv_bias: bool = False,
|
|
||||||
qk_norm: bool = True,
|
|
||||||
drop_path_rate: float = 0.0,
|
|
||||||
include_pi: bool = False,
|
|
||||||
scale_factor: float = 1.0039506158752403
|
|
||||||
):
|
|
||||||
|
|
||||||
super().__init__()
|
|
||||||
|
|
||||||
self.latent_shape = (num_latents, embed_dim)
|
|
||||||
self.scale_factor = scale_factor
|
|
||||||
|
|
||||||
self.fourier_embedder = FourierEmbedder(num_freq = num_frequencies, include_pi = include_pi)
|
|
||||||
|
|
||||||
self.encoder = PointCrossAttention(layers = num_encoder_layers,
|
|
||||||
num_latents = num_latents,
|
|
||||||
downsample_ratio = downsample_ratio,
|
|
||||||
heads = heads,
|
|
||||||
pc_size = pc_size,
|
|
||||||
width = width,
|
|
||||||
point_feats = point_feats,
|
|
||||||
fourier_embedder = self.fourier_embedder,
|
|
||||||
pc_sharpedge_size = pc_sharpedge_size)
|
|
||||||
|
|
||||||
self.transformer = Transformer(
|
|
||||||
width=width,
|
|
||||||
layers=num_decoder_layers,
|
|
||||||
heads=heads,
|
|
||||||
qkv_bias=qkv_bias,
|
|
||||||
qk_norm=qk_norm,
|
|
||||||
drop_path_rate=drop_path_rate
|
|
||||||
)
|
|
||||||
|
|
||||||
self.geo_decoder = CrossAttentionDecoder(
|
|
||||||
fourier_embedder = self.fourier_embedder,
|
|
||||||
out_channels = 1,
|
|
||||||
mlp_expand_ratio = geo_decoder_mlp_expand_ratio,
|
|
||||||
downsample_ratio = geo_decoder_downsample_ratio,
|
|
||||||
enable_ln_post = geo_decoder_ln_post,
|
|
||||||
width=width // geo_decoder_downsample_ratio,
|
|
||||||
heads=heads // geo_decoder_downsample_ratio,
|
|
||||||
qkv_bias = qkv_bias,
|
|
||||||
qk_norm= qk_norm
|
|
||||||
)
|
|
||||||
|
|
||||||
self.pre_kl = nn.Linear(width, embed_dim * 2)
|
|
||||||
self.post_kl = nn.Linear(embed_dim, width)
|
|
||||||
|
|
||||||
self.volume_decoder = VanillaVolumeDecoder()
|
|
||||||
self.surface_extractor = SurfaceExtractor()
|
|
||||||
|
|
||||||
|
|
||||||
def forward(self):
|
|
||||||
pass
|
|
||||||
|
|
||||||
def encode(self, surface):
|
|
||||||
|
|
||||||
pc, feats = surface[:, :, :3], surface[:, :, 3:]
|
|
||||||
latents = self.encoder(pc, feats)
|
|
||||||
|
|
||||||
moments = self.pre_kl(latents)
|
|
||||||
posterior = DiagonalGaussianDistribution(moments, feature_dim = -1)
|
|
||||||
|
|
||||||
latents = posterior.sample()
|
|
||||||
|
|
||||||
return latents
|
|
||||||
|
|
||||||
def decode(self, latents, **kwargs):
|
|
||||||
to_mesh = kwargs.pop("to_mesh", True)
|
|
||||||
latents = self.post_kl(latents)
|
|
||||||
latents = self.transformer(latents)
|
|
||||||
|
|
||||||
if not to_mesh:
|
|
||||||
return latents
|
|
||||||
|
|
||||||
grid_logits = self.volume_decoder(latents = latents, geo_decoder = self.geo_decoder, **kwargs)
|
|
||||||
mesh = self.surface_extractor(grid_logits, **kwargs)
|
|
||||||
|
|
||||||
return mesh
|
|
||||||
|
|
||||||
def load_vae(vae):
|
|
||||||
|
|
||||||
DEBUG = False
|
|
||||||
|
|
||||||
checkpoint = "model.fp16.ckpt"
|
|
||||||
missing, unexpected = vae.load_state_dict(torch.load(checkpoint), strict = not DEBUG)
|
|
||||||
|
|
||||||
if DEBUG:
|
|
||||||
print(f"Missing {len(missing)}: ", missing)
|
|
||||||
print(f"\nUnexpected {len(unexpected)}: ", unexpected)
|
|
||||||
|
|
||||||
return vae
|
|
||||||
|
|
||||||
|
|
||||||
18
comfy/sd.py
18
comfy/sd.py
@ -432,8 +432,8 @@ class VAE:
|
|||||||
self.memory_used_encode = lambda shape, dtype: 6000 * shape[3] * shape[4] * model_management.dtype_size(dtype)
|
self.memory_used_encode = lambda shape, dtype: 6000 * shape[3] * shape[4] * model_management.dtype_size(dtype)
|
||||||
self.memory_used_decode = lambda shape, dtype: 7000 * shape[3] * shape[4] * (8 * 8) * model_management.dtype_size(dtype)
|
self.memory_used_decode = lambda shape, dtype: 7000 * shape[3] * shape[4] * (8 * 8) * model_management.dtype_size(dtype)
|
||||||
|
|
||||||
# Hunyuan 3d v2 2.1
|
# Hunyuan 3d v2 2.0 & 2.1
|
||||||
elif 'geo_decoder.cross_attn_decoder.mlp.c_proj.weight' in sd:
|
elif "geo_decoder.cross_attn_decoder.ln_1.bias" in sd:
|
||||||
|
|
||||||
self.latent_dim = 1
|
self.latent_dim = 1
|
||||||
|
|
||||||
@ -451,19 +451,7 @@ class VAE:
|
|||||||
self.memory_used_decode = lambda shape, dtype, num_layers = 16, kv_cache_multiplier = 2: \
|
self.memory_used_decode = lambda shape, dtype, num_layers = 16, kv_cache_multiplier = 2: \
|
||||||
estimate_memory(shape, dtype, num_layers, kv_cache_multiplier)
|
estimate_memory(shape, dtype, num_layers, kv_cache_multiplier)
|
||||||
|
|
||||||
self.first_stage_model = comfy.ldm.hunyuan3dv2_1.vae.VAE()
|
self.first_stage_model = comfy.ldm.hunyuan3d.vae.ShapeVAE()
|
||||||
self.working_dtypes = [torch.float16, torch.bfloat16, torch.float32]
|
|
||||||
|
|
||||||
elif "geo_decoder.cross_attn_decoder.ln_1.bias" in sd:
|
|
||||||
self.latent_dim = 1
|
|
||||||
ln_post = "geo_decoder.ln_post.weight" in sd
|
|
||||||
inner_size = sd["geo_decoder.output_proj.weight"].shape[1]
|
|
||||||
downsample_ratio = sd["post_kl.weight"].shape[0] // inner_size
|
|
||||||
mlp_expand = sd["geo_decoder.cross_attn_decoder.mlp.c_fc.weight"].shape[0] // inner_size
|
|
||||||
self.memory_used_encode = lambda shape, dtype: (1000 * shape[2]) * model_management.dtype_size(dtype) # TODO
|
|
||||||
self.memory_used_decode = lambda shape, dtype: (1024 * 1024 * 1024 * 2.0) * model_management.dtype_size(dtype) # TODO
|
|
||||||
ddconfig = {"embed_dim": 64, "num_freqs": 8, "include_pi": False, "heads": 16, "width": 1024, "num_decoder_layers": 16, "qkv_bias": False, "qk_norm": True, "geo_decoder_mlp_expand_ratio": mlp_expand, "geo_decoder_downsample_ratio": downsample_ratio, "geo_decoder_ln_post": ln_post}
|
|
||||||
self.first_stage_model = comfy.ldm.hunyuan3d.vae.ShapeVAE(**ddconfig)
|
|
||||||
self.working_dtypes = [torch.float16, torch.bfloat16, torch.float32]
|
self.working_dtypes = [torch.float16, torch.bfloat16, torch.float32]
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@ -11,35 +11,16 @@ from comfy.cli_args import args
|
|||||||
class EmptyLatentHunyuan3Dv2:
|
class EmptyLatentHunyuan3Dv2:
|
||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(s):
|
def INPUT_TYPES(s):
|
||||||
return {
|
return {"required": {"resolution": ("INT", {"default": 3072, "min": 1, "max": 8192}),
|
||||||
"required": {
|
"batch_size": ("INT", {"default": 1, "min": 1, "max": 4096, "tooltip": "The number of latent images in the batch."}),
|
||||||
"resolution": ("INT", {"default": 3072, "min": 1, "max": 8192}),
|
}}
|
||||||
"batch_size": ("INT", {
|
|
||||||
"default": 1,
|
|
||||||
"min": 1,
|
|
||||||
"max": 4096,
|
|
||||||
"tooltip": "The number of latent images in the batch."
|
|
||||||
}),
|
|
||||||
"version": (["2.0", "2.1"], {
|
|
||||||
"default": "2.1",
|
|
||||||
"tooltip": "Choose latent layout version. 2.0: (B, C, N), 2.1: (B, N, C)"
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
RETURN_TYPES = ("LATENT",)
|
RETURN_TYPES = ("LATENT",)
|
||||||
FUNCTION = "generate"
|
FUNCTION = "generate"
|
||||||
|
|
||||||
CATEGORY = "latent/3d"
|
CATEGORY = "latent/3d"
|
||||||
|
|
||||||
def generate(self, resolution, batch_size, version):
|
def generate(self, resolution, batch_size):
|
||||||
embed_dim = 64
|
latent = torch.zeros([batch_size, resolution, 64], device=comfy.model_management.intermediate_device())
|
||||||
if version == "2.0":
|
|
||||||
latent = torch.zeros([batch_size, embed_dim, resolution],
|
|
||||||
device = comfy.model_management.intermediate_device())
|
|
||||||
else: # version = "2.1"
|
|
||||||
latent = torch.zeros([batch_size, resolution, embed_dim],
|
|
||||||
device = comfy.model_management.intermediate_device())
|
|
||||||
|
|
||||||
return ({"samples": latent, "type": "hunyuan3dv2"}, )
|
return ({"samples": latent, "type": "hunyuan3dv2"}, )
|
||||||
|
|
||||||
class Hunyuan3Dv2Conditioning:
|
class Hunyuan3Dv2Conditioning:
|
||||||
@ -97,53 +78,23 @@ class Hunyuan3Dv2ConditioningMultiView:
|
|||||||
class VOXEL:
|
class VOXEL:
|
||||||
def __init__(self, data):
|
def __init__(self, data):
|
||||||
self.data = data
|
self.data = data
|
||||||
|
|
||||||
class VAEDecodeHunyuan3D:
|
class VAEDecodeHunyuan3D:
|
||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(cls):
|
def INPUT_TYPES(s):
|
||||||
return {
|
return {"required": {"samples": ("LATENT", ),
|
||||||
"required": {
|
"vae": ("VAE", ),
|
||||||
"samples": ("LATENT",),
|
"num_chunks": ("INT", {"default": 8000, "min": 1000, "max": 500000}),
|
||||||
"vae": ("VAE",),
|
"octree_resolution": ("INT", {"default": 256, "min": 16, "max": 512}),
|
||||||
"version": (["2.0", "2.1"], {
|
}}
|
||||||
"default": "2.1",
|
RETURN_TYPES = ("VOXEL",)
|
||||||
"tooltip": "2.0 returns voxel grid; 2.1 returns implicit SDF function."
|
|
||||||
}),
|
|
||||||
"num_chunks": ("INT", {
|
|
||||||
"default": 8000, "min": 1000, "max": 500000,
|
|
||||||
}),
|
|
||||||
"octree_resolution": ("INT", {
|
|
||||||
"default": 256, "min": 16, "max": 512,
|
|
||||||
}),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
RETURN_TYPES = ("VOXEL", "MESH")
|
|
||||||
RETURN_NAMES = ("voxel", "mesh")
|
|
||||||
|
|
||||||
FUNCTION = "decode"
|
FUNCTION = "decode"
|
||||||
|
|
||||||
CATEGORY = "latent/3d"
|
CATEGORY = "latent/3d"
|
||||||
|
|
||||||
def decode(self, vae, samples, version, num_chunks, octree_resolution):
|
def decode(self, vae, samples, num_chunks, octree_resolution):
|
||||||
|
voxels = VOXEL(vae.decode(samples["samples"], vae_options={"num_chunks": num_chunks, "octree_resolution": octree_resolution}))
|
||||||
if version == "2.0":
|
return (voxels, )
|
||||||
voxel = vae.decode(samples["samples"], vae_options={
|
|
||||||
"num_chunks": num_chunks,
|
|
||||||
"octree_resolution": octree_resolution
|
|
||||||
})
|
|
||||||
return (VOXEL(voxel), None)
|
|
||||||
|
|
||||||
mesh = vae.decode(samples["samples"],vae_options={
|
|
||||||
"num_chunks": num_chunks,
|
|
||||||
"octree_resolution": octree_resolution,
|
|
||||||
"to_mesh": True
|
|
||||||
})
|
|
||||||
|
|
||||||
# ensure batch dim
|
|
||||||
if mesh.vertices.ndim == 2:
|
|
||||||
mesh.vertices = mesh.vertices[np.newaxis, ...]
|
|
||||||
mesh.faces = mesh.faces[np.newaxis, ...]
|
|
||||||
|
|
||||||
return (None, mesh)
|
|
||||||
|
|
||||||
def voxel_to_mesh(voxels, threshold=0.5, device=None):
|
def voxel_to_mesh(voxels, threshold=0.5, device=None):
|
||||||
if device is None:
|
if device is None:
|
||||||
@ -275,13 +226,9 @@ def voxel_to_mesh_surfnet(voxels, threshold=0.5, device=None):
|
|||||||
[0, 0, 1], [1, 0, 1], [0, 1, 1], [1, 1, 1]
|
[0, 0, 1], [1, 0, 1], [0, 1, 1], [1, 1, 1]
|
||||||
], device=device)
|
], device=device)
|
||||||
|
|
||||||
corner_values = torch.zeros((cell_positions.shape[0], 8), device=device)
|
pos = cell_positions.unsqueeze(1) + corner_offsets.unsqueeze(0)
|
||||||
for c, (dz, dy, dx) in enumerate(corner_offsets):
|
z_idx, y_idx, x_idx = pos.unbind(-1)
|
||||||
corner_values[:, c] = padded[
|
corner_values = padded[z_idx, y_idx, x_idx]
|
||||||
cell_positions[:, 0] + dz,
|
|
||||||
cell_positions[:, 1] + dy,
|
|
||||||
cell_positions[:, 2] + dx
|
|
||||||
]
|
|
||||||
|
|
||||||
corner_signs = corner_values > threshold
|
corner_signs = corner_values > threshold
|
||||||
has_inside = torch.any(corner_signs, dim=1)
|
has_inside = torch.any(corner_signs, dim=1)
|
||||||
@ -512,12 +459,8 @@ def save_glb(vertices, faces, filepath, metadata=None):
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
# Convert tensors to numpy arrays
|
# Convert tensors to numpy arrays
|
||||||
if isinstance(vertices, torch.tensor) and isinstance(faces, torch.tensor):
|
vertices_np = vertices.cpu().numpy().astype(np.float32)
|
||||||
vertices_np = vertices.cpu().numpy().astype(np.float32)
|
faces_np = faces.cpu().numpy().astype(np.uint32)
|
||||||
faces_np = faces.cpu().numpy().astype(np.uint32)
|
|
||||||
else:
|
|
||||||
vertices_np = vertices.astype(np.float32)
|
|
||||||
faces_np = faces.astype(np.uint32)
|
|
||||||
|
|
||||||
vertices_buffer = vertices_np.tobytes()
|
vertices_buffer = vertices_np.tobytes()
|
||||||
indices_buffer = faces_np.tobytes()
|
indices_buffer = faces_np.tobytes()
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user