Add new 1.1 models
This commit is contained in:
+16
-2
@@ -19,10 +19,14 @@ class MarigoldModelLoader:
|
||||
[
|
||||
'prs-eth/marigold-v1-0',
|
||||
'prs-eth/marigold-depth-lcm-v1-0',
|
||||
'prs-eth/marigold-depth-v1-1',
|
||||
'prs-eth/marigold-normals-v0-1',
|
||||
'prs-eth/marigold-normals-lcm-v0-1',
|
||||
'prs-eth/marigold-normals-v1-1',
|
||||
'GonzaloMG/marigold-e2e-ft-depth',
|
||||
'GonzaloMG/marigold-e2e-ft-normals'
|
||||
'GonzaloMG/marigold-e2e-ft-normals',
|
||||
'prs-eth/marigold-iid-lighting-v1-1',
|
||||
'prs-eth/marigold-iid-appearance-v1-1'
|
||||
],
|
||||
{
|
||||
"default": 'marigold-lcm-v1-0'
|
||||
@@ -46,7 +50,7 @@ ComfyUI/models/diffusers -folder
|
||||
def load(self, model):
|
||||
try:
|
||||
|
||||
from diffusers import MarigoldDepthPipeline, MarigoldNormalsPipeline
|
||||
from diffusers import MarigoldDepthPipeline, MarigoldNormalsPipeline, MarigoldIntrinsicsPipeline
|
||||
except:
|
||||
raise Exception("diffusers>=0.28 is required for v2 nodes")
|
||||
|
||||
@@ -75,6 +79,12 @@ ComfyUI/models/diffusers -folder
|
||||
checkpoint_path,
|
||||
variant=variant,
|
||||
torch_dtype=torch.float16).to(device)
|
||||
elif "iid" in model:
|
||||
modeltype = "intrinsics"
|
||||
self.marigold_pipeline = MarigoldIntrinsicsPipeline.from_pretrained(
|
||||
checkpoint_path,
|
||||
variant=variant,
|
||||
torch_dtype=torch.float16).to(device)
|
||||
else:
|
||||
modeltype = "depth"
|
||||
self.marigold_pipeline = MarigoldDepthPipeline.from_pretrained(
|
||||
@@ -171,6 +181,7 @@ Uses Diffusers 0.28.0 Marigold pipelines.
|
||||
processing_resolution = processing_resolution,
|
||||
**pipe_kwargs
|
||||
)
|
||||
#print("processed", processed[0].shape)
|
||||
|
||||
pbar.update(1)
|
||||
if pred_type == "normals":
|
||||
@@ -186,6 +197,9 @@ Uses Diffusers 0.28.0 Marigold pipelines.
|
||||
if pred_type == "normals":
|
||||
processed_out = torch.stack(processed_out_list, dim=0)
|
||||
processed_out = processed_out.permute(0, 2, 3, 1)
|
||||
elif pred_type == "intrinsics":
|
||||
processed_out = torch.cat(processed_out_list, dim=0)
|
||||
processed_out = processed_out.permute(0, 2, 3, 1)
|
||||
else:
|
||||
processed_out = torch.cat(processed_out_list, dim=0)
|
||||
processed_out = processed_out.permute(0, 2, 3, 1).repeat(1, 1, 1, 3)
|
||||
|
||||
+1
-1
@@ -1,5 +1,5 @@
|
||||
accelerate>=0.22.0
|
||||
diffusers>=0.28.0
|
||||
diffusers>=0.33.0
|
||||
matplotlib
|
||||
scipy
|
||||
torch>=2.0.1
|
||||
|
||||
Reference in New Issue
Block a user