Initial changes to calc_cond_batch to eventually support hook_patches

This commit is contained in:
kosinkadink1@gmail.com 2024-09-13 18:31:52 +09:00
parent 069ec7a64b
commit 3cbd40ada3

View File

@ -1,11 +1,13 @@
from .k_diffusion import sampling as k_diffusion_sampling from .k_diffusion import sampling as k_diffusion_sampling
from .extra_samplers import uni_pc from .extra_samplers import uni_pc
from typing import Dict, List, Tuple
import torch import torch
import collections import collections
from comfy import model_management from comfy import model_management
import math import math
import logging import logging
import comfy.sampler_helpers import comfy.sampler_helpers
import comfy.hooks
import scipy.stats import scipy.stats
import numpy import numpy
@ -141,7 +143,9 @@ def cond_cat(c_list):
def calc_cond_batch(model, conds, x_in, timestep, model_options): def calc_cond_batch(model, conds, x_in, timestep, model_options):
out_conds = [] out_conds = []
out_counts = [] out_counts = []
to_run = [] # separate conds by matching hooks
# TODO: implement default_conds support
hooked_to_run: Dict[comfy.hooks.HookWeightGroup,List[Tuple[Tuple,int]]] = {}
for i in range(len(conds)): for i in range(len(conds)):
out_conds.append(torch.zeros_like(x_in)) out_conds.append(torch.zeros_like(x_in))
@ -150,12 +154,15 @@ def calc_cond_batch(model, conds, x_in, timestep, model_options):
cond = conds[i] cond = conds[i]
if cond is not None: if cond is not None:
for x in cond: for x in cond:
p = get_area_and_mult(x, x_in, timestep) p = comfy.samplers.get_area_and_mult(x, x_in, timestep)
if p is None: if p is None:
continue continue
hook: comfy.hooks.HookWeightGroup = x.get('hooks', None)
hooked_to_run.setdefault(hook, list())
hooked_to_run[hook] += [(p, i)]
to_run += [(p, i)] # run every hooked_to_run separately
for hooks, to_run in hooked_to_run.items():
while len(to_run) > 0: while len(to_run) > 0:
first = to_run[0] first = to_run[0]
first_shape = first[0][0].shape first_shape = first[0][0].shape
@ -174,6 +181,7 @@ def calc_cond_batch(model, conds, x_in, timestep, model_options):
if model.memory_required(input_shape) * 1.5 < free_memory: if model.memory_required(input_shape) * 1.5 < free_memory:
to_batch = batch_amount to_batch = batch_amount
break break
# TODO: add apply_hooks call here, once a ModelPatcher ref is added to BaseModel
input_x = [] input_x = []
mult = [] mult = []