Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
26 changes: 26 additions & 0 deletions generate.py
Original file line number Diff line number Diff line change
Expand Up @@ -81,6 +81,32 @@ def main():
pipeline = Trellis2ImageTo3DPipeline.from_pretrained("microsoft/TRELLIS.2-4B")
print(f"Loaded in {time.time() - t0:.0f}s")

# --- DINOv2 fallback (use when DINOv3 gated access is pending) ---
# DINOv2 vitl14 @ 448px → 1025 tokens = same as DINOv3 vitl16 @ 512px
import os
if os.environ.get("USE_DINOV2_FALLBACK", "0") == "1":
print("INFO: Using DINOv2 vitl14 fallback (DINOv3 access pending)")
import sys as _sys
_sys.path.insert(0, str(__file__).replace("generate.py", "TRELLIS.2"))
from trellis2.modules.image_feature_extractor import DinoV2FeatureExtractor
_orig_call = DinoV2FeatureExtractor.__call__
def _dinov2_call_448(self, image, **kwargs):
from PIL import Image as _PILImage
import numpy as _np, torch as _torch
if isinstance(image, list):
image = [i.resize((448, 448), _PILImage.LANCZOS) for i in image]
image = [_np.array(i.convert("RGB")).astype(_np.float32) / 255 for i in image]
image = [_torch.from_numpy(i).permute(2, 0, 1).float() for i in image]
image = _torch.stack(image).to(self.device)
import torch.nn.functional as _F
image = self.transform(image).to(self.device)
features = self.model(image, is_training=True)["x_prenorm"]
return _F.layer_norm(features, features.shape[-1:])
DinoV2FeatureExtractor.__call__ = _dinov2_call_448
pipeline.image_cond_model = DinoV2FeatureExtractor("dinov2_vitl14")
print("INFO: DINOv2 vitl14 loaded (448px → 1025 tokens)")
# --- end fallback ---

# Move to MPS
pipeline.to(torch.device("mps"))
print("Device: MPS")
Expand Down