Use accelerate for faster model loading

This commit is contained in:
kijai
2024-06-28 23:46:35 +03:00
parent 15f3c93370
commit 047a4ecfd0
7 changed files with 52 additions and 32 deletions
+4 -3
View File
@@ -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:
@@ -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:
+12 -9
View File
@@ -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(),
)
+9 -8
View File
@@ -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)
+15 -2
View File
@@ -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
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)
dtype = torch.float16 if "fp16" in model else torch.float32
self.model = self.model.to(dtype).to(device).eval()
self.model.eval()
da_model = {
"model": self.model,
"dtype": dtype,
+2 -2
View File
@@ -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"
+1
View File
@@ -1 +1,2 @@
huggingface_hub
accelerate