Improve low_vram stability - Thanks to @ManglerFTW

This commit is contained in:
Bruno Fargnoli
2025-12-26 11:27:16 +01:00
parent f30fd4eab9
commit 5db44f5a55
2 changed files with 71 additions and 8 deletions
+2
View File
@@ -829,6 +829,7 @@ class Trellis2PostProcessAndUnWrapAndRasterizer:
else:
resolution = int(dual_contouring_resolution)
print('Performing Dual Contouring ...')
# Perform Dual Contouring remeshing (rebuilds topology)
cumesh.init(*CuMesh.remeshing.remesh_narrow_band_dc(
vertices, faces,
@@ -1059,6 +1060,7 @@ class Trellis2Remesh:
else:
resolution = int(dual_contouring_resolution)
print('Performing Dual Contouring ...')
# Perform Dual Contouring remeshing (rebuilds topology)
cumesh.init(*CuMesh.remeshing.remesh_narrow_band_dc(
vertices, faces,
+69 -8
View File
@@ -90,6 +90,19 @@ class Trellis2ImageTo3DPipeline(Pipeline):
}
self._device = 'cpu'
def _cond_to(self, cond: dict, device: torch.device) -> dict:
# Move only tensors; keep other items unchanged
return {k: (v.to(device) if torch.is_tensor(v) else v) for k, v in cond.items()}
def _cond_cpu(self, cond: dict) -> dict:
return {k: (v.cpu() if torch.is_tensor(v) else v) for k, v in cond.items()}
def _cleanup_cuda(self):
import gc
gc.collect()
if torch.cuda.is_available():
torch.cuda.empty_cache()
@classmethod
def from_pretrained(cls, path: str, config_file: str = "pipeline.json", keep_models_loaded = True) -> "Trellis2ImageTo3DPipeline":
"""
@@ -412,6 +425,8 @@ class Trellis2ImageTo3DPipeline(Pipeline):
num_samples (int): The number of samples to generate.
sampler_params (dict): Additional parameters for the sampler.
"""
if self.low_vram:
cond = self._cond_to(cond, self.device)
# Sample sparse structure latent
flow_model = self.models['sparse_structure_flow_model']
reso = flow_model.resolution
@@ -430,7 +445,7 @@ class Trellis2ImageTo3DPipeline(Pipeline):
).samples
if self.low_vram:
flow_model.cpu()
self._cleanup_cuda()
# Decode sparse structure latent
decoder = self.models['sparse_structure_decoder']
if self.low_vram:
@@ -443,6 +458,12 @@ class Trellis2ImageTo3DPipeline(Pipeline):
decoded = torch.nn.functional.max_pool3d(decoded.float(), ratio, ratio, 0) > 0.5
coords = torch.argwhere(decoded)[:, [0, 2, 3, 4]].int()
coords = coords.cpu()
del decoded
del z_s
if self.low_vram:
cond = self._cond_cpu(cond)
self._cleanup_cuda()
return coords
def sample_shape_slat(
@@ -460,10 +481,14 @@ class Trellis2ImageTo3DPipeline(Pipeline):
coords (torch.Tensor): The coordinates of the sparse structure.
sampler_params (dict): Additional parameters for the sampler.
"""
if self.low_vram:
cond = self._cond_to(cond, self.device)
coords_dev = coords.to(self.device)
# Sample structured latent
noise = SparseTensor(
feats=torch.randn(coords.shape[0], flow_model.in_channels).to(self.device),
coords=coords,
feats=torch.randn(coords.shape[0], flow_model.in_channels, device=self.device),
coords=coords_dev,
)
sampler_params = {**self.shape_slat_sampler_params, **sampler_params}
if self.low_vram:
@@ -478,11 +503,17 @@ class Trellis2ImageTo3DPipeline(Pipeline):
).samples
if self.low_vram:
flow_model.cpu()
self._cleanup_cuda()
std = torch.tensor(self.shape_slat_normalization['std'])[None].to(slat.device)
mean = torch.tensor(self.shape_slat_normalization['mean'])[None].to(slat.device)
slat = slat * std + mean
del coords_dev
if self.low_vram:
cond = self._cond_cpu(cond)
self._cleanup_cuda()
return slat
def sample_shape_slat_cascade(
@@ -506,9 +537,16 @@ class Trellis2ImageTo3DPipeline(Pipeline):
sampler_params (dict): Additional parameters for the sampler.
"""
# LR
if self.low_vram:
lr_cond = self._cond_to(lr_cond, self.device)
cond = self._cond_to(cond, self.device)
coords_dev = coords.to(self.device)
# Sample structured latent
noise = SparseTensor(
feats=torch.randn(coords.shape[0], flow_model_lr.in_channels).to(self.device),
coords=coords,
feats=torch.randn(coords.shape[0], flow_model.in_channels, device=self.device),
coords=coords_dev,
)
sampler_params = {**self.shape_slat_sampler_params, **sampler_params}
if self.low_vram:
@@ -523,11 +561,17 @@ class Trellis2ImageTo3DPipeline(Pipeline):
).samples
if self.low_vram:
flow_model_lr.cpu()
self._cleanup_cuda()
std = torch.tensor(self.shape_slat_normalization['std'])[None].to(slat.device)
mean = torch.tensor(self.shape_slat_normalization['mean'])[None].to(slat.device)
slat = slat * std + mean
# Upsample
del coords_dev
if self.low_vram:
lr_cond = self._cond_cpu(lr_cond)
self._cleanup_cuda()
# Upsample
self.load_shape_slat_decoder()
if self.low_vram:
self.models['shape_slat_decoder'].to(self.device)
@@ -554,10 +598,11 @@ class Trellis2ImageTo3DPipeline(Pipeline):
break
hr_resolution -= 128
coords_dev = coords.to(self.device)
# Sample structured latent
noise = SparseTensor(
feats=torch.randn(coords.shape[0], flow_model.in_channels).to(self.device),
coords=coords,
feats=torch.randn(coords.shape[0], flow_model.in_channels, device=self.device),
coords=coords_dev,
)
sampler_params = {**self.shape_slat_sampler_params, **sampler_params}
if self.low_vram:
@@ -572,11 +617,17 @@ class Trellis2ImageTo3DPipeline(Pipeline):
).samples
if self.low_vram:
flow_model.cpu()
self._cleanup_cuda()
std = torch.tensor(self.shape_slat_normalization['std'])[None].to(slat.device)
mean = torch.tensor(self.shape_slat_normalization['mean'])[None].to(slat.device)
slat = slat * std + mean
del coords_dev
if self.low_vram:
cond = self._cond_cpu(cond)
self._cleanup_cuda()
return slat, hr_resolution
def decode_shape_slat(
@@ -626,6 +677,8 @@ class Trellis2ImageTo3DPipeline(Pipeline):
shape_slat (SparseTensor): The structured latent for shape
sampler_params (dict): Additional parameters for the sampler.
"""
if self.low_vram:
cond = self._cond_to(cond, self.device)
# Sample structured latent
std = torch.tensor(self.shape_slat_normalization['std'])[None].to(shape_slat.device)
mean = torch.tensor(self.shape_slat_normalization['mean'])[None].to(shape_slat.device)
@@ -647,11 +700,15 @@ class Trellis2ImageTo3DPipeline(Pipeline):
).samples
if self.low_vram:
flow_model.cpu()
self._cleanup_cuda()
std = torch.tensor(self.tex_slat_normalization['std'])[None].to(slat.device)
mean = torch.tensor(self.tex_slat_normalization['mean'])[None].to(slat.device)
slat = slat * std + mean
if self.low_vram:
cond = self._cond_cpu(cond)
self._cleanup_cuda()
return slat
def decode_tex_slat(
@@ -703,7 +760,11 @@ class Trellis2ImageTo3DPipeline(Pipeline):
resolution (int): The resolution of the output.
"""
meshes, subs = self.decode_shape_slat(shape_slat, resolution)
if self.low_vram:
self._cleanup_cuda()
tex_voxels = self.decode_tex_slat(tex_slat, subs)
if self.low_vram:
self._cleanup_cuda()
out_mesh = []
for m, v in zip(meshes, tex_voxels):
m.fill_holes()