diff --git a/nodes.py b/nodes.py index 9dd6e05..116b84f 100644 --- a/nodes.py +++ b/nodes.py @@ -32,14 +32,19 @@ class DownloadAndLoadDepthAnythingV2Model: 'depth_anything_v2_vitb_fp32.safetensors', 'depth_anything_v2_vitl_fp16.safetensors', 'depth_anything_v2_vitl_fp32.safetensors', + 'depth_anything_v2_vitg_fp32.safetensors', 'depth_anything_v2_metric_hypersim_vitl_fp32.safetensors', 'depth_anything_v2_metric_vkitti_vitl_fp32.safetensors' ], - { - "default": 'depth_anything_v2_vitl_fp32.safetensors' - }), - }, - } + { + "default": 'depth_anything_v2_vitl_fp32.safetensors' + }), + }, + "optional": { + "precision": (["auto", "bf16", "fp16", "fp32"], {"default": "auto"},) + } + } + RETURN_TYPES = ("DAMODEL",) RETURN_NAMES = ("da_v2_model",) @@ -52,59 +57,70 @@ https://huggingface.co/Kijai/DepthAnythingV2-safetensors/tree/main fp16 reduces quality by a LOT, not recommended. """ - def loadmodel(self, model): + def loadmodel(self, model, precision="fp32"): device = mm.get_torch_device() - dtype = torch.float16 if "fp16" in model else torch.float32 + if precision == "auto": + dtype = torch.float16 if "fp16" in model else torch.float32 + elif precision == "bf16": + dtype = torch.bfloat16 + elif precision == "fp16": + dtype = torch.float16 + elif precision == "fp32": + dtype = torch.float32 + model_configs = { 'vits': {'encoder': 'vits', 'features': 64, 'out_channels': [48, 96, 192, 384]}, 'vitb': {'encoder': 'vitb', 'features': 128, 'out_channels': [96, 192, 384, 768]}, 'vitl': {'encoder': 'vitl', 'features': 256, 'out_channels': [256, 512, 1024, 1024]}, - #'vitg': {'encoder': 'vitg', 'features': 384, 'out_channels': [1536, 1536, 1536, 1536]} + 'vitg': {'encoder': 'vitg', 'features': 384, 'out_channels': [1536, 1536, 1536, 1536]} } - custom_config = { - 'model_name': model, - } - if not hasattr(self, 'model') or self.model == None or custom_config != self.current_config: - self.current_config = custom_config - download_path = os.path.join(folder_paths.models_dir, "depthanything") - model_path = os.path.join(download_path, model) + + download_path = os.path.join(folder_paths.models_dir, "depthanything") + model_path = os.path.join(download_path, model) - if not os.path.exists(model_path): - print(f"Downloading model to: {model_path}") - from huggingface_hub import snapshot_download - snapshot_download(repo_id="Kijai/DepthAnythingV2-safetensors", - allow_patterns=[f"*{model}*"], - local_dir=download_path, - local_dir_use_symlinks=False) + if "vitg" in model: + repo = "Nap/depth_anything_v2_vitg" + else: + repo = "Kijai/DepthAnythingV2-safetensors" - print(f"Loading model from: {model_path}") + if not os.path.exists(model_path): + print(f"Downloading model to: {model_path}") + from huggingface_hub import snapshot_download + snapshot_download(repo_id=repo, + allow_patterns=[f"*{model}*"], + local_dir=download_path, + local_dir_use_symlinks=False) - if "vitl" in model: - encoder = "vitl" - elif "vitb" in model: - encoder = "vitb" - elif "vits" in model: - encoder = "vits" + print(f"Loading model from: {model_path}") - if "hypersim" in model: - max_depth = 20.0 + if "vitg" in model: + encoder = "vitg" + elif "vitl" in model: + encoder = "vitl" + elif "vitb" in model: + encoder = "vitb" + elif "vits" in model: + encoder = "vits" + + if "hypersim" in model: + max_depth = 20.0 + else: + max_depth = 80.0 + + with (init_empty_weights() if is_accelerate_available else nullcontext()): + if 'metric' in model: + self.model = DepthAnythingV2(**{**model_configs[encoder], 'is_metric': True, 'max_depth': max_depth}) else: - max_depth = 80.0 + self.model = DepthAnythingV2(**model_configs[encoder]) + + state_dict = load_torch_file(model_path) + if is_accelerate_available: + for key in state_dict: + set_module_tensor_to_device(self.model, key, device=device, dtype=dtype, value=state_dict[key]) + else: + self.model.load_state_dict(state_dict) - with (init_empty_weights() if is_accelerate_available else nullcontext()): - if 'metric' in model: - self.model = DepthAnythingV2(**{**model_configs[encoder], 'is_metric': True, 'max_depth': max_depth}) - else: - self.model = DepthAnythingV2(**model_configs[encoder]) - - state_dict = load_torch_file(model_path) - if is_accelerate_available: - for key in state_dict: - set_module_tensor_to_device(self.model, key, device=device, dtype=dtype, value=state_dict[key]) - else: - self.model.load_state_dict(state_dict) - - self.model.eval() + self.model.eval() da_model = { "model": self.model, @@ -162,15 +178,15 @@ https://depth-anything-v2.github.io depth = (depth - depth.min()) / (depth.max() - depth.min()) out.append(depth.cpu()) pbar.update(1) - model.to(offload_device) - depth_out = torch.cat(out, dim=0) - depth_out = depth_out.unsqueeze(-1).repeat(1, 1, 1, 3).cpu().float() + + model.to(offload_device) + mm.soft_empty_cache() + depth_out = torch.cat(out, dim=0) + depth_out = depth_out.unsqueeze(-1).repeat(1, 1, 1, 3).cpu().float() final_H = (orig_H // 2) * 2 final_W = (orig_W // 2) * 2 - - if depth_out.shape[1] != final_H or depth_out.shape[2] != final_W: depth_out = F.interpolate(depth_out.permute(0, 3, 1, 2), size=(final_H, final_W), mode="bilinear").permute(0, 2, 3, 1) depth_out = (depth_out - depth_out.min()) / (depth_out.max() - depth_out.min())