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 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:
+12 -9
View File
@@ -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(),
) )
+9 -8
View File
@@ -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)
+15 -2
View File
@@ -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
View File
@@ -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
View File
@@ -1 +1,2 @@
huggingface_hub huggingface_hub
accelerate