diff --git a/depth_anything_v2/dinov2_layers/attention.py b/depth_anything_v2/dinov2_layers/attention.py index 815a2bf..a209583 100644 --- a/depth_anything_v2/dinov2_layers/attention.py +++ b/depth_anything_v2/dinov2_layers/attention.py @@ -12,7 +12,8 @@ import logging from torch import Tensor from torch import nn - +import comfy.ops +ops = comfy.ops.manual_cast logger = logging.getLogger("dinov2") @@ -41,9 +42,9 @@ class Attention(nn.Module): head_dim = dim // num_heads self.scale = head_dim**-0.5 - self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias) + self.qkv = ops.Linear(dim, dim * 3, bias=qkv_bias) self.attn_drop = nn.Dropout(attn_drop) - self.proj = nn.Linear(dim, dim, bias=proj_bias) + self.proj = ops.Linear(dim, dim, bias=proj_bias) self.proj_drop = nn.Dropout(proj_drop) def forward(self, x: Tensor) -> Tensor: diff --git a/depth_anything_v2/dinov2_layers/patch_embed.py b/depth_anything_v2/dinov2_layers/patch_embed.py index 574abe4..5a43eeb 100644 --- a/depth_anything_v2/dinov2_layers/patch_embed.py +++ b/depth_anything_v2/dinov2_layers/patch_embed.py @@ -12,7 +12,8 @@ from typing import Callable, Optional, Tuple, Union from torch import Tensor import torch.nn as nn - +import comfy.ops +ops = comfy.ops.manual_cast def make_2tuple(x): if isinstance(x, tuple): @@ -63,7 +64,7 @@ class PatchEmbed(nn.Module): self.flatten_embedding = flatten_embedding - self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_HW, stride=patch_HW) + self.proj = ops.Conv2d(in_chans, embed_dim, kernel_size=patch_HW, stride=patch_HW) self.norm = norm_layer(embed_dim) if norm_layer else nn.Identity() def forward(self, x: Tensor) -> Tensor: diff --git a/depth_anything_v2/dpt.py b/depth_anything_v2/dpt.py index 1e8040c..e79a02c 100644 --- a/depth_anything_v2/dpt.py +++ b/depth_anything_v2/dpt.py @@ -5,6 +5,9 @@ import torch.nn.functional as F from .dinov2 import DINOv2 from .util.blocks import FeatureFusionBlock, _make_scratch +import comfy.ops +ops = comfy.ops.manual_cast + def _make_fusion_block(features, use_bn, size=None): return FeatureFusionBlock( features, @@ -21,7 +24,7 @@ class ConvBlock(nn.Module): super().__init__() self.conv_block = nn.Sequential( - nn.Conv2d(in_feature, out_feature, kernel_size=3, stride=1, padding=1), + ops.Conv2d(in_feature, out_feature, kernel_size=3, stride=1, padding=1), nn.BatchNorm2d(out_feature), nn.ReLU(True) ) @@ -45,7 +48,7 @@ class DPTHead(nn.Module): self.is_metric=is_metric self.projects = nn.ModuleList([ - nn.Conv2d( + ops.Conv2d( in_channels=in_channels, out_channels=out_channel, kernel_size=1, @@ -68,7 +71,7 @@ class DPTHead(nn.Module): stride=2, padding=0), nn.Identity(), - nn.Conv2d( + ops.Conv2d( in_channels=out_channels[3], out_channels=out_channels[3], kernel_size=3, @@ -81,7 +84,7 @@ class DPTHead(nn.Module): for _ in range(len(self.projects)): self.readout_projects.append( nn.Sequential( - nn.Linear(2 * in_channels, in_channels), + ops.Linear(2 * in_channels, in_channels), nn.GELU())) self.scratch = _make_scratch( @@ -101,19 +104,19 @@ class DPTHead(nn.Module): head_features_1 = features 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_conv1 = ops.Conv2d(head_features_1, head_features_1 // 2, kernel_size=3, stride=1, padding=1) 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), + ops.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), + ops.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), + ops.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), + ops.Conv2d(head_features_2, 1, kernel_size=1, stride=1, padding=0), nn.ReLU(True), nn.Identity(), ) diff --git a/depth_anything_v2/util/blocks.py b/depth_anything_v2/util/blocks.py index 382ea18..331a4aa 100644 --- a/depth_anything_v2/util/blocks.py +++ b/depth_anything_v2/util/blocks.py @@ -1,5 +1,6 @@ import torch.nn as nn - +import comfy.ops +ops = comfy.ops.manual_cast def _make_scratch(in_shape, out_shape, groups=1, expand=False): scratch = nn.Module() @@ -17,11 +18,11 @@ def _make_scratch(in_shape, out_shape, groups=1, expand=False): if len(in_shape) >= 4: out_shape4 = out_shape * 8 - scratch.layer1_rn = nn.Conv2d(in_shape[0], out_shape1, kernel_size=3, stride=1, padding=1, bias=False, groups=groups) - scratch.layer2_rn = nn.Conv2d(in_shape[1], out_shape2, kernel_size=3, stride=1, padding=1, bias=False, groups=groups) - scratch.layer3_rn = nn.Conv2d(in_shape[2], out_shape3, kernel_size=3, stride=1, padding=1, bias=False, groups=groups) + scratch.layer1_rn = ops.Conv2d(in_shape[0], out_shape1, kernel_size=3, stride=1, padding=1, bias=False, groups=groups) + scratch.layer2_rn = ops.Conv2d(in_shape[1], out_shape2, kernel_size=3, stride=1, padding=1, bias=False, groups=groups) + scratch.layer3_rn = ops.Conv2d(in_shape[2], out_shape3, kernel_size=3, stride=1, padding=1, bias=False, groups=groups) if len(in_shape) >= 4: - scratch.layer4_rn = nn.Conv2d(in_shape[3], out_shape4, kernel_size=3, stride=1, padding=1, bias=False, groups=groups) + scratch.layer4_rn = ops.Conv2d(in_shape[3], out_shape4, kernel_size=3, stride=1, padding=1, bias=False, groups=groups) return scratch @@ -42,9 +43,9 @@ class ResidualConvUnit(nn.Module): self.groups=1 - self.conv1 = nn.Conv2d(features, features, kernel_size=3, stride=1, padding=1, bias=True, groups=self.groups) + self.conv1 = ops.Conv2d(features, features, kernel_size=3, stride=1, padding=1, bias=True, groups=self.groups) - self.conv2 = nn.Conv2d(features, features, kernel_size=3, stride=1, padding=1, bias=True, groups=self.groups) + self.conv2 = ops.Conv2d(features, features, kernel_size=3, stride=1, padding=1, bias=True, groups=self.groups) if self.bn == True: self.bn1 = nn.BatchNorm2d(features) @@ -111,7 +112,7 @@ class FeatureFusionBlock(nn.Module): if self.expand == True: out_features = features // 2 - self.out_conv = nn.Conv2d(features, out_features, kernel_size=1, stride=1, padding=0, bias=True, groups=1) + self.out_conv = ops.Conv2d(features, out_features, kernel_size=1, stride=1, padding=0, bias=True, groups=1) self.resConfUnit1 = ResidualConvUnit(features, activation, bn) self.resConfUnit2 = ResidualConvUnit(features, activation, bn) diff --git a/nodes.py b/nodes.py index 36ce76e..7ebc4e9 100644 --- a/nodes.py +++ b/nodes.py @@ -11,6 +11,14 @@ 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: + pass + class DownloadAndLoadDepthAnythingV2Model: @classmethod def INPUT_TYPES(s): @@ -45,6 +53,7 @@ fp16 reduces quality by a LOT, not recommended. def loadmodel(self, model): device = mm.get_torch_device() + dtype = torch.float16 if "fp16" in model else torch.float32 model_configs = { 'vits': {'encoder': 'vits', 'features': 64, 'out_channels': [48, 96, 192, 384]}, 'vitb': {'encoder': 'vitb', 'features': 128, 'out_channels': [96, 192, 384, 768]}, @@ -81,16 +90,20 @@ fp16 reduces quality by a LOT, not recommended. 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]) + 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) - self.model.load_state_dict(state_dict) - dtype = torch.float16 if "fp16" in model else torch.float32 + 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 = self.model.to(dtype).to(device).eval() + self.model.eval() da_model = { "model": self.model, "dtype": dtype, diff --git a/pyproject.toml b/pyproject.toml index 48d1185..5b98871 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,9 +1,9 @@ [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.0" +version = "1.0.1" license = "LICENSE" -dependencies = ["huggingface_hub"] +dependencies = ["huggingface_hub", "accelerate"] [project.urls] Repository = "https://github.com/kijai/ComfyUI-DepthAnythingV2" diff --git a/requirements.txt b/requirements.txt index b2642b3..ee9cb40 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1 +1,2 @@ -huggingface_hub \ No newline at end of file +huggingface_hub +accelerate \ No newline at end of file