Use accelerate for faster model loading
This commit is contained in:
@@ -12,7 +12,8 @@ import logging
|
|||||||
|
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
from torch import nn
|
from torch import nn
|
||||||
|
import comfy.ops
|
||||||
|
ops = comfy.ops.manual_cast
|
||||||
|
|
||||||
logger = logging.getLogger("dinov2")
|
logger = logging.getLogger("dinov2")
|
||||||
|
|
||||||
@@ -41,9 +42,9 @@ class Attention(nn.Module):
|
|||||||
head_dim = dim // num_heads
|
head_dim = dim // num_heads
|
||||||
self.scale = head_dim**-0.5
|
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.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)
|
self.proj_drop = nn.Dropout(proj_drop)
|
||||||
|
|
||||||
def forward(self, x: Tensor) -> Tensor:
|
def forward(self, x: Tensor) -> Tensor:
|
||||||
|
|||||||
@@ -12,7 +12,8 @@ from typing import Callable, Optional, Tuple, Union
|
|||||||
|
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
|
import comfy.ops
|
||||||
|
ops = comfy.ops.manual_cast
|
||||||
|
|
||||||
def make_2tuple(x):
|
def make_2tuple(x):
|
||||||
if isinstance(x, tuple):
|
if isinstance(x, tuple):
|
||||||
@@ -63,7 +64,7 @@ class PatchEmbed(nn.Module):
|
|||||||
|
|
||||||
self.flatten_embedding = flatten_embedding
|
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()
|
self.norm = norm_layer(embed_dim) if norm_layer else nn.Identity()
|
||||||
|
|
||||||
def forward(self, x: Tensor) -> Tensor:
|
def forward(self, x: Tensor) -> Tensor:
|
||||||
|
|||||||
@@ -5,6 +5,9 @@ import torch.nn.functional as F
|
|||||||
from .dinov2 import DINOv2
|
from .dinov2 import DINOv2
|
||||||
from .util.blocks import FeatureFusionBlock, _make_scratch
|
from .util.blocks import FeatureFusionBlock, _make_scratch
|
||||||
|
|
||||||
|
import comfy.ops
|
||||||
|
ops = comfy.ops.manual_cast
|
||||||
|
|
||||||
def _make_fusion_block(features, use_bn, size=None):
|
def _make_fusion_block(features, use_bn, size=None):
|
||||||
return FeatureFusionBlock(
|
return FeatureFusionBlock(
|
||||||
features,
|
features,
|
||||||
@@ -21,7 +24,7 @@ class ConvBlock(nn.Module):
|
|||||||
super().__init__()
|
super().__init__()
|
||||||
|
|
||||||
self.conv_block = nn.Sequential(
|
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.BatchNorm2d(out_feature),
|
||||||
nn.ReLU(True)
|
nn.ReLU(True)
|
||||||
)
|
)
|
||||||
@@ -45,7 +48,7 @@ class DPTHead(nn.Module):
|
|||||||
self.is_metric=is_metric
|
self.is_metric=is_metric
|
||||||
|
|
||||||
self.projects = nn.ModuleList([
|
self.projects = nn.ModuleList([
|
||||||
nn.Conv2d(
|
ops.Conv2d(
|
||||||
in_channels=in_channels,
|
in_channels=in_channels,
|
||||||
out_channels=out_channel,
|
out_channels=out_channel,
|
||||||
kernel_size=1,
|
kernel_size=1,
|
||||||
@@ -68,7 +71,7 @@ class DPTHead(nn.Module):
|
|||||||
stride=2,
|
stride=2,
|
||||||
padding=0),
|
padding=0),
|
||||||
nn.Identity(),
|
nn.Identity(),
|
||||||
nn.Conv2d(
|
ops.Conv2d(
|
||||||
in_channels=out_channels[3],
|
in_channels=out_channels[3],
|
||||||
out_channels=out_channels[3],
|
out_channels=out_channels[3],
|
||||||
kernel_size=3,
|
kernel_size=3,
|
||||||
@@ -81,7 +84,7 @@ class DPTHead(nn.Module):
|
|||||||
for _ in range(len(self.projects)):
|
for _ in range(len(self.projects)):
|
||||||
self.readout_projects.append(
|
self.readout_projects.append(
|
||||||
nn.Sequential(
|
nn.Sequential(
|
||||||
nn.Linear(2 * in_channels, in_channels),
|
ops.Linear(2 * in_channels, in_channels),
|
||||||
nn.GELU()))
|
nn.GELU()))
|
||||||
|
|
||||||
self.scratch = _make_scratch(
|
self.scratch = _make_scratch(
|
||||||
@@ -101,19 +104,19 @@ class DPTHead(nn.Module):
|
|||||||
head_features_1 = features
|
head_features_1 = features
|
||||||
head_features_2 = 32
|
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:
|
if self.is_metric:
|
||||||
self.scratch.output_conv2 = nn.Sequential(
|
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.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()
|
nn.Sigmoid()
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
self.scratch.output_conv2 = nn.Sequential(
|
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.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.ReLU(True),
|
||||||
nn.Identity(),
|
nn.Identity(),
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -1,5 +1,6 @@
|
|||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
|
import comfy.ops
|
||||||
|
ops = comfy.ops.manual_cast
|
||||||
|
|
||||||
def _make_scratch(in_shape, out_shape, groups=1, expand=False):
|
def _make_scratch(in_shape, out_shape, groups=1, expand=False):
|
||||||
scratch = nn.Module()
|
scratch = nn.Module()
|
||||||
@@ -17,11 +18,11 @@ def _make_scratch(in_shape, out_shape, groups=1, expand=False):
|
|||||||
if len(in_shape) >= 4:
|
if len(in_shape) >= 4:
|
||||||
out_shape4 = out_shape * 8
|
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.layer1_rn = ops.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.layer2_rn = ops.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.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:
|
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
|
return scratch
|
||||||
|
|
||||||
@@ -42,9 +43,9 @@ class ResidualConvUnit(nn.Module):
|
|||||||
|
|
||||||
self.groups=1
|
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:
|
if self.bn == True:
|
||||||
self.bn1 = nn.BatchNorm2d(features)
|
self.bn1 = nn.BatchNorm2d(features)
|
||||||
@@ -111,7 +112,7 @@ class FeatureFusionBlock(nn.Module):
|
|||||||
if self.expand == True:
|
if self.expand == True:
|
||||||
out_features = features // 2
|
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.resConfUnit1 = ResidualConvUnit(features, activation, bn)
|
||||||
self.resConfUnit2 = ResidualConvUnit(features, activation, bn)
|
self.resConfUnit2 = ResidualConvUnit(features, activation, bn)
|
||||||
|
|||||||
@@ -11,6 +11,14 @@ import folder_paths
|
|||||||
|
|
||||||
from .depth_anything_v2.dpt import DepthAnythingV2
|
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:
|
class DownloadAndLoadDepthAnythingV2Model:
|
||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(s):
|
def INPUT_TYPES(s):
|
||||||
@@ -45,6 +53,7 @@ fp16 reduces quality by a LOT, not recommended.
|
|||||||
|
|
||||||
def loadmodel(self, model):
|
def loadmodel(self, model):
|
||||||
device = mm.get_torch_device()
|
device = mm.get_torch_device()
|
||||||
|
dtype = torch.float16 if "fp16" in model else torch.float32
|
||||||
model_configs = {
|
model_configs = {
|
||||||
'vits': {'encoder': 'vits', 'features': 64, 'out_channels': [48, 96, 192, 384]},
|
'vits': {'encoder': 'vits', 'features': 64, 'out_channels': [48, 96, 192, 384]},
|
||||||
'vitb': {'encoder': 'vitb', 'features': 128, 'out_channels': [96, 192, 384, 768]},
|
'vitb': {'encoder': 'vitb', 'features': 128, 'out_channels': [96, 192, 384, 768]},
|
||||||
@@ -81,16 +90,20 @@ fp16 reduces quality by a LOT, not recommended.
|
|||||||
else:
|
else:
|
||||||
max_depth = 80.0
|
max_depth = 80.0
|
||||||
|
|
||||||
|
with (init_empty_weights() if is_accelerate_available else nullcontext()):
|
||||||
if 'metric' in model:
|
if 'metric' in model:
|
||||||
self.model = DepthAnythingV2(**{**model_configs[encoder], 'is_metric': True, 'max_depth': max_depth})
|
self.model = DepthAnythingV2(**{**model_configs[encoder], 'is_metric': True, 'max_depth': max_depth})
|
||||||
else:
|
else:
|
||||||
self.model = DepthAnythingV2(**model_configs[encoder])
|
self.model = DepthAnythingV2(**model_configs[encoder])
|
||||||
|
|
||||||
state_dict = load_torch_file(model_path)
|
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)
|
||||||
dtype = torch.float16 if "fp16" in model else torch.float32
|
|
||||||
|
|
||||||
self.model = self.model.to(dtype).to(device).eval()
|
self.model.eval()
|
||||||
da_model = {
|
da_model = {
|
||||||
"model": self.model,
|
"model": self.model,
|
||||||
"dtype": dtype,
|
"dtype": dtype,
|
||||||
|
|||||||
+2
-2
@@ -1,9 +1,9 @@
|
|||||||
[project]
|
[project]
|
||||||
name = "comfyui-depthanythingv2"
|
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)"
|
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"
|
license = "LICENSE"
|
||||||
dependencies = ["huggingface_hub"]
|
dependencies = ["huggingface_hub", "accelerate"]
|
||||||
|
|
||||||
[project.urls]
|
[project.urls]
|
||||||
Repository = "https://github.com/kijai/ComfyUI-DepthAnythingV2"
|
Repository = "https://github.com/kijai/ComfyUI-DepthAnythingV2"
|
||||||
|
|||||||
@@ -1 +1,2 @@
|
|||||||
huggingface_hub
|
huggingface_hub
|
||||||
|
accelerate
|
||||||
Reference in New Issue
Block a user