From 5db44f5a55b68e9a0a0d3d70821eeaeb3bcd9159 Mon Sep 17 00:00:00 2001 From: Bruno Fargnoli Date: Fri, 26 Dec 2025 11:27:16 +0100 Subject: [PATCH] Improve low_vram stability - Thanks to @ManglerFTW --- nodes.py | 2 + trellis2/pipelines/trellis2_image_to_3d.py | 77 +++++++++++++++++++--- 2 files changed, 71 insertions(+), 8 deletions(-) diff --git a/nodes.py b/nodes.py index fe5d7d3..13cc28e 100644 --- a/nodes.py +++ b/nodes.py @@ -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, diff --git a/trellis2/pipelines/trellis2_image_to_3d.py b/trellis2/pipelines/trellis2_image_to_3d.py index 65d5601..3b79d57 100644 --- a/trellis2/pipelines/trellis2_image_to_3d.py +++ b/trellis2/pipelines/trellis2_image_to_3d.py @@ -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()