Added processing by chunk

This commit is contained in:
Bruno Fargnoli
2026-01-02 18:15:14 +01:00
parent 0f5c7c305a
commit 848162efc4
7 changed files with 387 additions and 81 deletions
+15 -2
View File
@@ -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()
+21 -6
View File
@@ -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]],
+153 -30
View File
@@ -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:
+38 -2
View File
@@ -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):
+88 -40
View File
@@ -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)
+20
View File
@@ -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))
+52 -1
View File
@@ -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()