fix RemapDepth on older pytorch

This commit is contained in:
kijai
2023-12-15 22:30:41 +02:00
parent 6b143fbee2
commit b7bf6f745f
2 changed files with 3 additions and 5 deletions
+1 -4
View File
@@ -26,7 +26,6 @@ def ensemble_depths(input_images, regularizer_strength=0.02, max_iter=2, tol=1e-
by aligning estimating the scale and shift
"""
device = input_images.device
dtype = np.float32
original_input = input_images.clone()
n_img = input_images.shape[0]
ori_shape = input_images.shape
@@ -48,7 +47,7 @@ def ensemble_depths(input_images, regularizer_strength=0.02, max_iter=2, tol=1e-
# objective function
def closure(x):
x = x.astype(dtype)
x = x.astype(np.float32)
l = len(x)
s = x[:int(l/2)]
t = x[int(l/2):]
@@ -101,6 +100,4 @@ def ensemble_depths(input_images, regularizer_strength=0.02, max_iter=2, tol=1e-
_max = torch.max(aligned_images)
aligned_images = (aligned_images - _min) / (_max - _min)
uncertainty /= (_max - _min)
return aligned_images, uncertainty
+2 -1
View File
@@ -310,7 +310,8 @@ class RemapDepth:
CATEGORY = "Marigold"
def remap(self, image, min, max, clamp):
if image.dtype == torch.float16:
image = image.to(torch.float32)
image = min + image * (max - min)
if clamp:
image = torch.clamp(image, min=0.0, max=1.0)