Spaces:
Runtime error
Runtime error
Upload 74 files
Browse files- trellis/.DS_Store +0 -0
- trellis/modules/attention/__init__.py +1 -1
- trellis/modules/attention/full_attn.py +10 -0
- trellis/modules/sparse/__init__.py +1 -1
- trellis/modules/sparse/attention/full_attn.py +18 -0
- trellis/modules/sparse/attention/serialized_attn.py +10 -0
- trellis/modules/sparse/attention/windowed_attn.py +10 -0
- trellis/modules/sparse/basic.py +65 -0
- trellis/pipelines/__init__.py +1 -0
- trellis/pipelines/samplers/classifier_free_guidance_mixin.py +31 -4
- trellis/pipelines/samplers/flow_euler.py +6 -175
- trellis/pipelines/samplers/guidance_interval_mixin.py +16 -5
- trellis/pipelines/trellis_hybrid_pipeline.py +694 -0
- trellis/pipelines/trellis_image_to_3d.py +12 -135
trellis/.DS_Store
ADDED
|
Binary file (8.2 kB). View file
|
|
|
trellis/modules/attention/__init__.py
CHANGED
|
@@ -14,7 +14,7 @@ def __from_env():
|
|
| 14 |
env_attn_backend = os.environ.get('ATTN_BACKEND')
|
| 15 |
env_sttn_debug = os.environ.get('ATTN_DEBUG')
|
| 16 |
|
| 17 |
-
if env_attn_backend is not None and env_attn_backend in ['xformers', 'flash_attn', 'sdpa', 'naive']:
|
| 18 |
BACKEND = env_attn_backend
|
| 19 |
if env_sttn_debug is not None:
|
| 20 |
DEBUG = env_sttn_debug == '1'
|
|
|
|
| 14 |
env_attn_backend = os.environ.get('ATTN_BACKEND')
|
| 15 |
env_sttn_debug = os.environ.get('ATTN_DEBUG')
|
| 16 |
|
| 17 |
+
if env_attn_backend is not None and env_attn_backend in ['xformers', 'flash_attn', 'flash_attn_3', 'sdpa', 'naive']:
|
| 18 |
BACKEND = env_attn_backend
|
| 19 |
if env_sttn_debug is not None:
|
| 20 |
DEBUG = env_sttn_debug == '1'
|
trellis/modules/attention/full_attn.py
CHANGED
|
@@ -7,6 +7,8 @@ if BACKEND == 'xformers':
|
|
| 7 |
import xformers.ops as xops
|
| 8 |
elif BACKEND == 'flash_attn':
|
| 9 |
import flash_attn
|
|
|
|
|
|
|
| 10 |
elif BACKEND == 'sdpa':
|
| 11 |
from torch.nn.functional import scaled_dot_product_attention as sdpa
|
| 12 |
elif BACKEND == 'naive':
|
|
@@ -118,6 +120,14 @@ def scaled_dot_product_attention(*args, **kwargs):
|
|
| 118 |
out = flash_attn.flash_attn_kvpacked_func(q, kv)
|
| 119 |
elif num_all_args == 3:
|
| 120 |
out = flash_attn.flash_attn_func(q, k, v)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 121 |
elif BACKEND == 'sdpa':
|
| 122 |
if num_all_args == 1:
|
| 123 |
q, k, v = qkv.unbind(dim=2)
|
|
|
|
| 7 |
import xformers.ops as xops
|
| 8 |
elif BACKEND == 'flash_attn':
|
| 9 |
import flash_attn
|
| 10 |
+
elif BACKEND == 'flash_attn_3':
|
| 11 |
+
import flash_attn_interface as flash_attn_3
|
| 12 |
elif BACKEND == 'sdpa':
|
| 13 |
from torch.nn.functional import scaled_dot_product_attention as sdpa
|
| 14 |
elif BACKEND == 'naive':
|
|
|
|
| 120 |
out = flash_attn.flash_attn_kvpacked_func(q, kv)
|
| 121 |
elif num_all_args == 3:
|
| 122 |
out = flash_attn.flash_attn_func(q, k, v)
|
| 123 |
+
elif BACKEND == 'flash_attn_3':
|
| 124 |
+
if num_all_args == 1:
|
| 125 |
+
out = flash_attn_3.flash_attn_qkvpacked_func(qkv)
|
| 126 |
+
elif num_all_args == 2:
|
| 127 |
+
k, v = kv.unbind(dim=2)
|
| 128 |
+
out = flash_attn_3.flash_attn_func(q, k, v)
|
| 129 |
+
elif num_all_args == 3:
|
| 130 |
+
out = flash_attn_3.flash_attn_func(q, k, v)
|
| 131 |
elif BACKEND == 'sdpa':
|
| 132 |
if num_all_args == 1:
|
| 133 |
q, k, v = qkv.unbind(dim=2)
|
trellis/modules/sparse/__init__.py
CHANGED
|
@@ -21,7 +21,7 @@ def __from_env():
|
|
| 21 |
BACKEND = env_sparse_backend
|
| 22 |
if env_sparse_debug is not None:
|
| 23 |
DEBUG = env_sparse_debug == '1'
|
| 24 |
-
if env_sparse_attn is not None and env_sparse_attn in ['xformers', 'flash_attn']:
|
| 25 |
ATTN = env_sparse_attn
|
| 26 |
|
| 27 |
print(f"[SPARSE] Backend: {BACKEND}, Attention: {ATTN}")
|
|
|
|
| 21 |
BACKEND = env_sparse_backend
|
| 22 |
if env_sparse_debug is not None:
|
| 23 |
DEBUG = env_sparse_debug == '1'
|
| 24 |
+
if env_sparse_attn is not None and env_sparse_attn in ['xformers', 'flash_attn', 'flash_attn_3']:
|
| 25 |
ATTN = env_sparse_attn
|
| 26 |
|
| 27 |
print(f"[SPARSE] Backend: {BACKEND}, Attention: {ATTN}")
|
trellis/modules/sparse/attention/full_attn.py
CHANGED
|
@@ -7,6 +7,8 @@ if ATTN == 'xformers':
|
|
| 7 |
import xformers.ops as xops
|
| 8 |
elif ATTN == 'flash_attn':
|
| 9 |
import flash_attn
|
|
|
|
|
|
|
| 10 |
else:
|
| 11 |
raise ValueError(f"Unknown attention module: {ATTN}")
|
| 12 |
|
|
@@ -206,6 +208,22 @@ def sparse_scaled_dot_product_attention(*args, **kwargs):
|
|
| 206 |
out = flash_attn.flash_attn_varlen_kvpacked_func(q, kv, cu_seqlens_q, cu_seqlens_kv, max(q_seqlen), max(kv_seqlen))
|
| 207 |
elif num_all_args == 3:
|
| 208 |
out = flash_attn.flash_attn_varlen_func(q, k, v, cu_seqlens_q, cu_seqlens_kv, max(q_seqlen), max(kv_seqlen))
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 209 |
else:
|
| 210 |
raise ValueError(f"Unknown attention module: {ATTN}")
|
| 211 |
|
|
|
|
| 7 |
import xformers.ops as xops
|
| 8 |
elif ATTN == 'flash_attn':
|
| 9 |
import flash_attn
|
| 10 |
+
elif ATTN == 'flash_attn_3':
|
| 11 |
+
import flash_attn_interface as flash_attn_3
|
| 12 |
else:
|
| 13 |
raise ValueError(f"Unknown attention module: {ATTN}")
|
| 14 |
|
|
|
|
| 208 |
out = flash_attn.flash_attn_varlen_kvpacked_func(q, kv, cu_seqlens_q, cu_seqlens_kv, max(q_seqlen), max(kv_seqlen))
|
| 209 |
elif num_all_args == 3:
|
| 210 |
out = flash_attn.flash_attn_varlen_func(q, k, v, cu_seqlens_q, cu_seqlens_kv, max(q_seqlen), max(kv_seqlen))
|
| 211 |
+
elif ATTN == 'flash_attn_3':
|
| 212 |
+
cu_seqlens_q = torch.cat([torch.tensor([0]), torch.cumsum(torch.tensor(q_seqlen), dim=0)]).int().to(device)
|
| 213 |
+
if num_all_args == 1:
|
| 214 |
+
q, k, v = qkv.unbind(dim=1)
|
| 215 |
+
cu_seqlens_kv = cu_seqlens_q.clone()
|
| 216 |
+
max_q_seqlen = max_kv_seqlen = max(q_seqlen)
|
| 217 |
+
elif num_all_args == 2:
|
| 218 |
+
k, v = kv.unbind(dim=1)
|
| 219 |
+
cu_seqlens_kv = torch.cat([torch.tensor([0]), torch.cumsum(torch.tensor(kv_seqlen), dim=0)]).int().to(device)
|
| 220 |
+
max_q_seqlen = max(q_seqlen)
|
| 221 |
+
max_kv_seqlen = max(kv_seqlen)
|
| 222 |
+
elif num_all_args == 3:
|
| 223 |
+
cu_seqlens_kv = torch.cat([torch.tensor([0]), torch.cumsum(torch.tensor(kv_seqlen), dim=0)]).int().to(device)
|
| 224 |
+
max_q_seqlen = max(q_seqlen)
|
| 225 |
+
max_kv_seqlen = max(kv_seqlen)
|
| 226 |
+
out = flash_attn_3.flash_attn_varlen_func(q, k, v, cu_seqlens_q, cu_seqlens_kv, max_q_seqlen, max_kv_seqlen)
|
| 227 |
else:
|
| 228 |
raise ValueError(f"Unknown attention module: {ATTN}")
|
| 229 |
|
trellis/modules/sparse/attention/serialized_attn.py
CHANGED
|
@@ -9,6 +9,8 @@ if ATTN == 'xformers':
|
|
| 9 |
import xformers.ops as xops
|
| 10 |
elif ATTN == 'flash_attn':
|
| 11 |
import flash_attn
|
|
|
|
|
|
|
| 12 |
else:
|
| 13 |
raise ValueError(f"Unknown attention module: {ATTN}")
|
| 14 |
|
|
@@ -168,6 +170,9 @@ def sparse_serialized_scaled_dot_product_self_attention(
|
|
| 168 |
out = xops.memory_efficient_attention(q, k, v) # [B, N, H, C]
|
| 169 |
elif ATTN == 'flash_attn':
|
| 170 |
out = flash_attn.flash_attn_qkvpacked_func(qkv_feats) # [B, N, H, C]
|
|
|
|
|
|
|
|
|
|
| 171 |
else:
|
| 172 |
raise ValueError(f"Unknown attention module: {ATTN}")
|
| 173 |
out = out.reshape(B * N, H, C) # [M, H, C]
|
|
@@ -183,6 +188,11 @@ def sparse_serialized_scaled_dot_product_self_attention(
|
|
| 183 |
cu_seqlens = torch.cat([torch.tensor([0]), torch.cumsum(torch.tensor(seq_lens), dim=0)], dim=0) \
|
| 184 |
.to(qkv.device).int()
|
| 185 |
out = flash_attn.flash_attn_varlen_qkvpacked_func(qkv_feats, cu_seqlens, max(seq_lens)) # [M, H, C]
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 186 |
|
| 187 |
out = out[bwd_indices] # [T, H, C]
|
| 188 |
|
|
|
|
| 9 |
import xformers.ops as xops
|
| 10 |
elif ATTN == 'flash_attn':
|
| 11 |
import flash_attn
|
| 12 |
+
elif ATTN == 'flash_attn_3':
|
| 13 |
+
import flash_attn_interface as flash_attn_3
|
| 14 |
else:
|
| 15 |
raise ValueError(f"Unknown attention module: {ATTN}")
|
| 16 |
|
|
|
|
| 170 |
out = xops.memory_efficient_attention(q, k, v) # [B, N, H, C]
|
| 171 |
elif ATTN == 'flash_attn':
|
| 172 |
out = flash_attn.flash_attn_qkvpacked_func(qkv_feats) # [B, N, H, C]
|
| 173 |
+
elif ATTN == 'flash_attn_3':
|
| 174 |
+
q, k, v = qkv_feats.unbind(dim=2) # [B, N, H, C]
|
| 175 |
+
out = flash_attn_3.flash_attn_func(q, k, v) # [B, N, H, C]
|
| 176 |
else:
|
| 177 |
raise ValueError(f"Unknown attention module: {ATTN}")
|
| 178 |
out = out.reshape(B * N, H, C) # [M, H, C]
|
|
|
|
| 188 |
cu_seqlens = torch.cat([torch.tensor([0]), torch.cumsum(torch.tensor(seq_lens), dim=0)], dim=0) \
|
| 189 |
.to(qkv.device).int()
|
| 190 |
out = flash_attn.flash_attn_varlen_qkvpacked_func(qkv_feats, cu_seqlens, max(seq_lens)) # [M, H, C]
|
| 191 |
+
elif ATTN == 'flash_attn_3':
|
| 192 |
+
q, k, v = qkv_feats.unbind(dim=1) # [M, H, C]
|
| 193 |
+
cu_seqlens = torch.cat([torch.tensor([0]), torch.cumsum(torch.tensor(seq_lens), dim=0)], dim=0) \
|
| 194 |
+
.to(qkv.device).int()
|
| 195 |
+
out = flash_attn_3.flash_attn_varlen_func(q, k, v, cu_seqlens, cu_seqlens, max(seq_lens), max(seq_lens)) # [M, H, C]
|
| 196 |
|
| 197 |
out = out[bwd_indices] # [T, H, C]
|
| 198 |
|
trellis/modules/sparse/attention/windowed_attn.py
CHANGED
|
@@ -8,6 +8,8 @@ if ATTN == 'xformers':
|
|
| 8 |
import xformers.ops as xops
|
| 9 |
elif ATTN == 'flash_attn':
|
| 10 |
import flash_attn
|
|
|
|
|
|
|
| 11 |
else:
|
| 12 |
raise ValueError(f"Unknown attention module: {ATTN}")
|
| 13 |
|
|
@@ -110,6 +112,9 @@ def sparse_windowed_scaled_dot_product_self_attention(
|
|
| 110 |
out = xops.memory_efficient_attention(q, k, v) # [B, N, H, C]
|
| 111 |
elif ATTN == 'flash_attn':
|
| 112 |
out = flash_attn.flash_attn_qkvpacked_func(qkv_feats) # [B, N, H, C]
|
|
|
|
|
|
|
|
|
|
| 113 |
else:
|
| 114 |
raise ValueError(f"Unknown attention module: {ATTN}")
|
| 115 |
out = out.reshape(B * N, H, C) # [M, H, C]
|
|
@@ -125,6 +130,11 @@ def sparse_windowed_scaled_dot_product_self_attention(
|
|
| 125 |
cu_seqlens = torch.cat([torch.tensor([0]), torch.cumsum(torch.tensor(seq_lens), dim=0)], dim=0) \
|
| 126 |
.to(qkv.device).int()
|
| 127 |
out = flash_attn.flash_attn_varlen_qkvpacked_func(qkv_feats, cu_seqlens, max(seq_lens)) # [M, H, C]
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 128 |
|
| 129 |
out = out[bwd_indices] # [T, H, C]
|
| 130 |
|
|
|
|
| 8 |
import xformers.ops as xops
|
| 9 |
elif ATTN == 'flash_attn':
|
| 10 |
import flash_attn
|
| 11 |
+
elif ATTN == 'flash_attn_3':
|
| 12 |
+
import flash_attn_interface as flash_attn_3
|
| 13 |
else:
|
| 14 |
raise ValueError(f"Unknown attention module: {ATTN}")
|
| 15 |
|
|
|
|
| 112 |
out = xops.memory_efficient_attention(q, k, v) # [B, N, H, C]
|
| 113 |
elif ATTN == 'flash_attn':
|
| 114 |
out = flash_attn.flash_attn_qkvpacked_func(qkv_feats) # [B, N, H, C]
|
| 115 |
+
elif ATTN == 'flash_attn_3':
|
| 116 |
+
q, k, v = qkv_feats.unbind(dim=2) # [B, N, H, C]
|
| 117 |
+
out = flash_attn_3.flash_attn_func(q, k, v) # [B, N, H, C]
|
| 118 |
else:
|
| 119 |
raise ValueError(f"Unknown attention module: {ATTN}")
|
| 120 |
out = out.reshape(B * N, H, C) # [M, H, C]
|
|
|
|
| 130 |
cu_seqlens = torch.cat([torch.tensor([0]), torch.cumsum(torch.tensor(seq_lens), dim=0)], dim=0) \
|
| 131 |
.to(qkv.device).int()
|
| 132 |
out = flash_attn.flash_attn_varlen_qkvpacked_func(qkv_feats, cu_seqlens, max(seq_lens)) # [M, H, C]
|
| 133 |
+
elif ATTN == 'flash_attn_3':
|
| 134 |
+
q, k, v = qkv_feats.unbind(dim=1) # [M, H, C]
|
| 135 |
+
cu_seqlens = torch.cat([torch.tensor([0]), torch.cumsum(torch.tensor(seq_lens), dim=0)], dim=0) \
|
| 136 |
+
.to(qkv.device).int()
|
| 137 |
+
out = flash_attn_3.flash_attn_varlen_func(q, k, v, cu_seqlens, cu_seqlens, max(seq_lens), max(seq_lens)) # [M, H, C]
|
| 138 |
|
| 139 |
out = out[bwd_indices] # [T, H, C]
|
| 140 |
|
trellis/modules/sparse/basic.py
CHANGED
|
@@ -133,6 +133,37 @@ class SparseTensor:
|
|
| 133 |
def dim(self) -> int:
|
| 134 |
return len(self.shape)
|
| 135 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 136 |
@property
|
| 137 |
def layout(self) -> List[slice]:
|
| 138 |
return self._layout
|
|
@@ -367,6 +398,40 @@ class SparseTensor:
|
|
| 367 |
feats = torch.cat(feats, dim=0).contiguous()
|
| 368 |
return SparseTensor(feats=feats, coords=coords)
|
| 369 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 370 |
def register_spatial_cache(self, key, value) -> None:
|
| 371 |
"""
|
| 372 |
Register a spatial cache.
|
|
|
|
| 133 |
def dim(self) -> int:
|
| 134 |
return len(self.shape)
|
| 135 |
|
| 136 |
+
@property
|
| 137 |
+
def ndim(self) -> int:
|
| 138 |
+
return self.dim()
|
| 139 |
+
|
| 140 |
+
@property
|
| 141 |
+
def dtype(self):
|
| 142 |
+
return self.feats.dtype
|
| 143 |
+
|
| 144 |
+
@property
|
| 145 |
+
def device(self):
|
| 146 |
+
return self.feats.device
|
| 147 |
+
|
| 148 |
+
@property
|
| 149 |
+
def seqlen(self) -> torch.LongTensor:
|
| 150 |
+
seqlen = self.get_spatial_cache('seqlen')
|
| 151 |
+
if seqlen is None:
|
| 152 |
+
seqlen = torch.tensor([l.stop - l.start for l in self.layout], dtype=torch.long, device=self.device)
|
| 153 |
+
self.register_spatial_cache('seqlen', seqlen)
|
| 154 |
+
return seqlen
|
| 155 |
+
|
| 156 |
+
@property
|
| 157 |
+
def cum_seqlen(self) -> torch.LongTensor:
|
| 158 |
+
cum_seqlen = self.get_spatial_cache('cum_seqlen')
|
| 159 |
+
if cum_seqlen is None:
|
| 160 |
+
cum_seqlen = torch.cat([
|
| 161 |
+
torch.tensor([0], dtype=torch.long, device=self.device),
|
| 162 |
+
self.seqlen.cumsum(dim=0)
|
| 163 |
+
], dim=0)
|
| 164 |
+
self.register_spatial_cache('cum_seqlen', cum_seqlen)
|
| 165 |
+
return cum_seqlen
|
| 166 |
+
|
| 167 |
@property
|
| 168 |
def layout(self) -> List[slice]:
|
| 169 |
return self._layout
|
|
|
|
| 398 |
feats = torch.cat(feats, dim=0).contiguous()
|
| 399 |
return SparseTensor(feats=feats, coords=coords)
|
| 400 |
|
| 401 |
+
def reduce(self, op: str, dim: Optional[Union[int, Tuple[int,...]]] = None, keepdim: bool = False) -> torch.Tensor:
|
| 402 |
+
if isinstance(dim, int):
|
| 403 |
+
dim = (dim,)
|
| 404 |
+
|
| 405 |
+
if op =='mean':
|
| 406 |
+
red = self.feats.mean(dim=dim, keepdim=keepdim)
|
| 407 |
+
elif op =='sum':
|
| 408 |
+
red = self.feats.sum(dim=dim, keepdim=keepdim)
|
| 409 |
+
elif op == 'prod':
|
| 410 |
+
red = self.feats.prod(dim=dim, keepdim=keepdim)
|
| 411 |
+
else:
|
| 412 |
+
raise ValueError(f"Unsupported reduce operation: {op}")
|
| 413 |
+
|
| 414 |
+
if dim is None or 0 in dim:
|
| 415 |
+
return red
|
| 416 |
+
|
| 417 |
+
red = torch.segment_reduce(red, reduce=op, lengths=self.seqlen)
|
| 418 |
+
return red
|
| 419 |
+
|
| 420 |
+
def mean(self, dim: Optional[Union[int, Tuple[int,...]]] = None, keepdim: bool = False) -> torch.Tensor:
|
| 421 |
+
return self.reduce(op='mean', dim=dim, keepdim=keepdim)
|
| 422 |
+
|
| 423 |
+
def sum(self, dim: Optional[Union[int, Tuple[int,...]]] = None, keepdim: bool = False) -> torch.Tensor:
|
| 424 |
+
return self.reduce(op='sum', dim=dim, keepdim=keepdim)
|
| 425 |
+
|
| 426 |
+
def prod(self, dim: Optional[Union[int, Tuple[int,...]]] = None, keepdim: bool = False) -> torch.Tensor:
|
| 427 |
+
return self.reduce(op='prod', dim=dim, keepdim=keepdim)
|
| 428 |
+
|
| 429 |
+
def std(self, dim: Optional[Union[int, Tuple[int,...]]] = None, keepdim: bool = False) -> torch.Tensor:
|
| 430 |
+
mean = self.mean(dim=dim, keepdim=True)
|
| 431 |
+
mean2 = self.replace(self.feats ** 2).mean(dim=dim, keepdim=True)
|
| 432 |
+
std = (mean2 - mean ** 2).sqrt()
|
| 433 |
+
return std
|
| 434 |
+
|
| 435 |
def register_spatial_cache(self, key, value) -> None:
|
| 436 |
"""
|
| 437 |
Register a spatial cache.
|
trellis/pipelines/__init__.py
CHANGED
|
@@ -1,5 +1,6 @@
|
|
| 1 |
from . import samplers
|
| 2 |
from .trellis_image_to_3d import TrellisImageTo3DPipeline, TrellisVGGTTo3DPipeline
|
|
|
|
| 3 |
|
| 4 |
def from_pretrained(path: str):
|
| 5 |
"""
|
|
|
|
| 1 |
from . import samplers
|
| 2 |
from .trellis_image_to_3d import TrellisImageTo3DPipeline, TrellisVGGTTo3DPipeline
|
| 3 |
+
from .trellis_hybrid_pipeline import TrellisHybridPipeline
|
| 4 |
|
| 5 |
def from_pretrained(path: str):
|
| 6 |
"""
|
trellis/pipelines/samplers/classifier_free_guidance_mixin.py
CHANGED
|
@@ -1,12 +1,39 @@
|
|
| 1 |
from typing import *
|
| 2 |
|
| 3 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 4 |
class ClassifierFreeGuidanceSamplerMixin:
|
| 5 |
"""
|
| 6 |
A mixin class for samplers that apply classifier-free guidance.
|
| 7 |
"""
|
| 8 |
|
| 9 |
-
def _inference_model(self, model, x_t, t, cond, neg_cond, cfg_strength, **kwargs):
|
| 10 |
-
|
| 11 |
-
|
| 12 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
from typing import *
|
| 2 |
|
| 3 |
|
| 4 |
+
# class ClassifierFreeGuidanceSamplerMixin:
|
| 5 |
+
# """
|
| 6 |
+
# A mixin class for samplers that apply classifier-free guidance.
|
| 7 |
+
# """
|
| 8 |
+
|
| 9 |
+
# def _inference_model(self, model, x_t, t, cond, neg_cond, cfg_strength, **kwargs):
|
| 10 |
+
# pred = super()._inference_model(model, x_t, t, cond, **kwargs)
|
| 11 |
+
# neg_pred = super()._inference_model(model, x_t, t, neg_cond, **kwargs)
|
| 12 |
+
# return (1 + cfg_strength) * pred - cfg_strength * neg_pred
|
| 13 |
+
|
| 14 |
class ClassifierFreeGuidanceSamplerMixin:
|
| 15 |
"""
|
| 16 |
A mixin class for samplers that apply classifier-free guidance.
|
| 17 |
"""
|
| 18 |
|
| 19 |
+
def _inference_model(self, model, x_t, t, cond, neg_cond, cfg_strength, guidance_rescale=0.0, **kwargs):
|
| 20 |
+
if cfg_strength == 1:
|
| 21 |
+
return super()._inference_model(model, x_t, t, cond, **kwargs)
|
| 22 |
+
elif cfg_strength == 0:
|
| 23 |
+
return super()._inference_model(model, x_t, t, neg_cond, **kwargs)
|
| 24 |
+
else:
|
| 25 |
+
pred_pos = super()._inference_model(model, x_t, t, cond, **kwargs)
|
| 26 |
+
pred_neg = super()._inference_model(model, x_t, t, neg_cond, **kwargs)
|
| 27 |
+
pred = cfg_strength * pred_pos + (1 - cfg_strength) * pred_neg
|
| 28 |
+
|
| 29 |
+
# CFG rescale
|
| 30 |
+
if guidance_rescale > 0:
|
| 31 |
+
x_0_pos = self._pred_to_xstart(x_t, t, pred_pos)
|
| 32 |
+
x_0_cfg = self._pred_to_xstart(x_t, t, pred)
|
| 33 |
+
std_pos = x_0_pos.std(dim=list(range(1, x_0_pos.ndim)), keepdim=True)
|
| 34 |
+
std_cfg = x_0_cfg.std(dim=list(range(1, x_0_cfg.ndim)), keepdim=True)
|
| 35 |
+
x_0_rescaled = x_0_cfg * (std_pos / std_cfg)
|
| 36 |
+
x_0 = guidance_rescale * x_0_rescaled + (1 - guidance_rescale) * x_0_cfg
|
| 37 |
+
pred = self._xstart_to_pred(x_t, t, x_0)
|
| 38 |
+
|
| 39 |
+
return pred
|
trellis/pipelines/samplers/flow_euler.py
CHANGED
|
@@ -50,6 +50,11 @@ class FlowEulerSampler(Sampler):
|
|
| 50 |
assert x_0.shape == x_t.shape
|
| 51 |
return (x_t - (1 - self.sigma_min) * x_0) / (self.sigma_min + (1 - self.sigma_min) * t)
|
| 52 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 53 |
|
| 54 |
def _inference_model(self, model, x_t, t, cond=None, **kwargs):
|
| 55 |
t = torch.tensor([1000 * t] * x_t.shape[0], device=x_t.device, dtype=torch.float32)
|
|
@@ -178,73 +183,6 @@ class FlowEulerSampler(Sampler):
|
|
| 178 |
torch.cuda.empty_cache()
|
| 179 |
return edict({"pred_x_prev": pred_x_prev, "pred_x_0": pred_x_0, "pred_eps": pred_eps})
|
| 180 |
|
| 181 |
-
def sample_slat_once_opt_delta_v(
|
| 182 |
-
self,
|
| 183 |
-
model,
|
| 184 |
-
slat_decoder_gs,
|
| 185 |
-
slat_decoder_mesh,
|
| 186 |
-
std,
|
| 187 |
-
mean,
|
| 188 |
-
dreamsim_model,
|
| 189 |
-
learning_rate,
|
| 190 |
-
input_images,
|
| 191 |
-
extrinsics,
|
| 192 |
-
intrinsics,
|
| 193 |
-
x_t,
|
| 194 |
-
t: float,
|
| 195 |
-
t_prev: float,
|
| 196 |
-
cond: Optional[Any] = None,
|
| 197 |
-
**kwargs
|
| 198 |
-
):
|
| 199 |
-
"""
|
| 200 |
-
Sample x_{t-1} from the model using Euler method.
|
| 201 |
-
|
| 202 |
-
Args:
|
| 203 |
-
model: The model to sample from.
|
| 204 |
-
x_t: The [N x C x ...] tensor of noisy inputs at time t.
|
| 205 |
-
t: The current timestep.
|
| 206 |
-
t_prev: The previous timestep.
|
| 207 |
-
cond: conditional information.
|
| 208 |
-
**kwargs: Additional arguments for model inference.
|
| 209 |
-
|
| 210 |
-
Returns:
|
| 211 |
-
a dict containing the following
|
| 212 |
-
- 'pred_x_prev': x_{t-1}.
|
| 213 |
-
- 'pred_x_0': a prediction of x_0.
|
| 214 |
-
"""
|
| 215 |
-
torch.cuda.empty_cache()
|
| 216 |
-
with torch.no_grad():
|
| 217 |
-
pred_x_0, pred_eps, pred_v = self._get_model_prediction(model, x_t, t, cond, **kwargs)
|
| 218 |
-
pred_v_opt_feat = torch.nn.Parameter(pred_v.feats.detach().clone())
|
| 219 |
-
optimizer = torch.optim.Adam([pred_v_opt_feat], betas=(0.5, 0.9), lr=learning_rate)
|
| 220 |
-
pred_v_opt = sp.SparseTensor(feats=pred_v_opt_feat, coords=pred_v.coords)
|
| 221 |
-
total_steps = 5
|
| 222 |
-
input_images = F.interpolate(input_images, size=(259, 259), mode='bilinear', align_corners=False)
|
| 223 |
-
with tqdm(total=total_steps, disable=True, desc='Appearance (opt): optimizing') as pbar:
|
| 224 |
-
for step in range(total_steps):
|
| 225 |
-
optimizer.zero_grad()
|
| 226 |
-
pred_x_0, _ = self._v_to_xstart_eps(x_t=x_t, t=t, v=pred_v_opt)
|
| 227 |
-
pred_gs = slat_decoder_gs(pred_x_0 * std + mean)
|
| 228 |
-
# pred_mesh = slat_decoder_mesh(pred_x_0 * std + mean)
|
| 229 |
-
rend_gs = render_utils.render_frames(pred_gs[0], extrinsics, intrinsics, {'resolution': 259, 'bg_color': (0, 0, 0)}, need_depth=True, opt=True)['color']
|
| 230 |
-
# rend_mesh = render_utils.render_frames_opt(pred_mesh[0], extrinsics, intrinsics, {'resolution': 518, 'bg_color': (0, 0, 0)}, need_depth=True, opt=True)['color']
|
| 231 |
-
rend_gs = torch.stack(rend_gs, dim=0)
|
| 232 |
-
loss_gs = loss_utils.l1_loss(rend_gs, input_images, size_average=False).mean(dim=(1,2,3)) + \
|
| 233 |
-
(1 - loss_utils.ssim(rend_gs, input_images, size_average=False)) + \
|
| 234 |
-
loss_utils.lpips(rend_gs, input_images, size_average=False).mean(dim=(1,2,3)) + \
|
| 235 |
-
dreamsim_model(rend_gs, input_images)
|
| 236 |
-
loss_gs = loss_gs[loss_gs <= 0.8].mean()
|
| 237 |
-
# loss_gs = (1 - loss_utils.ssim(rend_gs, input_images)) + loss_utils.lpips(rend_gs, input_images) + dreamsim_model(rend_gs, input_images).mean()
|
| 238 |
-
# loss_mesh = loss_utils.l1_loss(rend_mesh, input_images) + 0.2 * (1 - loss_utils.ssim(rend_mesh, input_images)) + 0.2 * loss_utils.lpips(rend_mesh, input_images)
|
| 239 |
-
loss = loss_gs + 0.2 * loss_utils.l1_loss(pred_v_opt_feat, pred_v.feats)
|
| 240 |
-
loss.backward()
|
| 241 |
-
optimizer.step()
|
| 242 |
-
pbar.set_postfix({'loss': loss.item()})
|
| 243 |
-
pbar.update()
|
| 244 |
-
|
| 245 |
-
pred_x_prev = x_t - (t - t_prev) * pred_v_opt.detach()
|
| 246 |
-
torch.cuda.empty_cache()
|
| 247 |
-
return edict({"pred_x_prev": pred_x_prev, "pred_x_0": pred_x_0, "pred_eps": pred_eps})
|
| 248 |
|
| 249 |
def sample_opt(
|
| 250 |
self,
|
|
@@ -342,66 +280,6 @@ class FlowEulerSampler(Sampler):
|
|
| 342 |
ret.samples = sample
|
| 343 |
return ret
|
| 344 |
|
| 345 |
-
def sample_slat_opt_delta_v(
|
| 346 |
-
self,
|
| 347 |
-
model,
|
| 348 |
-
slat_decoder_gs,
|
| 349 |
-
slat_decoder_mesh,
|
| 350 |
-
std,
|
| 351 |
-
mean,
|
| 352 |
-
dreamsim_model,
|
| 353 |
-
apperance_learning_rate,
|
| 354 |
-
start_t,
|
| 355 |
-
input_images,
|
| 356 |
-
extrinsics,
|
| 357 |
-
intrinsics,
|
| 358 |
-
noise,
|
| 359 |
-
cond: Optional[Any] = None,
|
| 360 |
-
steps: int = 50,
|
| 361 |
-
rescale_t: float = 1.0,
|
| 362 |
-
verbose: bool = True,
|
| 363 |
-
**kwargs
|
| 364 |
-
):
|
| 365 |
-
"""
|
| 366 |
-
Generate samples from the model using Euler method.
|
| 367 |
-
|
| 368 |
-
Args:
|
| 369 |
-
model: The model to sample from.
|
| 370 |
-
noise: The initial noise tensor.
|
| 371 |
-
cond: conditional information.
|
| 372 |
-
steps: The number of steps to sample.
|
| 373 |
-
rescale_t: The rescale factor for t.
|
| 374 |
-
verbose: If True, show a progress bar.
|
| 375 |
-
**kwargs: Additional arguments for model_inference.
|
| 376 |
-
|
| 377 |
-
Returns:
|
| 378 |
-
a dict containing the following
|
| 379 |
-
- 'samples': the model samples.
|
| 380 |
-
- 'pred_x_t': a list of prediction of x_t.
|
| 381 |
-
- 'pred_x_0': a list of prediction of x_0.
|
| 382 |
-
"""
|
| 383 |
-
sample = noise
|
| 384 |
-
t_seq = np.linspace(1, 0, steps + 1)
|
| 385 |
-
t_seq = rescale_t * t_seq / (1 + (rescale_t - 1) * t_seq)
|
| 386 |
-
t_pairs = list((t_seq[i], t_seq[i + 1]) for i in range(steps))
|
| 387 |
-
ret = edict({"samples": None, "pred_x_t": [], "pred_x_0": []})
|
| 388 |
-
# def cosine_anealing(step, total_steps, start_lr, end_lr):
|
| 389 |
-
# return end_lr + 0.5 * (start_lr - end_lr) * (1 + np.cos(np.pi * step / total_steps))
|
| 390 |
-
for i, (t, t_prev) in enumerate(tqdm(t_pairs, desc="Sampling", disable=not verbose)):
|
| 391 |
-
if t > start_t:
|
| 392 |
-
out = self.sample_once(model, sample, t, t_prev, cond, **kwargs)
|
| 393 |
-
sample = out.pred_x_prev
|
| 394 |
-
ret.pred_x_t.append(out.pred_x_prev)
|
| 395 |
-
ret.pred_x_0.append(out.pred_x_0)
|
| 396 |
-
else:
|
| 397 |
-
# learning_rate = cosine_anealing(i - int(np.where(t_seq <= start_t)[0].min()), int(steps - np.where(t_seq <= start_t)[0].min()), apperance_learning_rate, 1e-5)
|
| 398 |
-
learning_rate = apperance_learning_rate
|
| 399 |
-
out = self.sample_slat_once_opt_delta_v(model, slat_decoder_gs, slat_decoder_mesh, std, mean, dreamsim_model, learning_rate, input_images, extrinsics, intrinsics, sample, t, t_prev, cond, **kwargs)
|
| 400 |
-
sample = out.pred_x_prev
|
| 401 |
-
ret.pred_x_t.append(out.pred_x_prev)
|
| 402 |
-
ret.pred_x_0.append(out.pred_x_0)
|
| 403 |
-
ret.samples = sample
|
| 404 |
-
return ret
|
| 405 |
|
| 406 |
@torch.no_grad()
|
| 407 |
def sample(
|
|
@@ -746,7 +624,7 @@ class FlowEulerCfgSampler(ClassifierFreeGuidanceSamplerMixin, FlowEulerSampler):
|
|
| 746 |
return super().sample(model, noise, cond, steps, rescale_t, verbose, neg_cond=neg_cond, cfg_strength=cfg_strength, **kwargs)
|
| 747 |
|
| 748 |
|
| 749 |
-
class FlowEulerGuidanceIntervalSampler(GuidanceIntervalSamplerMixin, FlowEulerSampler):
|
| 750 |
"""
|
| 751 |
Generate samples from a flow-matching model using Euler sampling with classifier-free guidance and interval.
|
| 752 |
"""
|
|
@@ -864,53 +742,6 @@ class FlowEulerGuidanceIntervalSampler(GuidanceIntervalSamplerMixin, FlowEulerSa
|
|
| 864 |
return super().sample_ss_opt_delta_v(model, ss_decoder, ss_learning_rate, ss_start_t, ss, noise, cond, steps, rescale_t, verbose, neg_cond=neg_cond, cfg_strength=cfg_strength, cfg_interval=cfg_interval, **kwargs)
|
| 865 |
|
| 866 |
|
| 867 |
-
def sample_slat_opt_delta_v(
|
| 868 |
-
self,
|
| 869 |
-
model,
|
| 870 |
-
slat_decoder_gs,
|
| 871 |
-
slat_decoder_mesh,
|
| 872 |
-
std,
|
| 873 |
-
mean,
|
| 874 |
-
dreamsim_model,
|
| 875 |
-
apperance_learning_rate,
|
| 876 |
-
start_t,
|
| 877 |
-
input_images,
|
| 878 |
-
extrinsics,
|
| 879 |
-
intrinsics,
|
| 880 |
-
noise,
|
| 881 |
-
cond,
|
| 882 |
-
neg_cond,
|
| 883 |
-
steps: int = 50,
|
| 884 |
-
rescale_t: float = 1.0,
|
| 885 |
-
cfg_strength: float = 3.0,
|
| 886 |
-
cfg_interval: Tuple[float, float] = (0.0, 1.0),
|
| 887 |
-
verbose: bool = True,
|
| 888 |
-
**kwargs
|
| 889 |
-
):
|
| 890 |
-
"""
|
| 891 |
-
Generate samples from the model using Euler method.
|
| 892 |
-
|
| 893 |
-
Args:
|
| 894 |
-
model: The model to sample from.
|
| 895 |
-
noise: The initial noise tensor.
|
| 896 |
-
cond: conditional information.
|
| 897 |
-
neg_cond: negative conditional information.
|
| 898 |
-
steps: The number of steps to sample.
|
| 899 |
-
rescale_t: The rescale factor for t.
|
| 900 |
-
cfg_strength: The strength of classifier-free guidance.
|
| 901 |
-
cfg_interval: The interval for classifier-free guidance.
|
| 902 |
-
verbose: If True, show a progress bar.
|
| 903 |
-
**kwargs: Additional arguments for model_inference.
|
| 904 |
-
|
| 905 |
-
Returns:
|
| 906 |
-
a dict containing the following
|
| 907 |
-
- 'samples': the model samples.
|
| 908 |
-
- 'pred_x_t': a list of prediction of x_t.
|
| 909 |
-
- 'pred_x_0': a list of prediction of x_0.
|
| 910 |
-
"""
|
| 911 |
-
return super().sample_slat_opt_delta_v(model, slat_decoder_gs, slat_decoder_mesh, std, mean, dreamsim_model, apperance_learning_rate, start_t, input_images, extrinsics, intrinsics,noise, cond, steps, rescale_t, verbose, neg_cond=neg_cond, cfg_strength=cfg_strength, cfg_interval=cfg_interval, **kwargs)
|
| 912 |
-
|
| 913 |
-
|
| 914 |
class LatentMatchGuidanceIntervalSampler(GuidanceIntervalSamplerMixin, LatentMatchSampler):
|
| 915 |
"""
|
| 916 |
Generate samples from a flow-matching model using Euler sampling with classifier-free guidance and interval.
|
|
|
|
| 50 |
assert x_0.shape == x_t.shape
|
| 51 |
return (x_t - (1 - self.sigma_min) * x_0) / (self.sigma_min + (1 - self.sigma_min) * t)
|
| 52 |
|
| 53 |
+
def _pred_to_xstart(self, x_t, t, pred):
|
| 54 |
+
return (1 - self.sigma_min) * x_t - (self.sigma_min + (1 - self.sigma_min) * t) * pred
|
| 55 |
+
|
| 56 |
+
def _xstart_to_pred(self, x_t, t, x_0):
|
| 57 |
+
return ((1 - self.sigma_min) * x_t - x_0) / (self.sigma_min + (1 - self.sigma_min) * t)
|
| 58 |
|
| 59 |
def _inference_model(self, model, x_t, t, cond=None, **kwargs):
|
| 60 |
t = torch.tensor([1000 * t] * x_t.shape[0], device=x_t.device, dtype=torch.float32)
|
|
|
|
| 183 |
torch.cuda.empty_cache()
|
| 184 |
return edict({"pred_x_prev": pred_x_prev, "pred_x_0": pred_x_0, "pred_eps": pred_eps})
|
| 185 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 186 |
|
| 187 |
def sample_opt(
|
| 188 |
self,
|
|
|
|
| 280 |
ret.samples = sample
|
| 281 |
return ret
|
| 282 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 283 |
|
| 284 |
@torch.no_grad()
|
| 285 |
def sample(
|
|
|
|
| 624 |
return super().sample(model, noise, cond, steps, rescale_t, verbose, neg_cond=neg_cond, cfg_strength=cfg_strength, **kwargs)
|
| 625 |
|
| 626 |
|
| 627 |
+
class FlowEulerGuidanceIntervalSampler(GuidanceIntervalSamplerMixin, ClassifierFreeGuidanceSamplerMixin, FlowEulerSampler):
|
| 628 |
"""
|
| 629 |
Generate samples from a flow-matching model using Euler sampling with classifier-free guidance and interval.
|
| 630 |
"""
|
|
|
|
| 742 |
return super().sample_ss_opt_delta_v(model, ss_decoder, ss_learning_rate, ss_start_t, ss, noise, cond, steps, rescale_t, verbose, neg_cond=neg_cond, cfg_strength=cfg_strength, cfg_interval=cfg_interval, **kwargs)
|
| 743 |
|
| 744 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 745 |
class LatentMatchGuidanceIntervalSampler(GuidanceIntervalSamplerMixin, LatentMatchSampler):
|
| 746 |
"""
|
| 747 |
Generate samples from a flow-matching model using Euler sampling with classifier-free guidance and interval.
|
trellis/pipelines/samplers/guidance_interval_mixin.py
CHANGED
|
@@ -1,15 +1,26 @@
|
|
| 1 |
from typing import *
|
| 2 |
|
| 3 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 4 |
class GuidanceIntervalSamplerMixin:
|
| 5 |
"""
|
| 6 |
A mixin class for samplers that apply classifier-free guidance with interval.
|
| 7 |
"""
|
| 8 |
|
| 9 |
-
def _inference_model(self, model, x_t, t, cond,
|
| 10 |
if cfg_interval[0] <= t <= cfg_interval[1]:
|
| 11 |
-
|
| 12 |
-
neg_pred = super()._inference_model(model, x_t, t, neg_cond, **kwargs)
|
| 13 |
-
return (1 + cfg_strength) * pred - cfg_strength * neg_pred
|
| 14 |
else:
|
| 15 |
-
return super()._inference_model(model, x_t, t, cond, **kwargs)
|
|
|
|
| 1 |
from typing import *
|
| 2 |
|
| 3 |
|
| 4 |
+
# class GuidanceIntervalSamplerMixin:
|
| 5 |
+
# """
|
| 6 |
+
# A mixin class for samplers that apply classifier-free guidance with interval.
|
| 7 |
+
# """
|
| 8 |
+
|
| 9 |
+
# def _inference_model(self, model, x_t, t, cond, neg_cond, cfg_strength, cfg_interval, **kwargs):
|
| 10 |
+
# if cfg_interval[0] <= t <= cfg_interval[1]:
|
| 11 |
+
# pred = super()._inference_model(model, x_t, t, cond, **kwargs)
|
| 12 |
+
# neg_pred = super()._inference_model(model, x_t, t, neg_cond, **kwargs)
|
| 13 |
+
# return (1 + cfg_strength) * pred - cfg_strength * neg_pred
|
| 14 |
+
# else:
|
| 15 |
+
# return super()._inference_model(model, x_t, t, cond, **kwargs)
|
| 16 |
+
|
| 17 |
class GuidanceIntervalSamplerMixin:
|
| 18 |
"""
|
| 19 |
A mixin class for samplers that apply classifier-free guidance with interval.
|
| 20 |
"""
|
| 21 |
|
| 22 |
+
def _inference_model(self, model, x_t, t, cond, cfg_strength, cfg_interval, **kwargs):
|
| 23 |
if cfg_interval[0] <= t <= cfg_interval[1]:
|
| 24 |
+
return super()._inference_model(model, x_t, t, cond, cfg_strength=cfg_strength, **kwargs)
|
|
|
|
|
|
|
| 25 |
else:
|
| 26 |
+
return super()._inference_model(model, x_t, t, cond, cfg_strength=1, **kwargs)
|
trellis/pipelines/trellis_hybrid_pipeline.py
ADDED
|
@@ -0,0 +1,694 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
Hybrid pipeline: ReconViaGen VGGT-based SS stage + TRELLIS.2 shape_slat / tex_slat stages.
|
| 3 |
+
|
| 4 |
+
Stage 1 – Sparse Structure (SS) : TrellisVGGTTo3DPipeline (ReconViaGen, VGGT-conditioned)
|
| 5 |
+
Stage 2 – Shape SLat : Trellis2ImageTo3DPipeline (TRELLIS.2, DINOv3-conditioned)
|
| 6 |
+
Stage 3 – Texture SLat : Trellis2ImageTo3DPipeline (TRELLIS.2, DINOv3-conditioned)
|
| 7 |
+
Stage 4 – Decode / GLB export : Trellis2ImageTo3DPipeline (TRELLIS.2, o_voxel)
|
| 8 |
+
"""
|
| 9 |
+
|
| 10 |
+
import sys
|
| 11 |
+
import os
|
| 12 |
+
|
| 13 |
+
# ── Make trellis2 importable from wheels/TRELLIS.2 ───────────────────────────
|
| 14 |
+
_TRELLIS2_ROOT = os.path.normpath(os.path.join(os.path.dirname(__file__), '..', '..', 'wheels', 'TRELLIS.2'))
|
| 15 |
+
|
| 16 |
+
if _TRELLIS2_ROOT not in sys.path:
|
| 17 |
+
sys.path.insert(0, _TRELLIS2_ROOT)
|
| 18 |
+
# o_voxel is installed into the conda env (no extra sys.path needed)
|
| 19 |
+
|
| 20 |
+
# ── Standard imports ──────────────────────────────────────────────────────────
|
| 21 |
+
from typing import *
|
| 22 |
+
import torch
|
| 23 |
+
import numpy as np
|
| 24 |
+
from PIL import Image
|
| 25 |
+
|
| 26 |
+
from trellis2.pipelines import Trellis2ImageTo3DPipeline
|
| 27 |
+
from trellis2.modules.sparse import SparseTensor
|
| 28 |
+
|
| 29 |
+
from .trellis_image_to_3d import TrellisVGGTTo3DPipeline
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
class TrellisHybridPipeline:
|
| 33 |
+
"""
|
| 34 |
+
Hybrid 3-stage pipeline:
|
| 35 |
+
|
| 36 |
+
1. SS – uses ReconViaGen's TrellisVGGTTo3DPipeline (VGGT + DINOv2 features)
|
| 37 |
+
2. Shape – uses TRELLIS.2's Trellis2ImageTo3DPipeline (shape_slat, DINOv3 features)
|
| 38 |
+
3. Texture – uses TRELLIS.2's Trellis2ImageTo3DPipeline (tex_slat, DINOv3 features)
|
| 39 |
+
|
| 40 |
+
The two sub-pipelines are kept separate so each can be moved between CPU/GPU
|
| 41 |
+
independently when ``low_vram=True``.
|
| 42 |
+
|
| 43 |
+
Args:
|
| 44 |
+
vggt_pipeline : Loaded TrellisVGGTTo3DPipeline (provides SS components).
|
| 45 |
+
trellis2_pipeline: Loaded Trellis2ImageTo3DPipeline (provides shape/tex slat + decode).
|
| 46 |
+
"""
|
| 47 |
+
|
| 48 |
+
def __init__(
|
| 49 |
+
self,
|
| 50 |
+
vggt_pipeline: TrellisVGGTTo3DPipeline,
|
| 51 |
+
trellis2_pipeline: Trellis2ImageTo3DPipeline,
|
| 52 |
+
low_vram: bool = False,
|
| 53 |
+
):
|
| 54 |
+
self.vggt_pipeline = vggt_pipeline
|
| 55 |
+
self.trellis2_pipeline = trellis2_pipeline
|
| 56 |
+
self.low_vram = low_vram
|
| 57 |
+
if low_vram:
|
| 58 |
+
trellis2_pipeline.low_vram = True
|
| 59 |
+
|
| 60 |
+
# ── Convenience wrappers ──────────────────────────────────────────────────
|
| 61 |
+
|
| 62 |
+
def _vggt_models_to(self, device) -> None:
|
| 63 |
+
"""Move vggt_pipeline inference models to *device* (for low-VRAM mode)."""
|
| 64 |
+
vp = self.vggt_pipeline
|
| 65 |
+
if hasattr(vp, 'VGGT_model') and vp.VGGT_model is not None:
|
| 66 |
+
vp.VGGT_model.to(device)
|
| 67 |
+
for model in vp.models.values():
|
| 68 |
+
model.to(device)
|
| 69 |
+
|
| 70 |
+
@property
|
| 71 |
+
def device(self):
|
| 72 |
+
return self.vggt_pipeline.device
|
| 73 |
+
|
| 74 |
+
@property
|
| 75 |
+
def pbr_attr_layout(self):
|
| 76 |
+
return self.trellis2_pipeline.pbr_attr_layout
|
| 77 |
+
|
| 78 |
+
def preprocess_image(self, image: Image.Image) -> Image.Image:
|
| 79 |
+
"""
|
| 80 |
+
Preprocess the input image — mirrors Trellis2ImageTo3DPipeline.preprocess_image:
|
| 81 |
+
if the image already has a real alpha channel it is used directly;
|
| 82 |
+
otherwise the background is removed with the rembg_model (BiRefNet).
|
| 83 |
+
The foreground is then cropped and composited onto a black background.
|
| 84 |
+
"""
|
| 85 |
+
t2p = self.trellis2_pipeline
|
| 86 |
+
|
| 87 |
+
# Use existing alpha if present, otherwise run background removal
|
| 88 |
+
has_alpha = False
|
| 89 |
+
if image.mode == 'RGBA':
|
| 90 |
+
alpha = np.array(image)[:, :, 3]
|
| 91 |
+
if not np.all(alpha == 255):
|
| 92 |
+
has_alpha = True
|
| 93 |
+
|
| 94 |
+
max_size = max(image.size)
|
| 95 |
+
scale = min(1, 1024 / max_size)
|
| 96 |
+
if scale < 1:
|
| 97 |
+
image = image.resize(
|
| 98 |
+
(int(image.width * scale), int(image.height * scale)),
|
| 99 |
+
Image.Resampling.LANCZOS,
|
| 100 |
+
)
|
| 101 |
+
|
| 102 |
+
if has_alpha:
|
| 103 |
+
output = image
|
| 104 |
+
else:
|
| 105 |
+
image = image.convert('RGB')
|
| 106 |
+
if t2p.low_vram:
|
| 107 |
+
t2p.rembg_model.to(t2p.device)
|
| 108 |
+
output = t2p.rembg_model(image)
|
| 109 |
+
if t2p.low_vram:
|
| 110 |
+
t2p.rembg_model.cpu()
|
| 111 |
+
|
| 112 |
+
output_np = np.array(output)
|
| 113 |
+
alpha = output_np[:, :, 3]
|
| 114 |
+
bbox = np.argwhere(alpha > 0.8 * 255)
|
| 115 |
+
bbox = np.min(bbox[:, 1]), np.min(bbox[:, 0]), np.max(bbox[:, 1]), np.max(bbox[:, 0])
|
| 116 |
+
center = (bbox[0] + bbox[2]) / 2, (bbox[1] + bbox[3]) / 2
|
| 117 |
+
size = max(bbox[2] - bbox[0], bbox[3] - bbox[1])
|
| 118 |
+
size = int(size * 1)
|
| 119 |
+
bbox = center[0] - size // 2, center[1] - size // 2, center[0] + size // 2, center[1] + size // 2
|
| 120 |
+
output = output.crop(bbox)
|
| 121 |
+
output = np.array(output).astype(np.float32) / 255
|
| 122 |
+
output = output[:, :, :3] * output[:, :, 3:4]
|
| 123 |
+
return Image.fromarray((output * 255).astype(np.uint8))
|
| 124 |
+
|
| 125 |
+
# ── Stage 1: SS via mesh voxelisation ────────────────────────────────────
|
| 126 |
+
|
| 127 |
+
@staticmethod
|
| 128 |
+
def _mesh_to_surface_coords(
|
| 129 |
+
mesh_result,
|
| 130 |
+
target_res: int,
|
| 131 |
+
device: torch.device,
|
| 132 |
+
simplify_ratio: float = 0.95,
|
| 133 |
+
fill_holes_max_hole_size: float = 0.04,
|
| 134 |
+
fill_holes_max_hole_nbe: int = 55,
|
| 135 |
+
fill_holes_resolution: int = 256,
|
| 136 |
+
fill_holes_num_views: int = 200,
|
| 137 |
+
) -> torch.Tensor:
|
| 138 |
+
"""
|
| 139 |
+
Convert a MeshExtractResult (surface triangle mesh in [-0.5, 0.5]^3) to
|
| 140 |
+
surface-only voxel coords [N, 4] = [batch_idx, x, y, z] in [0, target_res).
|
| 141 |
+
|
| 142 |
+
Pipeline
|
| 143 |
+
--------
|
| 144 |
+
1. Decimate + Fill holes – postprocess_mesh (pyvista decimate → GPU _fill_holes),
|
| 145 |
+
same approach as postprocessing_utils.py lines 426-437
|
| 146 |
+
2. Voxelize – Open3D surface voxelization directly at target_res
|
| 147 |
+
"""
|
| 148 |
+
from trellis.utils.postprocessing_utils import postprocess_mesh
|
| 149 |
+
import open3d as o3d
|
| 150 |
+
|
| 151 |
+
verts_np = mesh_result.vertices.cpu().numpy()
|
| 152 |
+
faces_np = mesh_result.faces.cpu().numpy()
|
| 153 |
+
|
| 154 |
+
# 1. ── Decimate + fill holes (mirrors postprocessing_utils.py L426-437) ─
|
| 155 |
+
verts_np, faces_np = postprocess_mesh(
|
| 156 |
+
verts_np, faces_np,
|
| 157 |
+
simplify=True,
|
| 158 |
+
simplify_ratio=simplify_ratio,
|
| 159 |
+
fill_holes=True,
|
| 160 |
+
fill_holes_max_hole_size=fill_holes_max_hole_size,
|
| 161 |
+
fill_holes_max_hole_nbe=fill_holes_max_hole_nbe,
|
| 162 |
+
fill_holes_resolution=fill_holes_resolution,
|
| 163 |
+
fill_holes_num_views=fill_holes_num_views,
|
| 164 |
+
)
|
| 165 |
+
|
| 166 |
+
# 2. ── Voxelize directly at target_res ───────────────────────────────
|
| 167 |
+
o3d_mesh = o3d.geometry.TriangleMesh()
|
| 168 |
+
o3d_mesh.vertices = o3d.utility.Vector3dVector(
|
| 169 |
+
np.clip(verts_np.astype(np.float64), -0.5 + 1e-6, 0.5 - 1e-6)
|
| 170 |
+
)
|
| 171 |
+
o3d_mesh.triangles = o3d.utility.Vector3iVector(faces_np.astype(np.int32))
|
| 172 |
+
voxel_grid = o3d.geometry.VoxelGrid.create_from_triangle_mesh_within_bounds(
|
| 173 |
+
o3d_mesh,
|
| 174 |
+
voxel_size=1.0 / target_res,
|
| 175 |
+
min_bound=(-0.5, -0.5, -0.5),
|
| 176 |
+
max_bound=( 0.5, 0.5, 0.5),
|
| 177 |
+
)
|
| 178 |
+
grid_idx = np.array(
|
| 179 |
+
[voxel.grid_index for voxel in voxel_grid.get_voxels()], dtype=np.int32
|
| 180 |
+
) # [N, 3] in [0, target_res)
|
| 181 |
+
|
| 182 |
+
if len(grid_idx) == 0:
|
| 183 |
+
return torch.zeros((0, 4), dtype=torch.int32, device=device)
|
| 184 |
+
|
| 185 |
+
idx = torch.tensor(grid_idx, dtype=torch.int32, device=device)
|
| 186 |
+
batch = torch.zeros(len(idx), 1, dtype=torch.int32, device=device)
|
| 187 |
+
coords = torch.cat([batch, idx], dim=1) # [N, 4]: [0, x, y, z]
|
| 188 |
+
return coords
|
| 189 |
+
|
| 190 |
+
@torch.no_grad()
|
| 191 |
+
def _run_ss_stage(
|
| 192 |
+
self,
|
| 193 |
+
images: List[Image.Image],
|
| 194 |
+
target_ss_res: int,
|
| 195 |
+
ss_sampler_params: dict,
|
| 196 |
+
slat_sampler_params: dict,
|
| 197 |
+
) -> torch.Tensor:
|
| 198 |
+
"""
|
| 199 |
+
Generate a rough mesh via vggt_pipeline, then voxelise it into
|
| 200 |
+
surface-only coords at target_ss_res^3 for the downstream shape/tex stages.
|
| 201 |
+
|
| 202 |
+
Returns:
|
| 203 |
+
coords : (N, 4) int tensor [batch_idx, x, y, z] in [0, target_ss_res)
|
| 204 |
+
"""
|
| 205 |
+
vp = self.vggt_pipeline
|
| 206 |
+
# vp.device is dynamic (inferred from model params), so when models are on
|
| 207 |
+
# CPU it returns 'cpu'. Hardcode the target cuda device instead.
|
| 208 |
+
cuda_device = torch.device('cuda')
|
| 209 |
+
|
| 210 |
+
if self.low_vram:
|
| 211 |
+
self._vggt_models_to(cuda_device)
|
| 212 |
+
|
| 213 |
+
outputs, _, _ = vp.run(
|
| 214 |
+
image=images,
|
| 215 |
+
formats=["mesh"],
|
| 216 |
+
preprocess_image=False,
|
| 217 |
+
sparse_structure_sampler_params=ss_sampler_params,
|
| 218 |
+
slat_sampler_params=slat_sampler_params,
|
| 219 |
+
)
|
| 220 |
+
mesh_result = outputs["mesh"][0]
|
| 221 |
+
coords = self._mesh_to_surface_coords(mesh_result, target_ss_res, cuda_device)
|
| 222 |
+
|
| 223 |
+
if self.low_vram:
|
| 224 |
+
self._vggt_models_to('cpu')
|
| 225 |
+
torch.cuda.empty_cache()
|
| 226 |
+
|
| 227 |
+
return coords
|
| 228 |
+
|
| 229 |
+
@torch.no_grad()
|
| 230 |
+
def _run_ss_stage_direct(
|
| 231 |
+
self,
|
| 232 |
+
images: List[Image.Image],
|
| 233 |
+
target_ss_res: int,
|
| 234 |
+
ss_sampler_params: dict,
|
| 235 |
+
) -> torch.Tensor:
|
| 236 |
+
"""
|
| 237 |
+
Run only ReconViaGen's sparse structure diffusion stage to obtain coords
|
| 238 |
+
directly, without proceeding to the SLAT/mesh stage.
|
| 239 |
+
|
| 240 |
+
Returns:
|
| 241 |
+
coords : (N, 4) int tensor [batch_idx, x, y, z] in [0, target_ss_res)
|
| 242 |
+
"""
|
| 243 |
+
vp = self.vggt_pipeline
|
| 244 |
+
cuda_device = torch.device('cuda')
|
| 245 |
+
|
| 246 |
+
if self.low_vram:
|
| 247 |
+
self._vggt_models_to(cuda_device)
|
| 248 |
+
|
| 249 |
+
with torch.no_grad():
|
| 250 |
+
with torch.cuda.amp.autocast(dtype=vp.VGGT_dtype):
|
| 251 |
+
aggregated_tokens_list, _ = vp.vggt_feat(images)
|
| 252 |
+
b, n, _, _ = aggregated_tokens_list[0].shape
|
| 253 |
+
image_cond = vp.encode_image(images).reshape(b, n, -1, 1024)
|
| 254 |
+
ss_cond = vp.get_ss_cond(image_cond[:, :, 5:], aggregated_tokens_list, 1)
|
| 255 |
+
|
| 256 |
+
ss_flow_model = vp.models['sparse_structure_flow_model']
|
| 257 |
+
sampler_params = {**vp.sparse_structure_sampler_params, **ss_sampler_params}
|
| 258 |
+
reso = ss_flow_model.resolution
|
| 259 |
+
ss_noise = torch.randn(1, ss_flow_model.in_channels, reso, reso, reso).to(cuda_device)
|
| 260 |
+
|
| 261 |
+
ss_latent = vp.sparse_structure_sampler.sample(
|
| 262 |
+
ss_flow_model,
|
| 263 |
+
ss_noise,
|
| 264 |
+
**ss_cond,
|
| 265 |
+
**sampler_params,
|
| 266 |
+
verbose=True,
|
| 267 |
+
).samples
|
| 268 |
+
|
| 269 |
+
decoder = vp.models['sparse_structure_decoder']
|
| 270 |
+
decoded = decoder(ss_latent) > 0
|
| 271 |
+
if target_ss_res != decoded.shape[2]:
|
| 272 |
+
ratio = decoded.shape[2] // target_ss_res
|
| 273 |
+
decoded = torch.nn.functional.max_pool3d(decoded.float(), ratio, ratio, 0) > 0.5
|
| 274 |
+
coords = torch.argwhere(decoded)[:, [0, 2, 3, 4]].int()
|
| 275 |
+
|
| 276 |
+
if self.low_vram:
|
| 277 |
+
self._vggt_models_to('cpu')
|
| 278 |
+
torch.cuda.empty_cache()
|
| 279 |
+
|
| 280 |
+
return coords
|
| 281 |
+
|
| 282 |
+
@staticmethod
|
| 283 |
+
def _translate_ss_params(ss_sampler_params: dict) -> dict:
|
| 284 |
+
"""
|
| 285 |
+
Translate ReconViaGen SS sampler param keys to TRELLIS.2 format.
|
| 286 |
+
|
| 287 |
+
ReconViaGen uses ``cfg_strength`` / ``cfg_interval``;
|
| 288 |
+
TRELLIS.2 expects ``guidance_strength`` / ``guidance_interval``.
|
| 289 |
+
All other keys (steps, guidance_rescale, rescale_t) pass through unchanged.
|
| 290 |
+
"""
|
| 291 |
+
mapping = {'cfg_strength': 'guidance_strength', 'cfg_interval': 'guidance_interval'}
|
| 292 |
+
return {mapping.get(k, k): v for k, v in ss_sampler_params.items()}
|
| 293 |
+
|
| 294 |
+
@torch.no_grad()
|
| 295 |
+
def _run_ss_stage_trellis2(
|
| 296 |
+
self,
|
| 297 |
+
images: List[Image.Image],
|
| 298 |
+
target_ss_res: int,
|
| 299 |
+
strategy: Optional[str] = None,
|
| 300 |
+
sampler_params: dict = {},
|
| 301 |
+
) -> torch.Tensor:
|
| 302 |
+
"""
|
| 303 |
+
Use TRELLIS.2's sparse structure flow model to generate voxel coordinates.
|
| 304 |
+
|
| 305 |
+
- Single image or strategy=None : calls t2p.sample_sparse_structure() directly.
|
| 306 |
+
- Multi-image with strategy : uses t2p._multi_image_sample() for fusion.
|
| 307 |
+
|
| 308 |
+
Args:
|
| 309 |
+
sampler_params: Override params in TRELLIS.2 key format
|
| 310 |
+
(guidance_strength, guidance_rescale, rescale_t, steps).
|
| 311 |
+
Use _translate_ss_params() to convert from ReconViaGen format.
|
| 312 |
+
|
| 313 |
+
Returns:
|
| 314 |
+
coords : (N, 4) int tensor [batch_idx, x, y, z] in [0, target_ss_res)
|
| 315 |
+
"""
|
| 316 |
+
t2p = self.trellis2_pipeline
|
| 317 |
+
|
| 318 |
+
if len(images) == 1 or strategy is None:
|
| 319 |
+
cond_512 = t2p.get_cond(images, 512)
|
| 320 |
+
return t2p.sample_sparse_structure(
|
| 321 |
+
cond_512, target_ss_res, num_samples=1,
|
| 322 |
+
sampler_params=sampler_params,
|
| 323 |
+
)
|
| 324 |
+
|
| 325 |
+
# Multi-image: per-image conds + fusion strategy
|
| 326 |
+
conds_512 = [t2p.get_cond([img], 512) for img in images]
|
| 327 |
+
flow_model_ss = t2p.models['sparse_structure_flow_model']
|
| 328 |
+
reso = flow_model_ss.resolution
|
| 329 |
+
noise = torch.randn(1, flow_model_ss.in_channels, reso, reso, reso).to(t2p.device)
|
| 330 |
+
if t2p.low_vram:
|
| 331 |
+
flow_model_ss.to(t2p.device)
|
| 332 |
+
z_s = t2p._multi_image_sample(
|
| 333 |
+
t2p.sparse_structure_sampler, flow_model_ss, noise,
|
| 334 |
+
conds_512, t2p.sparse_structure_sampler_params, sampler_params,
|
| 335 |
+
strategy=strategy, verbose=True,
|
| 336 |
+
tqdm_desc="Sampling sparse structure (TRELLIS.2)",
|
| 337 |
+
)
|
| 338 |
+
if t2p.low_vram:
|
| 339 |
+
flow_model_ss.cpu()
|
| 340 |
+
|
| 341 |
+
decoder_ss = t2p.models['sparse_structure_decoder']
|
| 342 |
+
if t2p.low_vram:
|
| 343 |
+
decoder_ss.to(t2p.device)
|
| 344 |
+
decoded = decoder_ss(z_s) > 0
|
| 345 |
+
if t2p.low_vram:
|
| 346 |
+
decoder_ss.cpu()
|
| 347 |
+
|
| 348 |
+
if target_ss_res != decoded.shape[2]:
|
| 349 |
+
ratio = decoded.shape[2] // target_ss_res
|
| 350 |
+
decoded = torch.nn.functional.max_pool3d(decoded.float(), ratio, ratio, 0) > 0.5
|
| 351 |
+
return torch.argwhere(decoded)[:, [0, 2, 3, 4]].int()
|
| 352 |
+
|
| 353 |
+
# ── Main run (single image) ───────────────────────────────────────────────
|
| 354 |
+
|
| 355 |
+
@torch.no_grad()
|
| 356 |
+
def run(
|
| 357 |
+
self,
|
| 358 |
+
images: List[Image.Image],
|
| 359 |
+
seed: int = 42,
|
| 360 |
+
ss_sampler_params: dict = {},
|
| 361 |
+
slat_sampler_params: dict = {},
|
| 362 |
+
shape_slat_sampler_params: dict = {},
|
| 363 |
+
tex_slat_sampler_params: dict = {},
|
| 364 |
+
pipeline_type: str = '1024',
|
| 365 |
+
preprocess_image: bool = True,
|
| 366 |
+
return_latent: bool = False,
|
| 367 |
+
max_num_tokens: int = 49152,
|
| 368 |
+
ss_source: str = 'direct',
|
| 369 |
+
):
|
| 370 |
+
"""
|
| 371 |
+
Run the full hybrid pipeline.
|
| 372 |
+
|
| 373 |
+
Args:
|
| 374 |
+
images : List of (preprocessed) PIL images.
|
| 375 |
+
seed : Random seed.
|
| 376 |
+
ss_sampler_params : Override params for the SS sampler (ReconViaGen keys:
|
| 377 |
+
cfg_strength, cfg_interval, guidance_rescale, rescale_t, steps).
|
| 378 |
+
slat_sampler_params : Override params for the ReconViaGen SLat sampler
|
| 379 |
+
(only used when ss_source='mesh').
|
| 380 |
+
shape_slat_sampler_params: Override params for the TRELLIS.2 shape-SLat sampler.
|
| 381 |
+
tex_slat_sampler_params : Override params for the TRELLIS.2 tex-SLat sampler.
|
| 382 |
+
pipeline_type : '512' – shape/tex slat at 512 resolution.
|
| 383 |
+
'1024' – coords at res 64, direct shape_slat_flow_model_1024
|
| 384 |
+
(no cascade, fast).
|
| 385 |
+
'1024_cascade' – LR 512 → upsample → HR 1024 (higher detail).
|
| 386 |
+
'1536_cascade' – LR 512 → upsample → HR 1536 (highest detail).
|
| 387 |
+
preprocess_image : Whether to preprocess images first.
|
| 388 |
+
return_latent : If True, also return (shape_slat, tex_slat, res).
|
| 389 |
+
ss_source : 'direct' – ReconViaGen SS diffusion → coords (fast).
|
| 390 |
+
'mesh' – ReconViaGen full pipeline → mesh → voxelize → coords.
|
| 391 |
+
'mvtrellis2' – TRELLIS.2 SS flow model → coords.
|
| 392 |
+
|
| 393 |
+
Returns:
|
| 394 |
+
List[MeshWithVoxel] (one per sample, usually 1)
|
| 395 |
+
optionally: (shape_slat, tex_slat, res)
|
| 396 |
+
"""
|
| 397 |
+
assert pipeline_type in ('512', '1024', '1024_cascade', '1536_cascade'), \
|
| 398 |
+
f"pipeline_type must be '512', '1024', '1024_cascade', or '1536_cascade', got '{pipeline_type}'"
|
| 399 |
+
assert ss_source in ('direct', 'mesh', 'mvtrellis2'), \
|
| 400 |
+
f"ss_source must be 'direct', 'mesh', or 'mvtrellis2', got '{ss_source}'"
|
| 401 |
+
|
| 402 |
+
torch.manual_seed(seed)
|
| 403 |
+
|
| 404 |
+
if preprocess_image:
|
| 405 |
+
images = [self.preprocess_image(img) for img in images]
|
| 406 |
+
|
| 407 |
+
# Target SS resolution: 32 for 512/cascade pipelines, 64 for direct-1024
|
| 408 |
+
target_ss_res = {'512': 32, '1024': 64, '1024_cascade': 32, '1536_cascade': 32}[pipeline_type]
|
| 409 |
+
|
| 410 |
+
# ── Stage 1: SS ───────────────────────────────────────────────────────
|
| 411 |
+
if ss_source == 'direct':
|
| 412 |
+
# ReconViaGen SS diffusion only → coords directly (no mesh)
|
| 413 |
+
coords = self._run_ss_stage_direct(images, target_ss_res, ss_sampler_params)
|
| 414 |
+
elif ss_source == 'mvtrellis2':
|
| 415 |
+
# TRELLIS.2 SS flow model → coords (single image, no fusion strategy)
|
| 416 |
+
coords = self._run_ss_stage_trellis2(
|
| 417 |
+
images, target_ss_res, strategy=None,
|
| 418 |
+
sampler_params=self._translate_ss_params(ss_sampler_params),
|
| 419 |
+
)
|
| 420 |
+
else: # 'mesh'
|
| 421 |
+
# ReconViaGen full pipeline → mesh → decimate/fill/voxelize → coords
|
| 422 |
+
coords = self._run_ss_stage(images, target_ss_res, ss_sampler_params, slat_sampler_params)
|
| 423 |
+
|
| 424 |
+
# ── Stages 2 & 3: shape_slat + tex_slat (TRELLIS.2) ──────────────────
|
| 425 |
+
t2p = self.trellis2_pipeline
|
| 426 |
+
|
| 427 |
+
cond_512 = t2p.get_cond(images, 512) if pipeline_type in ('512', '1024_cascade', '1536_cascade') else None
|
| 428 |
+
cond_1024 = t2p.get_cond(images, 1024) if pipeline_type != '512' else None
|
| 429 |
+
|
| 430 |
+
if pipeline_type == '512':
|
| 431 |
+
shape_slat = t2p.sample_shape_slat(
|
| 432 |
+
cond_512,
|
| 433 |
+
t2p.models['shape_slat_flow_model_512'],
|
| 434 |
+
coords,
|
| 435 |
+
shape_slat_sampler_params,
|
| 436 |
+
)
|
| 437 |
+
tex_slat = t2p.sample_tex_slat(
|
| 438 |
+
cond_512,
|
| 439 |
+
t2p.models['tex_slat_flow_model_512'],
|
| 440 |
+
shape_slat,
|
| 441 |
+
tex_slat_sampler_params,
|
| 442 |
+
)
|
| 443 |
+
res = 512
|
| 444 |
+
elif pipeline_type == '1024':
|
| 445 |
+
# Direct 1024: coords at res 64 → shape_slat_flow_model_1024 (no cascade)
|
| 446 |
+
shape_slat = t2p.sample_shape_slat(
|
| 447 |
+
cond_1024,
|
| 448 |
+
t2p.models['shape_slat_flow_model_1024'],
|
| 449 |
+
coords,
|
| 450 |
+
shape_slat_sampler_params,
|
| 451 |
+
)
|
| 452 |
+
tex_slat = t2p.sample_tex_slat(
|
| 453 |
+
cond_1024,
|
| 454 |
+
t2p.models['tex_slat_flow_model_1024'],
|
| 455 |
+
shape_slat,
|
| 456 |
+
tex_slat_sampler_params,
|
| 457 |
+
)
|
| 458 |
+
res = 1024
|
| 459 |
+
elif pipeline_type == '1024_cascade':
|
| 460 |
+
shape_slat, res = t2p.sample_shape_slat_cascade(
|
| 461 |
+
cond_512, cond_1024,
|
| 462 |
+
t2p.models['shape_slat_flow_model_512'], t2p.models['shape_slat_flow_model_1024'],
|
| 463 |
+
512, 1024,
|
| 464 |
+
coords,
|
| 465 |
+
shape_slat_sampler_params,
|
| 466 |
+
max_num_tokens
|
| 467 |
+
)
|
| 468 |
+
tex_slat = t2p.sample_tex_slat(
|
| 469 |
+
cond_1024,
|
| 470 |
+
t2p.models['tex_slat_flow_model_1024'],
|
| 471 |
+
shape_slat,
|
| 472 |
+
tex_slat_sampler_params,
|
| 473 |
+
)
|
| 474 |
+
elif pipeline_type == '1536_cascade':
|
| 475 |
+
shape_slat, res = t2p.sample_shape_slat_cascade(
|
| 476 |
+
cond_512, cond_1024,
|
| 477 |
+
t2p.models['shape_slat_flow_model_512'], t2p.models['shape_slat_flow_model_1024'],
|
| 478 |
+
512, 1536,
|
| 479 |
+
coords, shape_slat_sampler_params,
|
| 480 |
+
max_num_tokens
|
| 481 |
+
)
|
| 482 |
+
tex_slat = t2p.sample_tex_slat(
|
| 483 |
+
cond_1024,
|
| 484 |
+
t2p.models['tex_slat_flow_model_1024'],
|
| 485 |
+
shape_slat,
|
| 486 |
+
tex_slat_sampler_params,
|
| 487 |
+
)
|
| 488 |
+
|
| 489 |
+
# ── Stage 4: Decode ───────────────────────────────────────────────────
|
| 490 |
+
torch.cuda.empty_cache()
|
| 491 |
+
out_mesh = t2p.decode_latent(shape_slat, tex_slat, res)
|
| 492 |
+
|
| 493 |
+
if return_latent:
|
| 494 |
+
return out_mesh, (shape_slat, tex_slat, res)
|
| 495 |
+
return out_mesh
|
| 496 |
+
|
| 497 |
+
# ── Multi-image run ───────────────────────────────────────────────────────
|
| 498 |
+
|
| 499 |
+
@torch.no_grad()
|
| 500 |
+
def run_multi_image(
|
| 501 |
+
self,
|
| 502 |
+
images: List[Image.Image],
|
| 503 |
+
strategy: str = 'average_right',
|
| 504 |
+
seed: int = 42,
|
| 505 |
+
ss_sampler_params: dict = {},
|
| 506 |
+
slat_sampler_params: dict = {},
|
| 507 |
+
shape_slat_sampler_params: dict = {},
|
| 508 |
+
tex_slat_sampler_params: dict = {},
|
| 509 |
+
pipeline_type: str = '1024',
|
| 510 |
+
preprocess_image: bool = True,
|
| 511 |
+
return_latent: bool = False,
|
| 512 |
+
max_num_tokens: int = 49152,
|
| 513 |
+
ss_source: str = 'direct',
|
| 514 |
+
):
|
| 515 |
+
"""
|
| 516 |
+
Multi-image variant.
|
| 517 |
+
|
| 518 |
+
The SS stage uses all images jointly (VGGT processes them together, or
|
| 519 |
+
TRELLIS.2 multi-image fusion when ss_source='mvtrellis2').
|
| 520 |
+
The shape/tex slat stages use TRELLIS.2's ``_multi_image_sample`` with
|
| 521 |
+
the chosen fusion ``strategy``.
|
| 522 |
+
|
| 523 |
+
Pipeline types:
|
| 524 |
+
- '512' : direct shape_slat_flow_model_512, res=512.
|
| 525 |
+
- '1024' : coords at res 64, direct shape_slat_flow_model_1024 (no cascade).
|
| 526 |
+
- '1024_cascade' : LR 512 → upsample → HR 1024 cascade.
|
| 527 |
+
- '1536_cascade' : LR 512 → upsample → HR 1536 cascade.
|
| 528 |
+
|
| 529 |
+
Args:
|
| 530 |
+
strategy: 'sequential' | 'average' | 'average_right' | 'weighted_average'
|
| 531 |
+
| 'adaptive_guidance_weight' | 'fixed_guidance_rescale'
|
| 532 |
+
(passed to Trellis2ImageTo3DPipeline._multi_image_sample)
|
| 533 |
+
"""
|
| 534 |
+
assert pipeline_type in ('512', '1024', '1024_cascade', '1536_cascade'), \
|
| 535 |
+
f"pipeline_type must be '512', '1024', '1024_cascade', or '1536_cascade', got '{pipeline_type}'"
|
| 536 |
+
|
| 537 |
+
# Single image → fall back to run()
|
| 538 |
+
if len(images) == 1:
|
| 539 |
+
return self.run(
|
| 540 |
+
images, seed=seed,
|
| 541 |
+
ss_sampler_params=ss_sampler_params,
|
| 542 |
+
slat_sampler_params=slat_sampler_params,
|
| 543 |
+
shape_slat_sampler_params=shape_slat_sampler_params,
|
| 544 |
+
tex_slat_sampler_params=tex_slat_sampler_params,
|
| 545 |
+
pipeline_type=pipeline_type,
|
| 546 |
+
preprocess_image=preprocess_image,
|
| 547 |
+
return_latent=return_latent,
|
| 548 |
+
max_num_tokens=max_num_tokens,
|
| 549 |
+
ss_source=ss_source,
|
| 550 |
+
)
|
| 551 |
+
|
| 552 |
+
torch.manual_seed(seed)
|
| 553 |
+
|
| 554 |
+
if preprocess_image:
|
| 555 |
+
images = [self.preprocess_image(img) for img in images]
|
| 556 |
+
|
| 557 |
+
target_ss_res = {'512': 32, '1024': 64, '1024_cascade': 32, '1536_cascade': 32}[pipeline_type]
|
| 558 |
+
|
| 559 |
+
# ── Stage 1: SS (all images fed jointly) ─────────────────────────────
|
| 560 |
+
if ss_source == 'direct':
|
| 561 |
+
coords = self._run_ss_stage_direct(images, target_ss_res, ss_sampler_params)
|
| 562 |
+
elif ss_source == 'mvtrellis2':
|
| 563 |
+
# TRELLIS.2 SS flow model with multi-image fusion strategy
|
| 564 |
+
coords = self._run_ss_stage_trellis2(
|
| 565 |
+
images, target_ss_res, strategy=strategy,
|
| 566 |
+
sampler_params=self._translate_ss_params(ss_sampler_params),
|
| 567 |
+
)
|
| 568 |
+
else: # 'mesh'
|
| 569 |
+
coords = self._run_ss_stage(images, target_ss_res, ss_sampler_params, slat_sampler_params)
|
| 570 |
+
|
| 571 |
+
# ── Per-image conditioning ────────────────────────────────────────────
|
| 572 |
+
t2p = self.trellis2_pipeline
|
| 573 |
+
conds_512 = [t2p.get_cond([img], 512) for img in images] if pipeline_type in ('512', '1024_cascade', '1536_cascade') else None
|
| 574 |
+
conds_1024 = [t2p.get_cond([img], 1024) for img in images] if pipeline_type != '512' else None
|
| 575 |
+
|
| 576 |
+
std_shape = torch.tensor(t2p.shape_slat_normalization['std'])[None]
|
| 577 |
+
mean_shape = torch.tensor(t2p.shape_slat_normalization['mean'])[None]
|
| 578 |
+
std_tex = torch.tensor(t2p.tex_slat_normalization['std'])[None]
|
| 579 |
+
mean_tex = torch.tensor(t2p.tex_slat_normalization['mean'])[None]
|
| 580 |
+
|
| 581 |
+
def _sample_shape(fm, conds, tqdm_desc):
|
| 582 |
+
noise_slat = SparseTensor(
|
| 583 |
+
feats=torch.randn(coords.shape[0], fm.in_channels).to(t2p.device),
|
| 584 |
+
coords=coords,
|
| 585 |
+
)
|
| 586 |
+
if t2p.low_vram:
|
| 587 |
+
fm.to(t2p.device)
|
| 588 |
+
slat = t2p._multi_image_sample(
|
| 589 |
+
t2p.shape_slat_sampler, fm, noise_slat,
|
| 590 |
+
conds, t2p.shape_slat_sampler_params, shape_slat_sampler_params,
|
| 591 |
+
strategy=strategy, verbose=True, tqdm_desc=tqdm_desc,
|
| 592 |
+
)
|
| 593 |
+
if t2p.low_vram:
|
| 594 |
+
fm.cpu()
|
| 595 |
+
return slat * std_shape.to(slat.device) + mean_shape.to(slat.device)
|
| 596 |
+
|
| 597 |
+
def _sample_tex(fm, conds, shape_slat):
|
| 598 |
+
s_norm = (shape_slat - mean_shape.to(shape_slat.device)) / std_shape.to(shape_slat.device)
|
| 599 |
+
in_ch = fm.in_channels if isinstance(fm, torch.nn.Module) else fm[0].in_channels
|
| 600 |
+
noise_tex = s_norm.replace(
|
| 601 |
+
feats=torch.randn(s_norm.coords.shape[0], in_ch - s_norm.feats.shape[1]).to(t2p.device)
|
| 602 |
+
)
|
| 603 |
+
if t2p.low_vram:
|
| 604 |
+
fm.to(t2p.device)
|
| 605 |
+
slat = t2p._multi_image_sample(
|
| 606 |
+
t2p.tex_slat_sampler, fm, noise_tex,
|
| 607 |
+
conds, t2p.tex_slat_sampler_params, tex_slat_sampler_params,
|
| 608 |
+
strategy=strategy, verbose=True, tqdm_desc="Sampling texture SLat",
|
| 609 |
+
concat_cond=s_norm,
|
| 610 |
+
)
|
| 611 |
+
if t2p.low_vram:
|
| 612 |
+
fm.cpu()
|
| 613 |
+
return slat * std_tex.to(slat.device) + mean_tex.to(slat.device)
|
| 614 |
+
|
| 615 |
+
if pipeline_type == '512':
|
| 616 |
+
shape_slat = _sample_shape(t2p.models['shape_slat_flow_model_512'], conds_512, "Sampling shape SLat")
|
| 617 |
+
tex_slat = _sample_tex(t2p.models['tex_slat_flow_model_512'], conds_512, shape_slat)
|
| 618 |
+
res = 512
|
| 619 |
+
|
| 620 |
+
elif pipeline_type == '1024':
|
| 621 |
+
# Direct 1024: coords at res 64 → shape_slat_flow_model_1024 (no cascade)
|
| 622 |
+
shape_slat = _sample_shape(t2p.models['shape_slat_flow_model_1024'], conds_1024, "Sampling shape SLat (1024 direct)")
|
| 623 |
+
tex_slat = _sample_tex(t2p.models['tex_slat_flow_model_1024'], conds_1024, shape_slat)
|
| 624 |
+
res = 1024
|
| 625 |
+
|
| 626 |
+
else: # '1024_cascade' or '1536_cascade' — cascade: LR (512) → upsample → HR (1024)
|
| 627 |
+
target_res = {'1024_cascade': 1024, '1536_cascade': 1536}[pipeline_type]
|
| 628 |
+
|
| 629 |
+
# LR stage: multi-image sample at 512 resolution
|
| 630 |
+
fm_lr = t2p.models['shape_slat_flow_model_512']
|
| 631 |
+
noise_lr = SparseTensor(
|
| 632 |
+
feats=torch.randn(coords.shape[0], fm_lr.in_channels).to(t2p.device),
|
| 633 |
+
coords=coords,
|
| 634 |
+
)
|
| 635 |
+
if t2p.low_vram:
|
| 636 |
+
fm_lr.to(t2p.device)
|
| 637 |
+
slat_lr = t2p._multi_image_sample(
|
| 638 |
+
t2p.shape_slat_sampler, fm_lr, noise_lr,
|
| 639 |
+
conds_512, t2p.shape_slat_sampler_params, shape_slat_sampler_params,
|
| 640 |
+
strategy=strategy, verbose=True, tqdm_desc="Sampling shape SLat (LR)",
|
| 641 |
+
)
|
| 642 |
+
if t2p.low_vram:
|
| 643 |
+
fm_lr.cpu()
|
| 644 |
+
slat_lr = slat_lr * std_shape.to(slat_lr.device) + mean_shape.to(slat_lr.device)
|
| 645 |
+
|
| 646 |
+
# Upsample LR slat to HR coords
|
| 647 |
+
if t2p.low_vram:
|
| 648 |
+
t2p.models['shape_slat_decoder'].to(t2p.device)
|
| 649 |
+
t2p.models['shape_slat_decoder'].low_vram = True
|
| 650 |
+
hr_coords_raw = t2p.models['shape_slat_decoder'].upsample(slat_lr, upsample_times=4)
|
| 651 |
+
if t2p.low_vram:
|
| 652 |
+
t2p.models['shape_slat_decoder'].cpu()
|
| 653 |
+
t2p.models['shape_slat_decoder'].low_vram = False
|
| 654 |
+
|
| 655 |
+
hr_resolution = target_res
|
| 656 |
+
while True:
|
| 657 |
+
quant_coords = torch.cat([
|
| 658 |
+
hr_coords_raw[:, :1],
|
| 659 |
+
((hr_coords_raw[:, 1:] + 0.5) / 512 * (hr_resolution // 16)).int(),
|
| 660 |
+
], dim=1)
|
| 661 |
+
hr_c = quant_coords.unique(dim=0)
|
| 662 |
+
if hr_c.shape[0] < max_num_tokens or hr_resolution == 1024:
|
| 663 |
+
if hr_resolution != target_res:
|
| 664 |
+
print(f"Due to the limited number of tokens, the resolution is reduced to {hr_resolution}.")
|
| 665 |
+
break
|
| 666 |
+
hr_resolution -= 128
|
| 667 |
+
|
| 668 |
+
# HR stage: multi-image sample at hr_resolution using conds_1024
|
| 669 |
+
fm_hr = t2p.models['shape_slat_flow_model_1024']
|
| 670 |
+
noise_hr = SparseTensor(
|
| 671 |
+
feats=torch.randn(hr_c.shape[0], fm_hr.in_channels).to(t2p.device),
|
| 672 |
+
coords=hr_c,
|
| 673 |
+
)
|
| 674 |
+
if t2p.low_vram:
|
| 675 |
+
fm_hr.to(t2p.device)
|
| 676 |
+
shape_slat = t2p._multi_image_sample(
|
| 677 |
+
t2p.shape_slat_sampler, fm_hr, noise_hr,
|
| 678 |
+
conds_1024, t2p.shape_slat_sampler_params, shape_slat_sampler_params,
|
| 679 |
+
strategy=strategy, verbose=True, tqdm_desc="Sampling shape SLat (HR)",
|
| 680 |
+
)
|
| 681 |
+
if t2p.low_vram:
|
| 682 |
+
fm_hr.cpu()
|
| 683 |
+
shape_slat = shape_slat * std_shape.to(shape_slat.device) + mean_shape.to(shape_slat.device)
|
| 684 |
+
res = hr_resolution
|
| 685 |
+
|
| 686 |
+
tex_slat = _sample_tex(t2p.models['tex_slat_flow_model_1024'], conds_1024, shape_slat)
|
| 687 |
+
|
| 688 |
+
# ── Stage 4: Decode ───────────────────────────────────────────────────
|
| 689 |
+
torch.cuda.empty_cache()
|
| 690 |
+
out_mesh = t2p.decode_latent(shape_slat, tex_slat, res)
|
| 691 |
+
|
| 692 |
+
if return_latent:
|
| 693 |
+
return out_mesh, (shape_slat, tex_slat, res)
|
| 694 |
+
return out_mesh
|
trellis/pipelines/trellis_image_to_3d.py
CHANGED
|
@@ -20,8 +20,7 @@ from scipy.spatial.transform import Rotation
|
|
| 20 |
from transformers import AutoModelForImageSegmentation
|
| 21 |
import rembg
|
| 22 |
# for app_refine.py, please uncomment these lines
|
| 23 |
-
|
| 24 |
-
# from tqdm import tqdm
|
| 25 |
|
| 26 |
def export_point_cloud(xyz, color):
|
| 27 |
# Convert tensors to numpy arrays if needed
|
|
@@ -256,7 +255,7 @@ class TrellisImageTo3DPipeline(Pipeline):
|
|
| 256 |
bbox = np.min(bbox[:, 1]), np.min(bbox[:, 0]), np.max(bbox[:, 1]), np.max(bbox[:, 0])
|
| 257 |
center = [(bbox[0] + bbox[2]) / 2, (bbox[1] + bbox[3]) / 2]
|
| 258 |
size = max(bbox[2] - bbox[0], bbox[3] - bbox[1])
|
| 259 |
-
size = int(size * 1
|
| 260 |
height, width = alpha.shape
|
| 261 |
if not recenter:
|
| 262 |
center = [width / 2, height / 2]
|
|
@@ -340,8 +339,8 @@ class TrellisImageTo3DPipeline(Pipeline):
|
|
| 340 |
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
|
| 341 |
])
|
| 342 |
|
| 343 |
-
input_images = transform_image(image).unsqueeze(0).to(self.device)
|
| 344 |
-
|
| 345 |
with torch.no_grad():
|
| 346 |
preds = self.birefnet_model(input_images)[-1].sigmoid().cpu()
|
| 347 |
|
|
@@ -616,68 +615,6 @@ class TrellisImageTo3DPipeline(Pipeline):
|
|
| 616 |
slat = slat * std + mean
|
| 617 |
return slat
|
| 618 |
|
| 619 |
-
def sample_slat_opt(
|
| 620 |
-
self,
|
| 621 |
-
apperance_learning_rate,
|
| 622 |
-
start_t,
|
| 623 |
-
input_images: torch.Tensor,
|
| 624 |
-
extrinsics: torch.Tensor,
|
| 625 |
-
intrinsics: torch.Tensor,
|
| 626 |
-
cond: dict,
|
| 627 |
-
coords: torch.Tensor,
|
| 628 |
-
sampler_params: dict = {},
|
| 629 |
-
) -> sp.SparseTensor:
|
| 630 |
-
"""
|
| 631 |
-
Sample structured latent with the given conditioning.
|
| 632 |
-
|
| 633 |
-
Args:
|
| 634 |
-
cond (dict): The conditioning information.
|
| 635 |
-
coords (torch.Tensor): The coordinates of the sparse structure.
|
| 636 |
-
sampler_params (dict): Additional parameters for the sampler.
|
| 637 |
-
"""
|
| 638 |
-
# Sample structured latent
|
| 639 |
-
flow_model = self.models['slat_flow_model']
|
| 640 |
-
slat_decoder_gs = self.models['slat_decoder_gs']
|
| 641 |
-
slat_decoder_mesh = self.models['slat_decoder_mesh']
|
| 642 |
-
noise = sp.SparseTensor(
|
| 643 |
-
feats=torch.randn(coords.shape[0], flow_model.in_channels).to(self.device),
|
| 644 |
-
coords=coords,
|
| 645 |
-
)
|
| 646 |
-
std = torch.tensor(self.slat_normalization['std'])[None].to(self.device)
|
| 647 |
-
mean = torch.tensor(self.slat_normalization['mean'])[None].to(self.device)
|
| 648 |
-
sampler_params = {**self.slat_sampler_params, **sampler_params}
|
| 649 |
-
slat = self.slat_sampler.sample_slat_opt_delta_v(
|
| 650 |
-
flow_model,
|
| 651 |
-
slat_decoder_gs,
|
| 652 |
-
slat_decoder_mesh,
|
| 653 |
-
std,
|
| 654 |
-
mean,
|
| 655 |
-
self.dreamsim_model,
|
| 656 |
-
apperance_learning_rate,
|
| 657 |
-
start_t,
|
| 658 |
-
input_images,
|
| 659 |
-
extrinsics,
|
| 660 |
-
intrinsics,
|
| 661 |
-
noise,
|
| 662 |
-
**cond,
|
| 663 |
-
**sampler_params,
|
| 664 |
-
verbose=True
|
| 665 |
-
).samples
|
| 666 |
-
|
| 667 |
-
slat = slat * std + mean
|
| 668 |
-
# from trellis.utils import render_utils, postprocessing_utils
|
| 669 |
-
# import imageio
|
| 670 |
-
# std = torch.tensor(self.slat_normalization['std'])[None].to(noise.device)
|
| 671 |
-
# mean = torch.tensor(self.slat_normalization['mean'])[None].to(noise.device)
|
| 672 |
-
# for i in range(sampler_params['steps']):
|
| 673 |
-
# latent = slat.pred_x_0[i] * std + mean
|
| 674 |
-
# outputs = self.decode_slat(latent, ["mesh", "gaussian"])
|
| 675 |
-
# video_geo = render_utils.render_video(outputs['mesh'][0], resolution=512, pitch=0, inverse_direction=True, num_frames=120)['normal']
|
| 676 |
-
# video_color = render_utils.render_video(outputs['gaussian'][0], resolution=512, pitch=0, inverse_direction=True, num_frames=120)['color']
|
| 677 |
-
# video = [np.concatenate([video_color[i], video_geo[i]], axis=1) for i in range(len(video_color))]
|
| 678 |
-
# imageio.mimsave('outputs/slat_iter_{i:02d}.mp4'.format(i=i), video, fps=15)
|
| 679 |
-
return slat
|
| 680 |
-
|
| 681 |
def get_input(self, batch_data):
|
| 682 |
std = torch.tensor(self.slat_normalization['std'])[None].to(self.device)
|
| 683 |
mean = torch.tensor(self.slat_normalization['mean'])[None].to(self.device)
|
|
@@ -936,13 +873,16 @@ class TrellisVGGTTo3DPipeline(TrellisImageTo3DPipeline):
|
|
| 936 |
):
|
| 937 |
|
| 938 |
torch.manual_seed(seed)
|
| 939 |
-
|
|
|
|
|
|
|
| 940 |
b, n, _, _ = aggregated_tokens_list[0].shape
|
| 941 |
image_cond = self.encode_image(image).reshape(b, n, -1, 1024)
|
| 942 |
|
| 943 |
# if coords is None:
|
| 944 |
ss_flow_model = self.models['sparse_structure_flow_model']
|
| 945 |
-
|
|
|
|
| 946 |
# Sample structured latent
|
| 947 |
ss_sampler_params = {**self.sparse_structure_sampler_params, **sparse_structure_sampler_params}
|
| 948 |
reso = ss_flow_model.resolution
|
|
@@ -965,69 +905,11 @@ class TrellisVGGTTo3DPipeline(TrellisImageTo3DPipeline):
|
|
| 965 |
# slat_steps = {**self.slat_sampler_params, **slat_sampler_params}.get('steps')
|
| 966 |
# with self.inject_sampler_multi_image('slat_sampler', len(image), slat_steps, mode=mode):
|
| 967 |
# slat = self.sample_slat(cond, coords, slat_sampler_params)
|
| 968 |
-
|
| 969 |
-
|
| 970 |
slat = self.sample_slat(slat_cond, coords, slat_sampler_params)
|
| 971 |
return self.decode_slat(slat, formats), coords, ss_noise
|
| 972 |
|
| 973 |
-
def run_refine(
|
| 974 |
-
self,
|
| 975 |
-
image: Union[torch.Tensor, list[Image.Image]],
|
| 976 |
-
ss_learning_rate: float,
|
| 977 |
-
ss_start_t: float,
|
| 978 |
-
apperance_learning_rate: float,
|
| 979 |
-
apperance_start_t: float,
|
| 980 |
-
extrinsics: torch.Tensor,
|
| 981 |
-
intrinsics: torch.Tensor,
|
| 982 |
-
ss_noise: torch.Tensor,
|
| 983 |
-
input_points: torch.Tensor,
|
| 984 |
-
ss_refine_type: str = 'No',
|
| 985 |
-
coords: torch.Tensor = None,
|
| 986 |
-
num_samples: int = 1,
|
| 987 |
-
seed: int = 42,
|
| 988 |
-
sparse_structure_sampler_params: dict = {},
|
| 989 |
-
slat_sampler_params: dict = {},
|
| 990 |
-
formats: List[str] = ['mesh'],
|
| 991 |
-
mode: Literal['stochastic', 'multidiffusion'] = 'stochastic',
|
| 992 |
-
):
|
| 993 |
-
|
| 994 |
-
torch.manual_seed(seed)
|
| 995 |
-
aggregated_tokens_list, input_images = self.vggt_feat(image)
|
| 996 |
-
b, n, _, _ = aggregated_tokens_list[0].shape
|
| 997 |
-
image_cond = self.encode_image(image).reshape(b, n, -1, 1024)
|
| 998 |
-
|
| 999 |
-
if coords is None:
|
| 1000 |
-
ss_cond = self.get_ss_cond(image_cond[:, :, 5:], aggregated_tokens_list, num_samples)
|
| 1001 |
-
ss = torch.zeros(64, 64, 64, dtype=torch.long, device=image_cond.device)
|
| 1002 |
-
ss = ss.index_put_((input_points[:,0], input_points[:,1], input_points[:,2]), torch.tensor(1, dtype=ss.dtype, device=ss.device))
|
| 1003 |
-
ss = ss[None, None]
|
| 1004 |
-
torch.cuda.empty_cache()
|
| 1005 |
-
# Sample structured latent
|
| 1006 |
-
if ss_refine_type == 'noise':
|
| 1007 |
-
coords = self.sample_sparse_structure_opt_noise(ss_cond, ss, ss_learning_rate, num_samples, sparse_structure_sampler_params, ss_noise)
|
| 1008 |
-
elif ss_refine_type == 'deltav':
|
| 1009 |
-
coords = self.sample_sparse_structure_opt(ss_cond, ss, ss_learning_rate, ss_start_t, num_samples, sparse_structure_sampler_params, ss_noise)
|
| 1010 |
-
torch.cuda.empty_cache()
|
| 1011 |
-
|
| 1012 |
-
# pcd = o3d.geometry.PointCloud()
|
| 1013 |
-
# pcd.points = o3d.utility.Vector3dVector(coords[:,1:].cpu().numpy() / 64 - 0.5)
|
| 1014 |
-
# o3d.io.write_point_cloud('outputs/after_coords.ply', pcd)
|
| 1015 |
-
|
| 1016 |
-
# cond = {
|
| 1017 |
-
# 'cond': image_cond.reshape(n, -1, 1024),
|
| 1018 |
-
# 'neg_cond': torch.zeros_like(image_cond.reshape(n, -1, 1024))[:1],
|
| 1019 |
-
# }
|
| 1020 |
-
|
| 1021 |
-
# slat_steps = {**self.slat_sampler_params, **slat_sampler_params}.get('steps')
|
| 1022 |
-
|
| 1023 |
-
# with self.inject_sampler_multi_image('slat_sampler', len(image), slat_steps, mode=mode):
|
| 1024 |
-
# # slat = self.sample_slat(cond, coords, slat_sampler_params)
|
| 1025 |
-
# slat = self.sample_slat_opt(apperance_learning_rate, apperance_start_t, input_images, extrinsics, intrinsics, cond, coords, slat_sampler_params)
|
| 1026 |
-
|
| 1027 |
-
slat_cond = self.get_slat_cond(image_cond, aggregated_tokens_list, num_samples)
|
| 1028 |
-
slat = self.sample_slat_opt(apperance_learning_rate, apperance_start_t, input_images, extrinsics, intrinsics, slat_cond, coords, slat_sampler_params)
|
| 1029 |
-
return self.decode_slat(slat, formats)
|
| 1030 |
-
|
| 1031 |
@staticmethod
|
| 1032 |
def from_pretrained(path: str) -> "TrellisVGGTTo3DPipeline":
|
| 1033 |
"""
|
|
@@ -1040,7 +922,7 @@ class TrellisVGGTTo3DPipeline(TrellisImageTo3DPipeline):
|
|
| 1040 |
new_pipeline = TrellisVGGTTo3DPipeline()
|
| 1041 |
new_pipeline.__dict__ = pipeline.__dict__
|
| 1042 |
args = pipeline._pretrained_args
|
| 1043 |
-
new_pipeline.VGGT_dtype = torch.
|
| 1044 |
VGGT_model = VGGT.from_pretrained("Stable-X/vggt-object-v0-1")
|
| 1045 |
new_pipeline.VGGT_model = VGGT_model.to(new_pipeline.device)
|
| 1046 |
del new_pipeline.VGGT_model.depth_head
|
|
@@ -1065,9 +947,4 @@ class TrellisVGGTTo3DPipeline(TrellisImageTo3DPipeline):
|
|
| 1065 |
|
| 1066 |
new_pipeline._init_image_cond_model(args['image_cond_model'])
|
| 1067 |
|
| 1068 |
-
# for app_refine.py, please uncomment these lines
|
| 1069 |
-
# model, _ = dreamsim(pretrained=True, device=new_pipeline.device, dreamsim_type="dino_vitb16", cache_dir="weights/dreamsim")
|
| 1070 |
-
# new_pipeline.dreamsim_model = model
|
| 1071 |
-
# new_pipeline.dreamsim_model.eval()
|
| 1072 |
-
|
| 1073 |
return new_pipeline
|
|
|
|
| 20 |
from transformers import AutoModelForImageSegmentation
|
| 21 |
import rembg
|
| 22 |
# for app_refine.py, please uncomment these lines
|
| 23 |
+
from tqdm import tqdm
|
|
|
|
| 24 |
|
| 25 |
def export_point_cloud(xyz, color):
|
| 26 |
# Convert tensors to numpy arrays if needed
|
|
|
|
| 255 |
bbox = np.min(bbox[:, 1]), np.min(bbox[:, 0]), np.max(bbox[:, 1]), np.max(bbox[:, 0])
|
| 256 |
center = [(bbox[0] + bbox[2]) / 2, (bbox[1] + bbox[3]) / 2]
|
| 257 |
size = max(bbox[2] - bbox[0], bbox[3] - bbox[1])
|
| 258 |
+
size = int(size * 1)
|
| 259 |
height, width = alpha.shape
|
| 260 |
if not recenter:
|
| 261 |
center = [width / 2, height / 2]
|
|
|
|
| 339 |
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
|
| 340 |
])
|
| 341 |
|
| 342 |
+
input_images = transform_image(image).unsqueeze(0).to(self.device, dtype=torch.float16)
|
| 343 |
+
|
| 344 |
with torch.no_grad():
|
| 345 |
preds = self.birefnet_model(input_images)[-1].sigmoid().cpu()
|
| 346 |
|
|
|
|
| 615 |
slat = slat * std + mean
|
| 616 |
return slat
|
| 617 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 618 |
def get_input(self, batch_data):
|
| 619 |
std = torch.tensor(self.slat_normalization['std'])[None].to(self.device)
|
| 620 |
mean = torch.tensor(self.slat_normalization['mean'])[None].to(self.device)
|
|
|
|
| 873 |
):
|
| 874 |
|
| 875 |
torch.manual_seed(seed)
|
| 876 |
+
with torch.no_grad():
|
| 877 |
+
with torch.cuda.amp.autocast(dtype=self.VGGT_dtype):
|
| 878 |
+
aggregated_tokens_list, _ = self.vggt_feat(image)
|
| 879 |
b, n, _, _ = aggregated_tokens_list[0].shape
|
| 880 |
image_cond = self.encode_image(image).reshape(b, n, -1, 1024)
|
| 881 |
|
| 882 |
# if coords is None:
|
| 883 |
ss_flow_model = self.models['sparse_structure_flow_model']
|
| 884 |
+
with torch.no_grad():
|
| 885 |
+
ss_cond = self.get_ss_cond(image_cond[:, :, 5:], aggregated_tokens_list, num_samples)
|
| 886 |
# Sample structured latent
|
| 887 |
ss_sampler_params = {**self.sparse_structure_sampler_params, **sparse_structure_sampler_params}
|
| 888 |
reso = ss_flow_model.resolution
|
|
|
|
| 905 |
# slat_steps = {**self.slat_sampler_params, **slat_sampler_params}.get('steps')
|
| 906 |
# with self.inject_sampler_multi_image('slat_sampler', len(image), slat_steps, mode=mode):
|
| 907 |
# slat = self.sample_slat(cond, coords, slat_sampler_params)
|
| 908 |
+
with torch.no_grad():
|
| 909 |
+
slat_cond = self.get_slat_cond(image_cond, aggregated_tokens_list, num_samples)
|
| 910 |
slat = self.sample_slat(slat_cond, coords, slat_sampler_params)
|
| 911 |
return self.decode_slat(slat, formats), coords, ss_noise
|
| 912 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 913 |
@staticmethod
|
| 914 |
def from_pretrained(path: str) -> "TrellisVGGTTo3DPipeline":
|
| 915 |
"""
|
|
|
|
| 922 |
new_pipeline = TrellisVGGTTo3DPipeline()
|
| 923 |
new_pipeline.__dict__ = pipeline.__dict__
|
| 924 |
args = pipeline._pretrained_args
|
| 925 |
+
new_pipeline.VGGT_dtype = torch.bfloat16 if torch.cuda.get_device_capability()[0] >= 8 else torch.float16
|
| 926 |
VGGT_model = VGGT.from_pretrained("Stable-X/vggt-object-v0-1")
|
| 927 |
new_pipeline.VGGT_model = VGGT_model.to(new_pipeline.device)
|
| 928 |
del new_pipeline.VGGT_model.depth_head
|
|
|
|
| 947 |
|
| 948 |
new_pipeline._init_image_cond_model(args['image_cond_model'])
|
| 949 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 950 |
return new_pipeline
|