diff --git a/nodes.py b/nodes.py index 1abdde8..0fdb74d 100644 --- a/nodes.py +++ b/nodes.py @@ -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() diff --git a/trellis2/models/sc_vaes/fdg_vae.py b/trellis2/models/sc_vaes/fdg_vae.py index c9b5b07..1bb1f55 100644 --- a/trellis2/models/sc_vaes/fdg_vae.py +++ b/trellis2/models/sc_vaes/fdg_vae.py @@ -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]], diff --git a/trellis2/models/sc_vaes/sparse_unet_vae.py b/trellis2/models/sc_vaes/sparse_unet_vae.py index b9902a1..ad38c50 100644 --- a/trellis2/models/sc_vaes/sparse_unet_vae.py +++ b/trellis2/models/sc_vaes/sparse_unet_vae.py @@ -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: diff --git a/trellis2/modules/norm.py b/trellis2/modules/norm.py index 78675d0..2484c54 100644 --- a/trellis2/modules/norm.py +++ b/trellis2/modules/norm.py @@ -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): diff --git a/trellis2/modules/sparse/attention/full_attn.py b/trellis2/modules/sparse/attention/full_attn.py index 4eb74d2..49de383 100644 --- a/trellis2/modules/sparse/attention/full_attn.py +++ b/trellis2/modules/sparse/attention/full_attn.py @@ -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) + diff --git a/trellis2/modules/sparse/linear.py b/trellis2/modules/sparse/linear.py index 4431770..bc51f79 100644 --- a/trellis2/modules/sparse/linear.py +++ b/trellis2/modules/sparse/linear.py @@ -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)) diff --git a/trellis2/pipelines/trellis2_image_to_3d.py b/trellis2/pipelines/trellis2_image_to_3d.py index 3c44acb..7928804 100644 --- a/trellis2/pipelines/trellis2_image_to_3d.py +++ b/trellis2/pipelines/trellis2_image_to_3d.py @@ -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()