From 0b3e70e23c4482772495bb97451919ef3ca2757e Mon Sep 17 00:00:00 2001 From: jason Date: Tue, 9 Jan 2024 00:12:52 -0600 Subject: [PATCH] batch pooled condition and reference controlnet when available --- __init__.py | 17 +++++++++++++++-- 1 file changed, 15 insertions(+), 2 deletions(-) diff --git a/__init__.py b/__init__.py index d08588f..9170eef 100644 --- a/__init__.py +++ b/__init__.py @@ -283,8 +283,21 @@ class StyleAlignedReferenceSampler: # Add reference conditioning to batch batched_condition = [] for i,condition in enumerate(positive): - bath_with_reference = torch.cat([ref_positive[i][0], condition[0].repeat([batch_size] + [1] * len(condition[0].shape[1:]))], dim=0) - batched_condition.append([bath_with_reference, condition[1]]) + additional = condition[1].copy() + batch_with_reference = torch.cat([ref_positive[i][0], condition[0].repeat([batch_size] + [1] * len(condition[0].shape[1:]))], dim=0) + if 'pooled_output' in additional and 'pooled_output' in ref_positive[i][1]: + # combine pooled output + pooled_output = torch.cat([ref_positive[i][1]['pooled_output'], additional['pooled_output'].repeat([batch_size] + + [1] * len(additional['pooled_output'].shape[1:]))], dim=0) + additional['pooled_output'] = pooled_output + if 'control' in additional and 'control' in ref_positive[i][1]: + # combine control conditioning + control_hint = torch.cat([ref_positive[i][1]['control'].cond_hint_original, additional['control'].cond_hint_original.repeat([batch_size] + + [1] * len(additional['control'].cond_hint_original.shape[1:]))], dim=0) + cloned_controlnet = additional['control'].copy() + cloned_controlnet.set_cond_hint(control_hint, strength=additional['control'].strength, timestep_percent_range=additional['control'].timestep_percent_range) + additional['control'] = cloned_controlnet + batched_condition.append([batch_with_reference, additional]) disable_pbar = not comfy.utils.PROGRESS_BAR_ENABLED samples = comfy.sample.sample_custom(