Loading facebook model using relative path + added low_vram and device in "load model" node

This commit is contained in:
Bruno Fargnoli
2025-12-19 13:07:28 +01:00
parent 89c4df5580
commit 5ee5fcd16c
2 changed files with 9 additions and 3 deletions
+5 -2
View File
@@ -78,6 +78,8 @@ class Trellis2LoadModel:
"required": {
"modelname": (["TRELLIS.2-4B"],),
"backend": (["flash_attn","xformers"],{"default":"xformers"}),
"device": (["cpu","cuda"],{"default":"cuda"}),
"low_vram": ("BOOLEAN",{"default":True}),
},
}
@@ -87,7 +89,7 @@ class Trellis2LoadModel:
CATEGORY = "Trellis2Wrapper"
OUTPUT_NODE = True
def process(self, modelname, backend):
def process(self, modelname, backend, device, low_vram):
os.environ['OPENCV_IO_ENABLE_OPENEXR'] = '1'
os.environ["PYTORCH_CUDA_ALLOC_CONF"] = "expandable_segments:True" # Can save GPU memory
os.environ['ATTN_BACKEND'] = backend
@@ -111,7 +113,8 @@ class Trellis2LoadModel:
raise Exception("Facebook Dinov3 model not found in models/facebook/dinov3-vitl16-pretrain-lvd1689m folder")
pipeline = Trellis2ImageTo3DPipeline.from_pretrained(model_path)
pipeline.cuda()
pipeline.low_vram = low_vram
pipeline.to(device)
return (pipeline,)
+4 -1
View File
@@ -49,6 +49,8 @@ class Trellis2ImageTo3DPipeline(Pipeline):
):
if models is None:
return
#models['sparse_structure_decoder'] = os.path.join("microsoft","TRELLIS-image-large","ckpts","ss_dec_conv3d_16l8_fp16.safetensors")
super().__init__(models)
self.sparse_structure_sampler = sparse_structure_sampler
self.shape_slat_sampler = shape_slat_sampler
@@ -95,7 +97,8 @@ class Trellis2ImageTo3DPipeline(Pipeline):
new_pipeline.shape_slat_normalization = args['shape_slat_normalization']
new_pipeline.tex_slat_normalization = args['tex_slat_normalization']
facebook_model_path = os.path.join(folder_paths.models_dir,"facebook","dinov3-vitl16-pretrain-lvd1689m")
#facebook_model_path = os.path.join(folder_paths.models_dir,"facebook","dinov3-vitl16-pretrain-lvd1689m")
facebook_model_path = os.path.join("models","facebook","dinov3-vitl16-pretrain-lvd1689m")
args['image_cond_model']['args']['model_name'] = facebook_model_path
new_pipeline.image_cond_model = getattr(image_feature_extractor, args['image_cond_model']['name'])(**args['image_cond_model']['args'])