Team Ai
Apppublic

hugging-apps/echo-memory

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
controlnet_unit.py92 linesDownload Raw Back to controlnets
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