From 18d9462f59791206aebe4e0aaa532e4d1a870c40 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sun, 23 Jun 2024 03:47:33 +0300 Subject: [PATCH] Add metric models --- depth_anything_v2/dpt.py | 42 ++++++++++++++++++++++++++++------------ nodes.py | 17 ++++++++++++++-- 2 files changed, 45 insertions(+), 14 deletions(-) diff --git a/depth_anything_v2/dpt.py b/depth_anything_v2/dpt.py index 57eca74..c221f7d 100644 --- a/depth_anything_v2/dpt.py +++ b/depth_anything_v2/dpt.py @@ -42,11 +42,13 @@ class DPTHead(nn.Module): features=256, use_bn=False, out_channels=[256, 512, 1024, 1024], - use_clstoken=False + use_clstoken=False, + is_metric=False ): super(DPTHead, self).__init__() self.use_clstoken = use_clstoken + self.is_metric=is_metric self.projects = nn.ModuleList([ nn.Conv2d( @@ -106,13 +108,21 @@ class DPTHead(nn.Module): head_features_2 = 32 self.scratch.output_conv1 = nn.Conv2d(head_features_1, head_features_1 // 2, kernel_size=3, stride=1, padding=1) - self.scratch.output_conv2 = nn.Sequential( - nn.Conv2d(head_features_1 // 2, head_features_2, kernel_size=3, stride=1, padding=1), - nn.ReLU(True), - nn.Conv2d(head_features_2, 1, kernel_size=1, stride=1, padding=0), - nn.ReLU(True), - nn.Identity(), - ) + if self.is_metric: + self.scratch.output_conv2 = nn.Sequential( + nn.Conv2d(head_features_1 // 2, head_features_2, kernel_size=3, stride=1, padding=1), + nn.ReLU(True), + nn.Conv2d(head_features_2, 1, kernel_size=1, stride=1, padding=0), + nn.Sigmoid() + ) + else: + self.scratch.output_conv2 = nn.Sequential( + nn.Conv2d(head_features_1 // 2, head_features_2, kernel_size=3, stride=1, padding=1), + nn.ReLU(True), + nn.Conv2d(head_features_2, 1, kernel_size=1, stride=1, padding=0), + nn.ReLU(True), + nn.Identity(), + ) def forward(self, out_features, patch_h, patch_w): out = [] @@ -157,7 +167,9 @@ class DepthAnythingV2(nn.Module): features=256, out_channels=[256, 512, 1024, 1024], use_bn=False, - use_clstoken=False + use_clstoken=False, + is_metric=False, + max_depth=20.0 ): super(DepthAnythingV2, self).__init__() @@ -168,18 +180,24 @@ class DepthAnythingV2(nn.Module): 'vitg': [9, 19, 29, 39] } + self.is_metric = is_metric + self.max_depth = max_depth + self.encoder = encoder self.pretrained = DINOv2(model_name=encoder) - self.depth_head = DPTHead(self.pretrained.embed_dim, features, use_bn, out_channels=out_channels, use_clstoken=use_clstoken) + self.depth_head = DPTHead(self.pretrained.embed_dim, features, use_bn, out_channels=out_channels, use_clstoken=use_clstoken, is_metric=is_metric) def forward(self, x): patch_h, patch_w = x.shape[-2] // 14, x.shape[-1] // 14 features = self.pretrained.get_intermediate_layers(x, self.intermediate_layer_idx[self.encoder], return_class_token=True) - depth = self.depth_head(features, patch_h, patch_w) - depth = F.relu(depth) + if self.is_metric: + depth = self.depth_head(features, patch_h, patch_w) * self.max_depth + else: + depth = self.depth_head(features, patch_h, patch_w) + depth = F.relu(depth) return depth.squeeze(1) diff --git a/nodes.py b/nodes.py index fc89be6..1695fce 100644 --- a/nodes.py +++ b/nodes.py @@ -24,6 +24,8 @@ class DownloadAndLoadDepthAnythingV2Model: 'depth_anything_v2_vitb_fp32.safetensors', 'depth_anything_v2_vitl_fp16.safetensors', 'depth_anything_v2_vitl_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' @@ -75,7 +77,15 @@ fp16 reduces quality by a LOT, not recommended. elif "vits" in model: encoder = "vits" - self.model = DepthAnythingV2(**model_configs[encoder]) + if "hypersim" in model: + max_depth = 20.0 + else: + max_depth = 80.0 + + 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) self.model.load_state_dict(state_dict) @@ -84,7 +94,8 @@ fp16 reduces quality by a LOT, not recommended. self.model = self.model.to(dtype).to(device).eval() da_model = { "model": self.model, - "dtype": dtype + "dtype": dtype, + "is_metric": self.model.is_metric } return (da_model,) @@ -146,6 +157,8 @@ https://depth-anything-v2.github.io 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="bicubic").permute(0, 2, 3, 1) + if da_model['is_metric']: + depth_out = 1 - depth_out return (depth_out,) NODE_CLASS_MAPPINGS = {