modelscope/DiffSynth-Painter
14
1import torch2import numpy as np3from .processors import Processor_id4 5 6class ControlNetConfigUnit:7 def __init__(self, processor_id: Processor_id, model_path, scale=1.0):8 self.processor_id = processor_id9 self.model_path = model_path10 self.scale = scale11 12 13class ControlNetUnit:14 def __init__(self, processor, model, scale=1.0):15 self.processor = processor16 self.model = model17 self.scale = scale18 19 20class MultiControlNetManager:21 def __init__(self, controlnet_units=[]):22 self.processors = [unit.processor for unit in controlnet_units]23 self.models = [unit.model for unit in controlnet_units]24 self.scales = [unit.scale for unit in controlnet_units]25 26 def process_image(self, image, processor_id=None):27 if processor_id is None:28 processed_image = [processor(image) for processor in self.processors]29 else:30 processed_image = [self.processors[processor_id](image)]31 processed_image = torch.concat([32 torch.Tensor(np.array(image_, dtype=np.float32) / 255).permute(2, 0, 1).unsqueeze(0)33 for image_ in processed_image34 ], dim=0)35 return processed_image36 37 def __call__(38 self,39 sample, timestep, encoder_hidden_states, conditionings,40 tiled=False, tile_size=64, tile_stride=32, **kwargs41 ):42 res_stack = None43 for processor, conditioning, model, scale in zip(self.processors, conditionings, self.models, self.scales):44 res_stack_ = model(45 sample, timestep, encoder_hidden_states, conditioning, **kwargs,46 tiled=tiled, tile_size=tile_size, tile_stride=tile_stride,47 processor_id=processor.processor_id48 )49 res_stack_ = [res * scale for res in res_stack_]50 if res_stack is None:51 res_stack = res_stack_52 else:53 res_stack = [i + j for i, j in zip(res_stack, res_stack_)]54 return res_stack55 