hugging-apps/echo-memory
0
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, skip_processor=False):8 self.processor_id = processor_id9 self.model_path = model_path10 self.scale = scale11 self.skip_processor = skip_processor12 13 14class ControlNetUnit:15 def __init__(self, processor, model, scale=1.0):16 self.processor = processor17 self.model = model18 self.scale = scale19 20 21class MultiControlNetManager:22 def __init__(self, controlnet_units=[]):23 self.processors = [unit.processor for unit in controlnet_units]24 self.models = [unit.model for unit in controlnet_units]25 self.scales = [unit.scale for unit in controlnet_units]26 27 def cpu(self):28 for model in self.models:29 model.cpu()30 31 def to(self, device):32 for model in self.models:33 model.to(device)34 for processor in self.processors:35 processor.to(device)36 37 def process_image(self, image, processor_id=None):38 if processor_id is None:39 processed_image = [processor(image) for processor in self.processors]40 else:41 processed_image = [self.processors[processor_id](image)]42 processed_image = torch.concat([43 torch.Tensor(np.array(image_, dtype=np.float32) / 255).permute(2, 0, 1).unsqueeze(0)44 for image_ in processed_image45 ], dim=0)46 return processed_image47 48 def __call__(49 self,50 sample, timestep, encoder_hidden_states, conditionings,51 tiled=False, tile_size=64, tile_stride=32, **kwargs52 ):53 res_stack = None54 for processor, conditioning, model, scale in zip(self.processors, conditionings, self.models, self.scales):55 res_stack_ = model(56 sample, timestep, encoder_hidden_states, conditioning, **kwargs,57 tiled=tiled, tile_size=tile_size, tile_stride=tile_stride,58 processor_id=processor.processor_id59 )60 res_stack_ = [res * scale for res in res_stack_]61 if res_stack is None:62 res_stack = res_stack_63 else:64 res_stack = [i + j for i, j in zip(res_stack, res_stack_)]65 return res_stack66 67 68class FluxMultiControlNetManager(MultiControlNetManager):69 def __init__(self, controlnet_units=[]):70 super().__init__(controlnet_units=controlnet_units)71 72 def process_image(self, image, processor_id=None):73 if processor_id is None:74 processed_image = [processor(image) for processor in self.processors]75 else:76 processed_image = [self.processors[processor_id](image)]77 return processed_image78 79 def __call__(self, conditionings, **kwargs):80 res_stack, single_res_stack = None, None81 for processor, conditioning, model, scale in zip(self.processors, conditionings, self.models, self.scales):82 res_stack_, single_res_stack_ = model(controlnet_conditioning=conditioning, processor_id=processor.processor_id, **kwargs)83 res_stack_ = [res * scale for res in res_stack_]84 single_res_stack_ = [res * scale for res in single_res_stack_]85 if res_stack is None:86 res_stack = res_stack_87 single_res_stack = single_res_stack_88 else:89 res_stack = [i + j for i, j in zip(res_stack, res_stack_)]90 single_res_stack = [i + j for i, j in zip(single_res_stack, single_res_stack_)]91 return res_stack, single_res_stack92 