Add metric models

This commit is contained in:
kijai
2024-06-23 03:47:33 +03:00
parent 981b4f60a5
commit 18d9462f59
2 changed files with 45 additions and 14 deletions
+30 -12
View File
@@ -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)
+15 -2
View File
@@ -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 = {