diff --git a/README.md b/README.md index ac10d32..55038b4 100644 --- a/README.md +++ b/README.md @@ -14,7 +14,7 @@ | Date | Description | | --- | --- | -| **2026-04-20** | Added node "Fill Holes Nicely with Meshlib" | +| **2026-04-20** | Added node "Fill Holes Nicely with Meshlib"
Fixed DinoV3 Features Extractor | | **2026-04-06** | Added new "DINO-Lock" functionality
New fill_holes in Sparse Generator
Thanks to Easymode on Discord | | **2026-04-05** | Added node "Extract Images from Video"
Can be used with "Sparse Generator with ReconViaGen" | | **2026-04-04** | Added node "Sparse Generator with ReconViaGen" | diff --git a/trellis2/modules/image_feature_extractor.py b/trellis2/modules/image_feature_extractor.py index 6f6a79c..b676330 100644 --- a/trellis2/modules/image_feature_extractor.py +++ b/trellis2/modules/image_feature_extractor.py @@ -83,11 +83,21 @@ class DinoV3FeatureExtractor: hidden_states = self.model.embeddings(image, bool_masked_pos=None) position_embeddings = self.model.rope_embeddings(image) - for i, layer_module in enumerate(self.model.layer): - hidden_states = layer_module( - hidden_states, - position_embeddings=position_embeddings, - ) + # transformers < 5 + if hasattr(self.model,'layer'): + for i, layer_module in enumerate(self.model.layer): + hidden_states = layer_module( + hidden_states, + position_embeddings=position_embeddings, + ) + elif hasattr(self.model,'model') and hasattr(self.model.model,'layer'): # transformers >= 5 + for i, layer_module in enumerate(self.model.model.layer): + hidden_states = layer_module( + hidden_states, + position_embeddings=position_embeddings, + ) + else: + raise Exception("Cannot extract features") return F.layer_norm(hidden_states, hidden_states.shape[-1:]) diff --git a/trellis2/trainers/flow_matching/mixins/image_conditioned.py b/trellis2/trainers/flow_matching/mixins/image_conditioned.py index ab8da40..f3889ae 100644 --- a/trellis2/trainers/flow_matching/mixins/image_conditioned.py +++ b/trellis2/trainers/flow_matching/mixins/image_conditioned.py @@ -85,11 +85,21 @@ class DinoV3FeatureExtractor: hidden_states = self.model.embeddings(image, bool_masked_pos=None) position_embeddings = self.model.rope_embeddings(image) - for i, layer_module in enumerate(self.model.layer): - hidden_states = layer_module( - hidden_states, - position_embeddings=position_embeddings, - ) + # transformers < 5 + if hasattr(self.model,'layer'): + for i, layer_module in enumerate(self.model.layer): + hidden_states = layer_module( + hidden_states, + position_embeddings=position_embeddings, + ) + elif hasattr(self.model,'model') and hasattr(self.model.model,'layer'): # transformers >= 5 + for i, layer_module in enumerate(self.model.model.layer): + hidden_states = layer_module( + hidden_states, + position_embeddings=position_embeddings, + ) + else: + raise Exception("Cannot extract features") return F.layer_norm(hidden_states, hidden_states.shape[-1:])