From cdde0e30b2d6df77df52c366dd96eb309703a2ec Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Wed, 2 Aug 2023 09:54:56 -0500 Subject: [PATCH] fixed control net weights not initializing when no weights passed in, removed commented-out code --- control.py | 9 +++------ 1 file changed, 3 insertions(+), 6 deletions(-) diff --git a/control.py b/control.py index 93572f3..b4e2cff 100644 --- a/control.py +++ b/control.py @@ -28,7 +28,7 @@ class ControlNetAdvanced(ControlBase): def __init__(self, control_model, weights: ControlNetWeightsType, global_average_pooling=False, device=None): super().__init__(device) self.control_model = control_model - self.weights = weights + self.weights = weights if weights else [1.0]*13 self.global_average_pooling = global_average_pooling def get_control(self, x_noisy, t, cond, batched_number): @@ -77,9 +77,7 @@ class ControlNetAdvanced(ControlBase): if self.global_average_pooling: x = torch.mean(x, dim=(2, 3), keepdim=True).repeat(1, 1, x.shape[2], x.shape[3]) - #multiplier = 1#0.825**float(12-i) - #print(f"$$$ multiplier: {multiplier}") - x *= self.strength*self.weights[i] + x *= self.strength * self.weights[i] # apply layer weight if x.dtype != output_dtype and not autocast_enabled: x = x.to(output_dtype) @@ -248,10 +246,9 @@ class T2IAdapterAdvanced(ControlBase): out = {'input':[]} autocast_enabled = torch.is_autocast_enabled() - #print(f"$$$$ t2i control_input len: {len(self.control_input)}") for i in range(len(self.control_input)): key = 'input' - x = self.control_input[i] * self.strength * self.weights[i] + x = self.control_input[i] * self.strength * self.weights[i] # apply layer weight if x.dtype != output_dtype and not autocast_enabled: x = x.to(output_dtype)