Fixes
This commit is contained in:
@@ -10,30 +10,15 @@
|
||||
from functools import partial
|
||||
import math
|
||||
import logging
|
||||
from typing import Sequence, Tuple, Union, Callable
|
||||
from typing import Sequence, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.utils.checkpoint
|
||||
from torch.nn.init import trunc_normal_
|
||||
|
||||
from .dinov2_layers import Mlp, PatchEmbed, SwiGLUFFNFused, MemEffAttention, NestedTensorBlock as Block
|
||||
|
||||
|
||||
logger = logging.getLogger("dinov2")
|
||||
|
||||
|
||||
def named_apply(fn: Callable, module: nn.Module, name="", depth_first=True, include_root=False) -> nn.Module:
|
||||
if not depth_first and include_root:
|
||||
fn(module=module, name=name)
|
||||
for child_name, child_module in module.named_children():
|
||||
child_name = ".".join((name, child_name)) if name else child_name
|
||||
named_apply(fn=fn, module=child_module, name=child_name, depth_first=depth_first, include_root=True)
|
||||
if depth_first and include_root:
|
||||
fn(module=module, name=name)
|
||||
return module
|
||||
|
||||
|
||||
class BlockChunk(nn.ModuleList):
|
||||
def forward(self, x):
|
||||
for b in self:
|
||||
@@ -167,14 +152,7 @@ class DinoVisionTransformer(nn.Module):
|
||||
|
||||
self.mask_token = nn.Parameter(torch.zeros(1, embed_dim))
|
||||
|
||||
self.init_weights()
|
||||
|
||||
def init_weights(self):
|
||||
trunc_normal_(self.pos_embed, std=0.02)
|
||||
nn.init.normal_(self.cls_token, std=1e-6)
|
||||
if self.register_tokens is not None:
|
||||
nn.init.normal_(self.register_tokens, std=1e-6)
|
||||
named_apply(init_weights_vit_timm, self)
|
||||
pass
|
||||
|
||||
def interpolate_pos_encoding(self, x, w, h):
|
||||
previous_dtype = x.dtype
|
||||
@@ -193,7 +171,6 @@ class DinoVisionTransformer(nn.Module):
|
||||
# DINOv2 with register modify the interpolate_offset from 0.1 to 0.0
|
||||
w0, h0 = w0 + self.interpolate_offset, h0 + self.interpolate_offset
|
||||
# w0, h0 = w0 + 0.1, h0 + 0.1
|
||||
|
||||
sqrt_N = math.sqrt(N)
|
||||
sx, sy = float(w0) / sqrt_N, float(h0) / sqrt_N
|
||||
patch_pos_embed = nn.functional.interpolate(
|
||||
@@ -328,13 +305,6 @@ class DinoVisionTransformer(nn.Module):
|
||||
return self.head(ret["x_norm_clstoken"])
|
||||
|
||||
|
||||
def init_weights_vit_timm(module: nn.Module, name: str = ""):
|
||||
"""ViT weight initialization, original timm impl (for reproducibility)"""
|
||||
if isinstance(module, nn.Linear):
|
||||
trunc_normal_(module.weight, std=0.02)
|
||||
if module.bias is not None:
|
||||
nn.init.zeros_(module.bias)
|
||||
|
||||
|
||||
def vit_small(patch_size=16, num_register_tokens=0, **kwargs):
|
||||
model = DinoVisionTransformer(
|
||||
|
||||
@@ -11,15 +11,6 @@ import folder_paths
|
||||
|
||||
from .depth_anything_v2.dpt import DepthAnythingV2
|
||||
|
||||
from contextlib import nullcontext
|
||||
try:
|
||||
from accelerate import init_empty_weights
|
||||
from accelerate.utils import set_module_tensor_to_device
|
||||
is_accelerate_available = True
|
||||
except:
|
||||
is_accelerate_available = False
|
||||
pass
|
||||
|
||||
class DownloadAndLoadDepthAnythingV2Model:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -51,14 +42,13 @@ class DownloadAndLoadDepthAnythingV2Model:
|
||||
FUNCTION = "loadmodel"
|
||||
CATEGORY = "DepthAnythingV2"
|
||||
DESCRIPTION = """
|
||||
Models autodownload to `ComfyUI\models\depthanything` from
|
||||
Models autodownload to `ComfyUI/models/depthanything` from
|
||||
https://huggingface.co/Kijai/DepthAnythingV2-safetensors/tree/main
|
||||
|
||||
fp16 reduces quality by a LOT, not recommended.
|
||||
"""
|
||||
|
||||
def loadmodel(self, model, precision="fp32"):
|
||||
device = mm.get_torch_device()
|
||||
if precision == "auto":
|
||||
dtype = torch.float16 if "fp16" in model else torch.float32
|
||||
elif precision == "bf16":
|
||||
@@ -107,18 +97,13 @@ fp16 reduces quality by a LOT, not recommended.
|
||||
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:
|
||||
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.load_state_dict(state_dict, strict=False)
|
||||
|
||||
self.model.eval()
|
||||
|
||||
@@ -182,15 +167,17 @@ https://depth-anything-v2.github.io
|
||||
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()
|
||||
depth_out = depth_out.unsqueeze(-1).expand(-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())
|
||||
depth_out = torch.clamp(depth_out, 0, 1)
|
||||
depth_min = depth_out.min()
|
||||
depth_max = depth_out.max()
|
||||
depth_out.sub_(depth_min).div_(depth_max - depth_min)
|
||||
depth_out.clamp_(0, 1)
|
||||
if da_model['is_metric']:
|
||||
depth_out = 1 - depth_out
|
||||
return (depth_out,)
|
||||
|
||||
+1
-1
@@ -1,7 +1,7 @@
|
||||
[project]
|
||||
name = "comfyui-depthanythingv2"
|
||||
description = "ComfyUI nodes to use [a/DepthAnythingV2](https://depth-anything-v2.github.io/)\nNOTE:Models autodownload to ComfyUI/models/depthanything from [a/https://huggingface.co/Kijai/DepthAnythingV2-safetensors/tree/main](https://huggingface.co/Kijai/DepthAnythingV2-safetensors/tree/main)"
|
||||
version = "1.0.1"
|
||||
version = "1.0.2"
|
||||
license = "LICENSE"
|
||||
dependencies = ["huggingface_hub", "accelerate"]
|
||||
|
||||
|
||||
Reference in New Issue
Block a user