Improve low_vram stability - Thanks to @ManglerFTW
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user