Add metric models
This commit is contained in:
+30
-12
@@ -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)
|
||||
|
||||
|
||||
@@ -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 = {
|
||||
|
||||
Reference in New Issue
Block a user