Added processing by chunk
This commit is contained in:
@@ -820,8 +820,15 @@ class Trellis2MeshWithVoxelAdvancedGenerator:
|
||||
sparse_structure_sampler_params = {"steps":sparse_structure_steps,"guidance_strength":sparse_structure_guidance_strength,"guidance_rescale":sparse_structure_guidance_rescale,"guidance_interval":sparse_structure_guidance_interval,"rescale_t":sparse_structure_rescale_t}
|
||||
shape_slat_sampler_params = {"steps":shape_steps,"guidance_strength":shape_guidance_strength,"guidance_rescale":shape_guidance_rescale,"guidance_interval":shape_guidance_interval,"rescale_t":shape_rescale_t}
|
||||
tex_slat_sampler_params = {"steps":texture_steps,"guidance_strength":texture_guidance_strength,"guidance_rescale":texture_guidance_rescale,"guidance_interval":texture_guidance_interval,"rescale_t":texture_rescale_t}
|
||||
|
||||
if generate_texture_slat:
|
||||
num_steps = 4
|
||||
else:
|
||||
num_steps = 3
|
||||
|
||||
pbar = ProgressBar(num_steps)
|
||||
|
||||
mesh = pipeline.run(image=image_in, seed=seed, pipeline_type=pipeline_type, sparse_structure_sampler_params = sparse_structure_sampler_params, shape_slat_sampler_params = shape_slat_sampler_params, tex_slat_sampler_params = tex_slat_sampler_params, max_num_tokens = max_num_tokens, sparse_structure_resolution = sparse_structure_resolution, max_views = max_views, generate_texture_slat=generate_texture_slat)[0]
|
||||
mesh = pipeline.run(image=image_in, seed=seed, pipeline_type=pipeline_type, sparse_structure_sampler_params = sparse_structure_sampler_params, shape_slat_sampler_params = shape_slat_sampler_params, tex_slat_sampler_params = tex_slat_sampler_params, max_num_tokens = max_num_tokens, sparse_structure_resolution = sparse_structure_resolution, max_views = max_views, generate_texture_slat=generate_texture_slat, pbar=pbar)[0]
|
||||
|
||||
return (mesh,)
|
||||
|
||||
@@ -857,6 +864,7 @@ class Trellis2PostProcessAndUnWrapAndRasterizer:
|
||||
OUTPUT_NODE = True
|
||||
|
||||
def process(self, mesh, mesh_cluster_threshold_cone_half_angle_rad, mesh_cluster_refine_iterations, mesh_cluster_global_iterations, mesh_cluster_smooth_strength, texture_size, remesh, remesh_band, remesh_project, target_face_num, simplify_method, fill_holes, fill_holes_max_perimeter, texture_alpha_mode, dual_contouring_resolution, double_side_material,remove_floaters):
|
||||
pbar = ProgressBar(5)
|
||||
mesh_copy = copy.deepcopy(mesh)
|
||||
|
||||
aabb = [[-0.5, -0.5, -0.5], [0.5, 0.5, 0.5]]
|
||||
@@ -915,6 +923,7 @@ class Trellis2PostProcessAndUnWrapAndRasterizer:
|
||||
# Build BVH for the current mesh to guide remeshing
|
||||
print(f"Building BVH for current mesh...")
|
||||
bvh = CuMesh.cuBVH(vertices, faces)
|
||||
pbar.update(1)
|
||||
|
||||
print("Cleaning mesh...")
|
||||
# --- Branch 1: Standard Pipeline (Simplification & Cleaning) ---
|
||||
@@ -992,6 +1001,7 @@ class Trellis2PostProcessAndUnWrapAndRasterizer:
|
||||
cumesh.init(torch.from_numpy(new_vertices).float().cuda(), torch.from_numpy(new_faces).int().cuda())
|
||||
|
||||
print(f"After simplifying: {cumesh.num_vertices} vertices, {cumesh.num_faces} faces")
|
||||
pbar.update(1)
|
||||
|
||||
print('Unwrapping ...')
|
||||
out_vertices, out_faces, out_uvs, out_vmaps = cumesh.uv_unwrap(
|
||||
@@ -1004,6 +1014,7 @@ class Trellis2PostProcessAndUnWrapAndRasterizer:
|
||||
return_vmaps=True,
|
||||
verbose=True,
|
||||
)
|
||||
pbar.update(1)
|
||||
|
||||
out_vertices = out_vertices.cuda()
|
||||
out_faces = out_faces.cuda()
|
||||
@@ -1054,6 +1065,7 @@ class Trellis2PostProcessAndUnWrapAndRasterizer:
|
||||
|
||||
# --- Texture Post-Processing & Material Construction ---
|
||||
print("Finalizing mesh...")
|
||||
pbar.update(1)
|
||||
|
||||
mask = mask.cpu().numpy()
|
||||
|
||||
@@ -1101,7 +1113,8 @@ class Trellis2PostProcessAndUnWrapAndRasterizer:
|
||||
vertex_normals=normals_np,
|
||||
process=False,
|
||||
visual=Trimesh.visual.TextureVisuals(uv=uvs_np,material=material)
|
||||
)
|
||||
)
|
||||
pbar.update(1)
|
||||
|
||||
del cumesh
|
||||
gc.collect()
|
||||
|
||||
@@ -15,6 +15,7 @@ from .sparse_unet_vae import (
|
||||
from .sparse_unet_vae import (
|
||||
SparseUnetVaeEncoder,
|
||||
SparseUnetVaeDecoder,
|
||||
chunked_apply,
|
||||
)
|
||||
from ...representations import Mesh
|
||||
from o_voxel.convert import flexible_dual_grid_to_mesh
|
||||
@@ -84,9 +85,16 @@ class FlexiDualGridVaeDecoder(SparseUnetVaeDecoder):
|
||||
decoded = super().forward(x, **kwargs)
|
||||
if self.training:
|
||||
h, subs_gt, subs = decoded
|
||||
vertices = h.replace((1 + 2 * self.voxel_margin) * F.sigmoid(h.feats[..., 0:3]) - self.voxel_margin)
|
||||
intersected_logits = h.replace(h.feats[..., 3:6])
|
||||
quad_lerp = h.replace(F.softplus(h.feats[..., 6:7]))
|
||||
if self.low_vram:
|
||||
vertices_feats = chunked_apply(lambda t: (1 + 2 * self.voxel_margin) * F.sigmoid(t[..., 0:3]) - self.voxel_margin, h.feats, self.chunk_size)
|
||||
vertices = h.replace(vertices_feats)
|
||||
intersected_logits = h.replace(h.feats[..., 3:6])
|
||||
quad_lerp_feats = chunked_apply(lambda t: F.softplus(t[..., 6:7]), h.feats, self.chunk_size)
|
||||
quad_lerp = h.replace(quad_lerp_feats)
|
||||
else:
|
||||
vertices = h.replace((1 + 2 * self.voxel_margin) * F.sigmoid(h.feats[..., 0:3]) - self.voxel_margin)
|
||||
intersected_logits = h.replace(h.feats[..., 3:6])
|
||||
quad_lerp = h.replace(F.softplus(h.feats[..., 6:7]))
|
||||
mesh = [Mesh(flexible_dual_grid_to_mesh(
|
||||
h.coords[:, 1:], v.feats, i.feats, q.feats,
|
||||
aabb=[[-0.5, -0.5, -0.5], [0.5, 0.5, 0.5]],
|
||||
@@ -97,9 +105,16 @@ class FlexiDualGridVaeDecoder(SparseUnetVaeDecoder):
|
||||
else:
|
||||
out_list = list(decoded) if isinstance(decoded, tuple) else [decoded]
|
||||
h = out_list[0]
|
||||
vertices = h.replace((1 + 2 * self.voxel_margin) * F.sigmoid(h.feats[..., 0:3]) - self.voxel_margin)
|
||||
intersected = h.replace(h.feats[..., 3:6] > 0)
|
||||
quad_lerp = h.replace(F.softplus(h.feats[..., 6:7]))
|
||||
if self.low_vram:
|
||||
vertices_feats = chunked_apply(lambda t: (1 + 2 * self.voxel_margin) * F.sigmoid(t[..., 0:3]) - self.voxel_margin, h.feats, self.chunk_size)
|
||||
vertices = h.replace(vertices_feats)
|
||||
intersected = h.replace(h.feats[..., 3:6] > 0)
|
||||
quad_lerp_feats = chunked_apply(lambda t: F.softplus(t[..., 6:7]), h.feats, self.chunk_size)
|
||||
quad_lerp = h.replace(quad_lerp_feats)
|
||||
else:
|
||||
vertices = h.replace((1 + 2 * self.voxel_margin) * F.sigmoid(h.feats[..., 0:3]) - self.voxel_margin)
|
||||
intersected = h.replace(h.feats[..., 3:6] > 0)
|
||||
quad_lerp = h.replace(F.softplus(h.feats[..., 6:7]))
|
||||
mesh = [Mesh(*flexible_dual_grid_to_mesh(
|
||||
h.coords[:, 1:], v.feats, i.feats, q.feats,
|
||||
aabb=[[-0.5, -0.5, -0.5], [0.5, 0.5, 0.5]],
|
||||
|
||||
@@ -8,6 +8,22 @@ from ...modules import sparse as sp
|
||||
from ...modules.norm import LayerNorm32
|
||||
|
||||
|
||||
def chunked_apply(module: nn.Module, x: torch.Tensor, chunk_size: int) -> torch.Tensor:
|
||||
if chunk_size <= 0 or x.shape[0] <= chunk_size:
|
||||
return module(x)
|
||||
|
||||
# Process first chunk to determine output shape and dtype
|
||||
out_0 = module(x[0:chunk_size])
|
||||
out_shape = (x.shape[0],) + out_0.shape[1:]
|
||||
out = torch.empty(out_shape, device=x.device, dtype=out_0.dtype)
|
||||
out[0:chunk_size] = out_0
|
||||
|
||||
# Process remaining chunks
|
||||
for i in range(chunk_size, x.shape[0], chunk_size):
|
||||
out[i:i+chunk_size] = module(x[i:i+chunk_size])
|
||||
return out
|
||||
|
||||
|
||||
class SparseResBlock3d(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
@@ -25,6 +41,8 @@ class SparseResBlock3d(nn.Module):
|
||||
self.upsample = upsample
|
||||
self.resample_mode = resample_mode
|
||||
self.use_checkpoint = use_checkpoint
|
||||
self.low_vram = False
|
||||
self.chunk_size = 65536
|
||||
|
||||
assert not (downsample and upsample), "Cannot downsample and upsample at the same time"
|
||||
|
||||
@@ -64,19 +82,26 @@ class SparseResBlock3d(nn.Module):
|
||||
return x
|
||||
|
||||
def _forward(self, x: sp.SparseTensor) -> sp.SparseTensor:
|
||||
subdiv = None
|
||||
if self.upsample:
|
||||
subdiv = self.to_subdiv(x)
|
||||
h = x.replace(self.norm1(x.feats))
|
||||
h = h.replace(F.silu(h.feats))
|
||||
if self.low_vram:
|
||||
h = x.replace(chunked_apply(self.norm1, x.feats, self.chunk_size))
|
||||
h = h.replace(chunked_apply(F.silu, h.feats, self.chunk_size))
|
||||
else:
|
||||
h = x.replace(self.norm1(x.feats))
|
||||
h = h.replace(F.silu(h.feats))
|
||||
if self.resample_mode == 'spatial2channel':
|
||||
h = self.conv1(h)
|
||||
h = self._updown(h, subdiv)
|
||||
x = self._updown(x, subdiv)
|
||||
if self.resample_mode == 'nearest':
|
||||
h = self.conv1(h)
|
||||
h = h.replace(self.norm2(h.feats))
|
||||
h = h.replace(F.silu(h.feats))
|
||||
if self.low_vram:
|
||||
h = h.replace(chunked_apply(self.norm2, h.feats, self.chunk_size))
|
||||
h = h.replace(chunked_apply(F.silu, h.feats, self.chunk_size))
|
||||
else:
|
||||
h = h.replace(self.norm2(h.feats))
|
||||
h = h.replace(F.silu(h.feats))
|
||||
h = self.conv2(h)
|
||||
h = h + self.skip_connection(x)
|
||||
if self.upsample:
|
||||
@@ -96,11 +121,15 @@ class SparseResBlockDownsample3d(nn.Module):
|
||||
channels: int,
|
||||
out_channels: Optional[int] = None,
|
||||
use_checkpoint: bool = False,
|
||||
low_vram: bool = False,
|
||||
chunk_size: int = 65536,
|
||||
):
|
||||
super().__init__()
|
||||
self.channels = channels
|
||||
self.out_channels = out_channels or channels
|
||||
self.use_checkpoint = use_checkpoint
|
||||
self.low_vram = low_vram
|
||||
self.chunk_size = chunk_size
|
||||
|
||||
self.norm1 = LayerNorm32(channels, elementwise_affine=True, eps=1e-6)
|
||||
self.norm2 = LayerNorm32(self.out_channels, elementwise_affine=False, eps=1e-6)
|
||||
@@ -110,13 +139,21 @@ class SparseResBlockDownsample3d(nn.Module):
|
||||
self.updown = sp.SparseDownsample(2)
|
||||
|
||||
def _forward(self, x: sp.SparseTensor) -> sp.SparseTensor:
|
||||
h = x.replace(self.norm1(x.feats))
|
||||
h = h.replace(F.silu(h.feats))
|
||||
if self.low_vram:
|
||||
h = x.replace(chunked_apply(self.norm1, x.feats, self.chunk_size))
|
||||
h = h.replace(chunked_apply(F.silu, h.feats, self.chunk_size))
|
||||
else:
|
||||
h = x.replace(self.norm1(x.feats))
|
||||
h = h.replace(F.silu(h.feats))
|
||||
h = self.updown(h)
|
||||
x = self.updown(x)
|
||||
h = self.conv1(h)
|
||||
h = h.replace(self.norm2(h.feats))
|
||||
h = h.replace(F.silu(h.feats))
|
||||
if self.low_vram:
|
||||
h = h.replace(chunked_apply(self.norm2, h.feats, self.chunk_size))
|
||||
h = h.replace(chunked_apply(F.silu, h.feats, self.chunk_size))
|
||||
else:
|
||||
h = h.replace(self.norm2(h.feats))
|
||||
h = h.replace(F.silu(h.feats))
|
||||
h = self.conv2(h)
|
||||
h = h + self.skip_connection(x)
|
||||
return h
|
||||
@@ -135,12 +172,16 @@ class SparseResBlockUpsample3d(nn.Module):
|
||||
out_channels: Optional[int] = None,
|
||||
use_checkpoint: bool = False,
|
||||
pred_subdiv: bool = True,
|
||||
low_vram: bool = False,
|
||||
chunk_size: int = 65536,
|
||||
):
|
||||
super().__init__()
|
||||
self.channels = channels
|
||||
self.out_channels = out_channels or channels
|
||||
self.use_checkpoint = use_checkpoint
|
||||
self.pred_subdiv = pred_subdiv
|
||||
self.low_vram = low_vram
|
||||
self.chunk_size = chunk_size
|
||||
|
||||
self.norm1 = LayerNorm32(channels, elementwise_affine=True, eps=1e-6)
|
||||
self.norm2 = LayerNorm32(self.out_channels, elementwise_affine=False, eps=1e-6)
|
||||
@@ -154,14 +195,22 @@ class SparseResBlockUpsample3d(nn.Module):
|
||||
def _forward(self, x: sp.SparseTensor, subdiv: sp.SparseTensor = None) -> sp.SparseTensor:
|
||||
if self.pred_subdiv:
|
||||
subdiv = self.to_subdiv(x)
|
||||
h = x.replace(self.norm1(x.feats))
|
||||
h = h.replace(F.silu(h.feats))
|
||||
if self.low_vram:
|
||||
h = x.replace(chunked_apply(self.norm1, x.feats, self.chunk_size))
|
||||
h = h.replace(chunked_apply(F.silu, h.feats, self.chunk_size))
|
||||
else:
|
||||
h = x.replace(self.norm1(x.feats))
|
||||
h = h.replace(F.silu(h.feats))
|
||||
subdiv_binarized = subdiv.replace(subdiv.feats > 0) if subdiv is not None else None
|
||||
h = self.updown(h, subdiv_binarized)
|
||||
x = self.updown(x, subdiv_binarized)
|
||||
h = self.conv1(h)
|
||||
h = h.replace(self.norm2(h.feats))
|
||||
h = h.replace(F.silu(h.feats))
|
||||
if self.low_vram:
|
||||
h = h.replace(chunked_apply(self.norm2, h.feats, self.chunk_size))
|
||||
h = h.replace(chunked_apply(F.silu, h.feats, self.chunk_size))
|
||||
else:
|
||||
h = h.replace(self.norm2(h.feats))
|
||||
h = h.replace(F.silu(h.feats))
|
||||
h = self.conv2(h)
|
||||
h = h + self.skip_connection(x)
|
||||
if self.pred_subdiv:
|
||||
@@ -182,11 +231,15 @@ class SparseResBlockS2C3d(nn.Module):
|
||||
channels: int,
|
||||
out_channels: Optional[int] = None,
|
||||
use_checkpoint: bool = False,
|
||||
low_vram: bool = False,
|
||||
chunk_size: int = 65536,
|
||||
):
|
||||
super().__init__()
|
||||
self.channels = channels
|
||||
self.out_channels = out_channels or channels
|
||||
self.use_checkpoint = use_checkpoint
|
||||
self.low_vram = low_vram
|
||||
self.chunk_size = chunk_size
|
||||
|
||||
self.norm1 = LayerNorm32(channels, elementwise_affine=True, eps=1e-6)
|
||||
self.norm2 = LayerNorm32(self.out_channels, elementwise_affine=False, eps=1e-6)
|
||||
@@ -196,13 +249,21 @@ class SparseResBlockS2C3d(nn.Module):
|
||||
self.updown = sp.SparseSpatial2Channel(2)
|
||||
|
||||
def _forward(self, x: sp.SparseTensor) -> sp.SparseTensor:
|
||||
h = x.replace(self.norm1(x.feats))
|
||||
h = h.replace(F.silu(h.feats))
|
||||
if self.low_vram:
|
||||
h = x.replace(chunked_apply(self.norm1, x.feats, self.chunk_size))
|
||||
h = h.replace(chunked_apply(F.silu, h.feats, self.chunk_size))
|
||||
else:
|
||||
h = x.replace(self.norm1(x.feats))
|
||||
h = h.replace(F.silu(h.feats))
|
||||
h = self.conv1(h)
|
||||
h = self.updown(h)
|
||||
x = self.updown(x)
|
||||
h = h.replace(self.norm2(h.feats))
|
||||
h = h.replace(F.silu(h.feats))
|
||||
if self.low_vram:
|
||||
h = h.replace(chunked_apply(self.norm2, h.feats, self.chunk_size))
|
||||
h = h.replace(chunked_apply(F.silu, h.feats, self.chunk_size))
|
||||
else:
|
||||
h = h.replace(self.norm2(h.feats))
|
||||
h = h.replace(F.silu(h.feats))
|
||||
h = self.conv2(h)
|
||||
h = h + self.skip_connection(x)
|
||||
return h
|
||||
@@ -221,12 +282,16 @@ class SparseResBlockC2S3d(nn.Module):
|
||||
out_channels: Optional[int] = None,
|
||||
use_checkpoint: bool = False,
|
||||
pred_subdiv: bool = True,
|
||||
low_vram: bool = False,
|
||||
chunk_size: int = 65536,
|
||||
):
|
||||
super().__init__()
|
||||
self.channels = channels
|
||||
self.out_channels = out_channels or channels
|
||||
self.use_checkpoint = use_checkpoint
|
||||
self.pred_subdiv = pred_subdiv
|
||||
self.low_vram = low_vram
|
||||
self.chunk_size = chunk_size
|
||||
|
||||
self.norm1 = LayerNorm32(channels, elementwise_affine=True, eps=1e-6)
|
||||
self.norm2 = LayerNorm32(self.out_channels, elementwise_affine=False, eps=1e-6)
|
||||
@@ -240,14 +305,22 @@ class SparseResBlockC2S3d(nn.Module):
|
||||
def _forward(self, x: sp.SparseTensor, subdiv: sp.SparseTensor = None) -> sp.SparseTensor:
|
||||
if self.pred_subdiv:
|
||||
subdiv = self.to_subdiv(x)
|
||||
h = x.replace(self.norm1(x.feats))
|
||||
h = h.replace(F.silu(h.feats))
|
||||
if self.low_vram:
|
||||
h = x.replace(chunked_apply(self.norm1, x.feats, self.chunk_size))
|
||||
h = h.replace(chunked_apply(F.silu, h.feats, self.chunk_size))
|
||||
else:
|
||||
h = x.replace(self.norm1(x.feats))
|
||||
h = h.replace(F.silu(h.feats))
|
||||
h = self.conv1(h)
|
||||
subdiv_binarized = subdiv.replace(subdiv.feats > 0) if subdiv is not None else None
|
||||
h = self.updown(h, subdiv_binarized)
|
||||
x = self.updown(x, subdiv_binarized)
|
||||
h = h.replace(self.norm2(h.feats))
|
||||
h = h.replace(F.silu(h.feats))
|
||||
if self.low_vram:
|
||||
h = h.replace(chunked_apply(self.norm2, h.feats, self.chunk_size))
|
||||
h = h.replace(chunked_apply(F.silu, h.feats, self.chunk_size))
|
||||
else:
|
||||
h = h.replace(self.norm2(h.feats))
|
||||
h = h.replace(F.silu(h.feats))
|
||||
h = self.conv2(h)
|
||||
h = h + self.skip_connection(x)
|
||||
if self.pred_subdiv:
|
||||
@@ -272,6 +345,8 @@ class SparseConvNeXtBlock3d(nn.Module):
|
||||
super().__init__()
|
||||
self.channels = channels
|
||||
self.use_checkpoint = use_checkpoint
|
||||
self.low_vram = False
|
||||
self.chunk_size = 65536
|
||||
|
||||
self.norm = LayerNorm32(channels, elementwise_affine=True, eps=1e-6)
|
||||
self.conv = sp.SparseConv3d(channels, channels, 3)
|
||||
@@ -283,8 +358,12 @@ class SparseConvNeXtBlock3d(nn.Module):
|
||||
|
||||
def _forward(self, x: sp.SparseTensor) -> sp.SparseTensor:
|
||||
h = self.conv(x)
|
||||
h = h.replace(self.norm(h.feats))
|
||||
h = h.replace(self.mlp(h.feats))
|
||||
if self.low_vram:
|
||||
h = h.replace(chunked_apply(self.norm, h.feats, self.chunk_size))
|
||||
h = h.replace(chunked_apply(self.mlp, h.feats, self.chunk_size))
|
||||
else:
|
||||
h = h.replace(self.norm(h.feats))
|
||||
h = h.replace(self.mlp(h.feats))
|
||||
return h + x
|
||||
|
||||
def forward(self, x: sp.SparseTensor) -> sp.SparseTensor:
|
||||
@@ -314,7 +393,6 @@ class SparseUnetVaeEncoder(nn.Module):
|
||||
self.model_channels = model_channels
|
||||
self.num_blocks = num_blocks
|
||||
self.dtype = torch.float16 if use_fp16 else torch.float32
|
||||
self.dtype = torch.float16 if use_fp16 else torch.float32
|
||||
|
||||
self.input_layer = sp.SparseLinear(in_channels, model_channels[0])
|
||||
self.to_latent = sp.SparseLinear(model_channels[-1], 2 * latent_channels)
|
||||
@@ -339,9 +417,22 @@ class SparseUnetVaeEncoder(nn.Module):
|
||||
)
|
||||
|
||||
self.initialize_weights()
|
||||
self._low_vram = False
|
||||
self.chunk_size = 65536
|
||||
if use_fp16:
|
||||
self.convert_to_fp16()
|
||||
|
||||
@property
|
||||
def low_vram(self) -> bool:
|
||||
return self._low_vram
|
||||
|
||||
@low_vram.setter
|
||||
def low_vram(self, value: bool):
|
||||
self._low_vram = value
|
||||
for m in self.modules():
|
||||
if hasattr(m, 'low_vram') and m is not self:
|
||||
m.low_vram = value
|
||||
|
||||
@property
|
||||
def device(self) -> torch.device:
|
||||
"""
|
||||
@@ -376,9 +467,18 @@ class SparseUnetVaeEncoder(nn.Module):
|
||||
for i, res in enumerate(self.blocks):
|
||||
for j, block in enumerate(res):
|
||||
h = block(h)
|
||||
h = h.type(x.dtype)
|
||||
h = h.replace(F.layer_norm(h.feats, h.feats.shape[-1:]))
|
||||
h = self.to_latent(h)
|
||||
|
||||
if self.low_vram:
|
||||
def fused_finalize(t):
|
||||
t = t.type(x.dtype)
|
||||
t = F.layer_norm(t, (t.shape[-1],))
|
||||
t = F.linear(t, self.to_latent.weight, self.to_latent.bias)
|
||||
return t
|
||||
h = h.replace(chunked_apply(fused_finalize, h.feats, self.chunk_size))
|
||||
else:
|
||||
h = h.type(x.dtype)
|
||||
h = h.replace(F.layer_norm(h.feats, h.feats.shape[-1:]))
|
||||
h = self.to_latent(h)
|
||||
|
||||
# Sample from the posterior distribution
|
||||
mean, logvar = h.feats.chunk(2, dim=-1)
|
||||
@@ -444,8 +544,21 @@ class SparseUnetVaeDecoder(nn.Module):
|
||||
)
|
||||
|
||||
self.initialize_weights()
|
||||
self._low_vram = False
|
||||
self.chunk_size = 65536
|
||||
if use_fp16:
|
||||
self.convert_to_fp16()
|
||||
|
||||
@property
|
||||
def low_vram(self) -> bool:
|
||||
return self._low_vram
|
||||
|
||||
@low_vram.setter
|
||||
def low_vram(self, value: bool):
|
||||
self._low_vram = value
|
||||
for m in self.modules():
|
||||
if hasattr(m, 'low_vram') and m is not self:
|
||||
m.low_vram = value
|
||||
|
||||
@property
|
||||
def device(self) -> torch.device:
|
||||
@@ -495,9 +608,19 @@ class SparseUnetVaeDecoder(nn.Module):
|
||||
h = block(h, subdiv=guide_subs[i] if guide_subs is not None else None)
|
||||
else:
|
||||
h = block(h)
|
||||
h = h.type(x.dtype)
|
||||
h = h.replace(F.layer_norm(h.feats, h.feats.shape[-1:]))
|
||||
h = self.output_layer(h)
|
||||
if self.low_vram:
|
||||
def fused_finalize(t):
|
||||
t = t.type(x.dtype)
|
||||
t = F.layer_norm(t, (t.shape[-1],))
|
||||
# Manually call linear to avoid creating intermediate specific type tensor if possible,
|
||||
# but utilizing the weights of output_layer
|
||||
t = F.linear(t, self.output_layer.weight, self.output_layer.bias)
|
||||
return t
|
||||
h = h.replace(chunked_apply(fused_finalize, h.feats, self.chunk_size))
|
||||
else:
|
||||
h = h.type(x.dtype)
|
||||
h = h.replace(F.layer_norm(h.feats, h.feats.shape[-1:]))
|
||||
h = self.output_layer(h)
|
||||
if self.training and self.pred_subdiv:
|
||||
return h, subs_gt, subs
|
||||
else:
|
||||
|
||||
@@ -3,23 +3,59 @@ import torch.nn as nn
|
||||
from .utils import manual_cast
|
||||
|
||||
|
||||
def chunked_apply(module, x: torch.Tensor, chunk_size: int) -> torch.Tensor:
|
||||
if chunk_size <= 0 or x.shape[0] <= chunk_size:
|
||||
return module(x)
|
||||
|
||||
# Process first chunk to determine output shape and dtype
|
||||
out_0 = module(x[0:chunk_size])
|
||||
out_shape = (x.shape[0],) + out_0.shape[1:]
|
||||
out = torch.empty(out_shape, device=x.device, dtype=out_0.dtype)
|
||||
out[0:chunk_size] = out_0
|
||||
|
||||
# Process remaining chunks
|
||||
for i in range(chunk_size, x.shape[0], chunk_size):
|
||||
out[i:i+chunk_size] = module(x[i:i+chunk_size])
|
||||
return out
|
||||
|
||||
|
||||
class LayerNorm32(nn.LayerNorm):
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.low_vram = False
|
||||
self.chunk_size = 65536
|
||||
|
||||
def _forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
x_dtype = x.dtype
|
||||
x = manual_cast(x, torch.float32)
|
||||
o = super().forward(x)
|
||||
return manual_cast(o, x_dtype)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
if self.low_vram:
|
||||
return chunked_apply(self._forward, x, self.chunk_size)
|
||||
return self._forward(x)
|
||||
|
||||
|
||||
class GroupNorm32(nn.GroupNorm):
|
||||
"""
|
||||
A GroupNorm layer that converts to float32 before the forward pass.
|
||||
"""
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.low_vram = False
|
||||
self.chunk_size = 65536
|
||||
|
||||
def _forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
x_dtype = x.dtype
|
||||
x = manual_cast(x, torch.float32)
|
||||
o = super().forward(x)
|
||||
return manual_cast(o, x_dtype)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
if self.low_vram:
|
||||
return chunked_apply(self._forward, x, self.chunk_size)
|
||||
return self._forward(x)
|
||||
|
||||
|
||||
class ChannelLayerNorm32(LayerNorm32):
|
||||
|
||||
@@ -169,52 +169,100 @@ def sparse_scaled_dot_product_attention(*args, **kwargs):
|
||||
k = k.reshape(N * L, H, CI) # [T_KV, H, Ci]
|
||||
v = v.reshape(N * L, H, CO) # [T_KV, H, Co]
|
||||
|
||||
# ---- attention dispatch with safe fallbacks (Windows-friendly) ----
|
||||
def _sdpa_varlen(q, k, v, q_seqlen, kv_seqlen):
|
||||
# q: [TQ, H, Cq], k: [TK, H, Cq], v: [TK, H, Cv]
|
||||
# returns: [TQ, H, Cv]
|
||||
outs = []
|
||||
q_off = 0
|
||||
kv_off = 0
|
||||
for n in range(len(q_seqlen)):
|
||||
qn = q_seqlen[n]
|
||||
kn = kv_seqlen[n]
|
||||
q_i = q[q_off:q_off + qn].transpose(0, 1) # [H, qn, C]
|
||||
k_i = k[kv_off:kv_off + kn].transpose(0, 1) # [H, kn, C]
|
||||
v_i = v[kv_off:kv_off + kn].transpose(0, 1) # [H, kn, Cv]
|
||||
|
||||
# SDPA expects [B, heads, L, C] or [heads, L, C]. We use [1, H, L, C]
|
||||
q_i = q_i.unsqueeze(0) # [1, H, qn, C]
|
||||
k_i = k_i.unsqueeze(0) # [1, H, kn, C]
|
||||
v_i = v_i.unsqueeze(0) # [1, H, kn, Cv]
|
||||
|
||||
out_i = torch.nn.functional.scaled_dot_product_attention(
|
||||
q_i, k_i, v_i,
|
||||
dropout_p=0.0,
|
||||
is_causal=False
|
||||
)[0] # [H, qn, Cv]
|
||||
|
||||
outs.append(out_i.transpose(0, 1)) # [qn, H, Cv]
|
||||
q_off += qn
|
||||
kv_off += kn
|
||||
|
||||
return torch.cat(outs, dim=0) # [TQ, H, Cv]
|
||||
|
||||
if num_all_args == 1:
|
||||
q, k, v = qkv.unbind(dim=1) # qkv: [T, 3, H, C]
|
||||
elif num_all_args == 2:
|
||||
k, v = kv.unbind(dim=1) # kv: [T, 2, H, C]
|
||||
# for num_all_args == 3, q/k/v already set
|
||||
|
||||
if config.ATTN == 'xformers':
|
||||
if 'xops' not in globals():
|
||||
import xformers.ops as xops
|
||||
if num_all_args == 1:
|
||||
q, k, v = qkv.unbind(dim=1)
|
||||
elif num_all_args == 2:
|
||||
k, v = kv.unbind(dim=1)
|
||||
q = q.unsqueeze(0)
|
||||
k = k.unsqueeze(0)
|
||||
v = v.unsqueeze(0)
|
||||
mask = xops.fmha.BlockDiagonalMask.from_seqlens(q_seqlen, kv_seqlen)
|
||||
out = xops.memory_efficient_attention(q, k, v, mask)[0]
|
||||
try:
|
||||
if 'xops' not in globals():
|
||||
import xformers.ops as xops
|
||||
q_x = q.unsqueeze(0)
|
||||
k_x = k.unsqueeze(0)
|
||||
v_x = v.unsqueeze(0)
|
||||
mask = xops.fmha.BlockDiagonalMask.from_seqlens(q_seqlen, kv_seqlen)
|
||||
out = xops.memory_efficient_attention(q_x, k_x, v_x, mask)[0]
|
||||
except Exception:
|
||||
# fallback to torch SDPA
|
||||
out = _sdpa_varlen(q, k, v, q_seqlen, kv_seqlen)
|
||||
|
||||
elif config.ATTN == 'flash_attn':
|
||||
if 'flash_attn' not in globals():
|
||||
import flash_attn
|
||||
cu_seqlens_q = torch.cat([torch.tensor([0]), torch.cumsum(torch.tensor(q_seqlen), dim=0)]).int().to(device)
|
||||
if num_all_args in [2, 3]:
|
||||
try:
|
||||
if 'flash_attn' not in globals():
|
||||
import flash_attn
|
||||
cu_seqlens_q = torch.cat([torch.tensor([0]), torch.cumsum(torch.tensor(q_seqlen), dim=0)]).int().to(device)
|
||||
cu_seqlens_kv = torch.cat([torch.tensor([0]), torch.cumsum(torch.tensor(kv_seqlen), dim=0)]).int().to(device)
|
||||
if num_all_args == 1:
|
||||
out = flash_attn.flash_attn_varlen_qkvpacked_func(qkv, cu_seqlens_q, max(q_seqlen))
|
||||
elif num_all_args == 2:
|
||||
out = flash_attn.flash_attn_varlen_kvpacked_func(q, kv, cu_seqlens_q, cu_seqlens_kv, max(q_seqlen), max(kv_seqlen))
|
||||
elif num_all_args == 3:
|
||||
out = flash_attn.flash_attn_varlen_func(q, k, v, cu_seqlens_q, cu_seqlens_kv, max(q_seqlen), max(kv_seqlen))
|
||||
if num_all_args == 1:
|
||||
# needs packed qkv [T, 3, H, C]
|
||||
out = flash_attn.flash_attn_varlen_qkvpacked_func(qkv, cu_seqlens_q, max(q_seqlen))
|
||||
elif num_all_args == 2:
|
||||
out = flash_attn.flash_attn_varlen_kvpacked_func(q, kv, cu_seqlens_q, cu_seqlens_kv, max(q_seqlen), max(kv_seqlen))
|
||||
elif num_all_args == 3:
|
||||
out = flash_attn.flash_attn_varlen_func(q, k, v, cu_seqlens_q, cu_seqlens_kv, max(q_seqlen), max(kv_seqlen))
|
||||
except Exception:
|
||||
# fallback to torch SDPA
|
||||
out = _sdpa_varlen(q, k, v, q_seqlen, kv_seqlen)
|
||||
|
||||
elif config.ATTN == 'flash_attn_3':
|
||||
if 'flash_attn_3' not in globals():
|
||||
import flash_attn_interface as flash_attn_3
|
||||
cu_seqlens_q = torch.cat([torch.tensor([0]), torch.cumsum(torch.tensor(q_seqlen), dim=0)]).int().to(device)
|
||||
if num_all_args == 1:
|
||||
q, k, v = qkv.unbind(dim=1)
|
||||
cu_seqlens_kv = cu_seqlens_q.clone()
|
||||
max_q_seqlen = max_kv_seqlen = max(q_seqlen)
|
||||
elif num_all_args == 2:
|
||||
k, v = kv.unbind(dim=1)
|
||||
cu_seqlens_kv = torch.cat([torch.tensor([0]), torch.cumsum(torch.tensor(kv_seqlen), dim=0)]).int().to(device)
|
||||
max_q_seqlen = max(q_seqlen)
|
||||
max_kv_seqlen = max(kv_seqlen)
|
||||
elif num_all_args == 3:
|
||||
cu_seqlens_kv = torch.cat([torch.tensor([0]), torch.cumsum(torch.tensor(kv_seqlen), dim=0)]).int().to(device)
|
||||
max_q_seqlen = max(q_seqlen)
|
||||
max_kv_seqlen = max(kv_seqlen)
|
||||
out = flash_attn_3.flash_attn_varlen_func(q, k, v, cu_seqlens_q, cu_seqlens_kv, max_q_seqlen, max_kv_seqlen)
|
||||
try:
|
||||
if 'flash_attn_3' not in globals():
|
||||
import flash_attn_interface as flash_attn_3
|
||||
cu_seqlens_q = torch.cat([torch.tensor([0]), torch.cumsum(torch.tensor(q_seqlen), dim=0)]).int().to(device)
|
||||
if num_all_args == 1:
|
||||
q3, k3, v3 = qkv.unbind(dim=1)
|
||||
cu_seqlens_kv = cu_seqlens_q.clone()
|
||||
max_q_seqlen = max_kv_seqlen = max(q_seqlen)
|
||||
out = flash_attn_3.flash_attn_varlen_func(q3, k3, v3, cu_seqlens_q, cu_seqlens_kv, max_q_seqlen, max_kv_seqlen)
|
||||
elif num_all_args == 2:
|
||||
k3, v3 = kv.unbind(dim=1)
|
||||
cu_seqlens_kv = torch.cat([torch.tensor([0]), torch.cumsum(torch.tensor(kv_seqlen), dim=0)]).int().to(device)
|
||||
out = flash_attn_3.flash_attn_varlen_func(q, k3, v3, cu_seqlens_q, cu_seqlens_kv, max(q_seqlen), max(kv_seqlen))
|
||||
elif num_all_args == 3:
|
||||
cu_seqlens_kv = torch.cat([torch.tensor([0]), torch.cumsum(torch.tensor(kv_seqlen), dim=0)]).int().to(device)
|
||||
out = flash_attn_3.flash_attn_varlen_func(q, k, v, cu_seqlens_q, cu_seqlens_kv, max(q_seqlen), max(kv_seqlen))
|
||||
except Exception:
|
||||
# fallback to torch SDPA
|
||||
out = _sdpa_varlen(q, k, v, q_seqlen, kv_seqlen)
|
||||
|
||||
else:
|
||||
raise ValueError(f"Unknown attention module: {config.ATTN}")
|
||||
|
||||
# final fallback: torch SDPA
|
||||
out = _sdpa_varlen(q, k, v, q_seqlen, kv_seqlen)
|
||||
|
||||
if s is not None:
|
||||
return s.replace(out)
|
||||
else:
|
||||
return out.reshape(N, L, H, -1)
|
||||
|
||||
|
||||
@@ -7,9 +7,29 @@ __all__ = [
|
||||
]
|
||||
|
||||
|
||||
def chunked_apply(module, x: torch.Tensor, chunk_size: int) -> torch.Tensor:
|
||||
if chunk_size <= 0 or x.shape[0] <= chunk_size:
|
||||
return module(x)
|
||||
|
||||
# Process first chunk to determine output shape and dtype
|
||||
out_0 = module(x[0:chunk_size])
|
||||
out_shape = (x.shape[0],) + out_0.shape[1:]
|
||||
out = torch.empty(out_shape, device=x.device, dtype=out_0.dtype)
|
||||
out[0:chunk_size] = out_0
|
||||
|
||||
# Process remaining chunks
|
||||
for i in range(chunk_size, x.shape[0], chunk_size):
|
||||
out[i:i+chunk_size] = module(x[i:i+chunk_size])
|
||||
return out
|
||||
|
||||
|
||||
class SparseLinear(nn.Linear):
|
||||
def __init__(self, in_features, out_features, bias=True):
|
||||
super(SparseLinear, self).__init__(in_features, out_features, bias)
|
||||
self.low_vram = False
|
||||
self.chunk_size = 65536
|
||||
|
||||
def forward(self, input: VarLenTensor) -> VarLenTensor:
|
||||
if self.low_vram:
|
||||
return input.replace(chunked_apply(super().forward, input.feats, self.chunk_size))
|
||||
return input.replace(super().forward(input.feats))
|
||||
|
||||
@@ -21,6 +21,8 @@ import nvdiffrast.torch as dr
|
||||
import cv2
|
||||
import flex_gemm
|
||||
|
||||
from comfy.utils import ProgressBar
|
||||
|
||||
|
||||
class Trellis2ImageTo3DPipeline(Pipeline):
|
||||
"""
|
||||
@@ -844,7 +846,8 @@ class Trellis2ImageTo3DPipeline(Pipeline):
|
||||
max_num_tokens: int = 49152,
|
||||
sparse_structure_resolution: int = 32,
|
||||
max_views: int = 4,
|
||||
generate_texture_slat = True
|
||||
generate_texture_slat = True,
|
||||
pbar = None
|
||||
) -> List[MeshWithVoxel]:
|
||||
"""
|
||||
Run the pipeline.
|
||||
@@ -897,6 +900,9 @@ class Trellis2ImageTo3DPipeline(Pipeline):
|
||||
cond_512 = self.get_cond(images, 512, max_views = max_views)
|
||||
cond_1024 = self.get_cond(images, 1024, max_views = max_views) if pipeline_type != '512' else None
|
||||
|
||||
if pbar is not None:
|
||||
pbar.update(1)
|
||||
|
||||
if not self.keep_models_loaded:
|
||||
self.unload_image_cond_model()
|
||||
|
||||
@@ -909,6 +915,9 @@ class Trellis2ImageTo3DPipeline(Pipeline):
|
||||
num_samples, sparse_structure_sampler_params
|
||||
)
|
||||
|
||||
if pbar is not None:
|
||||
pbar.update(1)
|
||||
|
||||
if not self.keep_models_loaded:
|
||||
self.unload_sparse_structure_model()
|
||||
|
||||
@@ -921,6 +930,9 @@ class Trellis2ImageTo3DPipeline(Pipeline):
|
||||
coords, shape_slat_sampler_params
|
||||
)
|
||||
|
||||
if pbar is not None:
|
||||
pbar.update(1)
|
||||
|
||||
if not self.keep_models_loaded:
|
||||
self.unload_shape_slat_flow_model_512()
|
||||
|
||||
@@ -931,6 +943,9 @@ class Trellis2ImageTo3DPipeline(Pipeline):
|
||||
cond_512, self.models['tex_slat_flow_model_512'],
|
||||
shape_slat, tex_slat_sampler_params
|
||||
)
|
||||
|
||||
if pbar is not None:
|
||||
pbar.update(1)
|
||||
|
||||
if not self.keep_models_loaded:
|
||||
self.unload_tex_slat_flow_model_512()
|
||||
@@ -944,6 +959,9 @@ class Trellis2ImageTo3DPipeline(Pipeline):
|
||||
coords, shape_slat_sampler_params
|
||||
)
|
||||
|
||||
if pbar is not None:
|
||||
pbar.update(1)
|
||||
|
||||
if not self.keep_models_loaded:
|
||||
self.unload_shape_slat_flow_model_1024()
|
||||
|
||||
@@ -954,6 +972,9 @@ class Trellis2ImageTo3DPipeline(Pipeline):
|
||||
cond_1024, self.models['tex_slat_flow_model_1024'],
|
||||
shape_slat, tex_slat_sampler_params
|
||||
)
|
||||
|
||||
if pbar is not None:
|
||||
pbar.update(1)
|
||||
|
||||
if not self.keep_models_loaded:
|
||||
self.unload_tex_slat_flow_model_1024()
|
||||
@@ -970,6 +991,9 @@ class Trellis2ImageTo3DPipeline(Pipeline):
|
||||
max_num_tokens
|
||||
)
|
||||
|
||||
if pbar is not None:
|
||||
pbar.update(1)
|
||||
|
||||
if not self.keep_models_loaded:
|
||||
self.unload_shape_slat_flow_model_512()
|
||||
self.unload_shape_slat_flow_model_1024()
|
||||
@@ -981,6 +1005,9 @@ class Trellis2ImageTo3DPipeline(Pipeline):
|
||||
cond_1024, self.models['tex_slat_flow_model_1024'],
|
||||
shape_slat, tex_slat_sampler_params
|
||||
)
|
||||
|
||||
if pbar is not None:
|
||||
pbar.update(1)
|
||||
|
||||
if not self.keep_models_loaded:
|
||||
self.unload_tex_slat_flow_model_1024()
|
||||
@@ -994,9 +1021,14 @@ class Trellis2ImageTo3DPipeline(Pipeline):
|
||||
coords, shape_slat_sampler_params,
|
||||
max_num_tokens
|
||||
)
|
||||
|
||||
if pbar is not None:
|
||||
pbar.update(1)
|
||||
|
||||
if not self.keep_models_loaded:
|
||||
self.unload_shape_slat_flow_model_512()
|
||||
self.unload_shape_slat_flow_model_1024()
|
||||
|
||||
if generate_texture_slat:
|
||||
self.unload_tex_slat_flow_model_512()
|
||||
self.load_tex_slat_flow_model_1024()
|
||||
@@ -1004,6 +1036,10 @@ class Trellis2ImageTo3DPipeline(Pipeline):
|
||||
cond_1024, self.models['tex_slat_flow_model_1024'],
|
||||
shape_slat, tex_slat_sampler_params
|
||||
)
|
||||
|
||||
if pbar is not None:
|
||||
pbar.update(1)
|
||||
|
||||
if not self.keep_models_loaded:
|
||||
self.unload_tex_slat_flow_model_1024()
|
||||
elif pipeline_type == '4096_cascade':
|
||||
@@ -1016,9 +1052,14 @@ class Trellis2ImageTo3DPipeline(Pipeline):
|
||||
coords, shape_slat_sampler_params,
|
||||
max_num_tokens
|
||||
)
|
||||
|
||||
if pbar is not None:
|
||||
pbar.update(1)
|
||||
|
||||
if not self.keep_models_loaded:
|
||||
self.unload_shape_slat_flow_model_512()
|
||||
self.unload_shape_slat_flow_model_1024()
|
||||
|
||||
if generate_texture_slat:
|
||||
self.unload_tex_slat_flow_model_512()
|
||||
self.load_tex_slat_flow_model_1024()
|
||||
@@ -1026,6 +1067,10 @@ class Trellis2ImageTo3DPipeline(Pipeline):
|
||||
cond_1024, self.models['tex_slat_flow_model_1024'],
|
||||
shape_slat, tex_slat_sampler_params
|
||||
)
|
||||
|
||||
if pbar is not None:
|
||||
pbar.update(1)
|
||||
|
||||
if not self.keep_models_loaded:
|
||||
self.unload_tex_slat_flow_model_1024()
|
||||
elif pipeline_type == '1536_cascade':
|
||||
@@ -1039,6 +1084,9 @@ class Trellis2ImageTo3DPipeline(Pipeline):
|
||||
max_num_tokens
|
||||
)
|
||||
|
||||
if pbar is not None:
|
||||
pbar.update(1)
|
||||
|
||||
if not self.keep_models_loaded:
|
||||
self.unload_shape_slat_flow_model_512()
|
||||
self.unload_shape_slat_flow_model_1024()
|
||||
@@ -1050,6 +1098,9 @@ class Trellis2ImageTo3DPipeline(Pipeline):
|
||||
cond_1024, self.models['tex_slat_flow_model_1024'],
|
||||
shape_slat, tex_slat_sampler_params
|
||||
)
|
||||
|
||||
if pbar is not None:
|
||||
pbar.update(1)
|
||||
|
||||
if not self.keep_models_loaded:
|
||||
self.unload_tex_slat_flow_model_1024()
|
||||
|
||||
Reference in New Issue
Block a user