Stable-X commited on
Commit
17dd905
·
verified ·
1 Parent(s): cf5ebdc

Upload 74 files

Browse files
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
- 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
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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, 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)
 
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
- # from dreamsim import dreamsim
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.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
- aggregated_tokens_list, _ = self.vggt_feat(image)
 
 
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
- ss_cond = self.get_ss_cond(image_cond[:, :, 5:], aggregated_tokens_list, num_samples)
 
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
- slat_cond = self.get_slat_cond(image_cond, aggregated_tokens_list, num_samples)
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.float32
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