diff --git a/nodes.py b/nodes.py index e9f7665..5e39b32 100644 --- a/nodes.py +++ b/nodes.py @@ -11,19 +11,32 @@ class MarigoldDepthEstimation: return {"required": { "image": ("IMAGE", ), "seed": ("INT", {"default": 123,"min": 0, "max": 0xffffffffffffffff, "step": 1}), - "denoise_steps": ("INT", {"default": 10, "min": 0, "max": 4096, "step": 1}), + "denoise_steps": ("INT", {"default": 10, "min": 1, "max": 4096, "step": 1}), "n_repeat": ("INT", {"default": 2, "min": 2, "max": 4096, "step": 1}), + "regularizer_strength": ("FLOAT", {"default": 0.02, "min": 0.001, "max": 4096, "step": 0.001}), + "reduction_method": ( + [ + 'median', + 'mean', + ], { + "default": 'median' + }), + "max_iter": ("INT", {"default": 5, "min": 1, "max": 4096, "step": 1}), + "tol": ("FLOAT", {"default": 1e-3, "min": 1e-6, "max": 1e-1, "step": 1e-6}), + "invert": ("BOOLEAN", {"default": True}), }, + } - RETURN_TYPES = ("IMAGE","IMAGE",) - RETURN_NAMES =("ensembled_image","depth_images",) + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES =("ensembled_image",) FUNCTION = "process" CATEGORY = "Marigold" - def process(self, image, seed, denoise_steps, n_repeat, invert): + def process(self, image, seed, denoise_steps, n_repeat, regularizer_strength, reduction_method, max_iter, tol,invert): + batch_size = image.shape[0] device = torch.device("cuda" if torch.cuda.is_available() else "cpu") torch.manual_seed(seed) image = image.permute(0, 3, 1, 2).to(device) @@ -40,44 +53,48 @@ class MarigoldDepthEstimation: self.marigold_pipeline = self.marigold_pipeline.to(device) self.marigold_pipeline.unet.eval() # Set the model to evaluation mode - print(image.shape) - depth_maps = [] + out = [] + for i in range(batch_size): + depth_maps = [] - with torch.no_grad(): - for _ in range(n_repeat): - depth_map = self.marigold_pipeline(image, num_inference_steps=denoise_steps) # Process the image tensor to get the depth map - depth_map = torch.clip(depth_map, -1.0, 1.0) - depth_map = (depth_map + 1.0) / 2.0 - depth_maps.append(depth_map) - - depth_predictions = torch.cat(depth_maps, dim=0).squeeze() - - torch.cuda.empty_cache() # clear vram cache for ensembling + with torch.no_grad(): + for _ in range(n_repeat): + depth_map = self.marigold_pipeline(image[i].unsqueeze(0), num_inference_steps=denoise_steps, show_pbar=True) # Process the image tensor to get the depth map + depth_map = torch.clip(depth_map, -1.0, 1.0) + depth_map = (depth_map + 1.0) / 2.0 + depth_maps.append(depth_map) + + depth_predictions = torch.cat(depth_maps, dim=0).squeeze() + + torch.cuda.empty_cache() # clear vram cache for ensembling - #ensemble parameters - regularizer_strength = 0.02 - max_iter = 5 - tol = 1e-3 - reduction_method = "median" - merging_max_res = None + #ensemble parameters + #regularizer_strength = 0.02 + #max_iter = 5 + #tol = 1e-3 + #reduction_method = "median" + merging_max_res = None + + # Test-time ensembling + if n_repeat > 1: + depth_map, pred_uncert = ensemble_depths( + depth_predictions, + regularizer_strength=regularizer_strength, + max_iter=max_iter, + tol=tol, + reduction=reduction_method, + max_res=merging_max_res, + device=device, + ) + depth_map = depth_map.unsqueeze(2).repeat(1, 1, 3) + out.append(depth_map) - # Test-time ensembling - if n_repeat > 1: - depth_map, pred_uncert = ensemble_depths( - depth_predictions, - regularizer_strength=regularizer_strength, - max_iter=max_iter, - tol=tol, - reduction=reduction_method, - max_res=merging_max_res, - device=device, - ) - - depth_map = depth_map.unsqueeze_(0).to(dtype=torch.float32) - if invert: - depth_map = 1 - depth_map - return (depth_map, depth_predictions,) + outstack = 1.0 - torch.stack(out, dim=0) + else: + outstack = torch.stack(out, dim=0) + + return (outstack,) NODE_CLASS_MAPPINGS = { "MarigoldDepthEstimation": MarigoldDepthEstimation, diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..10b21e0 --- /dev/null +++ b/requirements.txt @@ -0,0 +1,6 @@ +accelerate>=0.22.0 +diffusers>=0.20.1 +matplotlib +scipy +torch>=2.0.1 +transformers>=4.32.1 \ No newline at end of file