This commit is contained in:
kijai
2026-03-06 18:38:50 +02:00
parent d505cbca99
commit 553187872e
3 changed files with 21 additions and 64 deletions
+2 -32
View File
@@ -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(
+7 -20
View File
@@ -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
View File
@@ -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"]