diff --git a/modules/hypernetworks/hypernetwork.py b/modules/hypernetworks/hypernetwork.py index 85c58cbb4..351e261ea 100644 --- a/modules/hypernetworks/hypernetwork.py +++ b/modules/hypernetworks/hypernetwork.py @@ -1,5 +1,4 @@ import datetime -import glob import html import os from collections import deque @@ -35,19 +34,14 @@ class HypernetworkModule(torch.nn.Module): def __init__(self, dim, state_dict=None, layer_structure=None, activation_func=None, weight_init='Normal', add_layer_norm=False, activate_output=False, dropout_structure=None): super().__init__() - self.multiplier = 1.0 - assert layer_structure is not None, "layer_structure must not be None" assert layer_structure[0] == 1, "Multiplier Sequence should start with size 1!" assert layer_structure[-1] == 1, "Multiplier Sequence should end with size 1!" - linears = [] for i in range(len(layer_structure) - 1): - # Add a fully-connected layer linears.append(torch.nn.Linear(int(dim * layer_structure[i]), int(dim * layer_structure[i+1]))) - # Add an activation func except last layer if activation_func == "linear" or activation_func is None or (i >= len(layer_structure) - 2 and not activate_output): pass @@ -55,20 +49,16 @@ class HypernetworkModule(torch.nn.Module): linears.append(self.activation_dict[activation_func]()) else: raise RuntimeError(f'hypernetwork uses an unsupported activation function: {activation_func}') - # Add layer normalization if add_layer_norm: linears.append(torch.nn.LayerNorm(int(dim * layer_structure[i+1]))) - # Everything should be now parsed into dropout structure, and applied here. # Since we only have dropouts after layers, dropout structure should start with 0 and end with 0. if dropout_structure is not None and dropout_structure[i+1] > 0: assert 0 < dropout_structure[i+1] < 1, "Dropout probability should be 0 or float between 0 and 1!" linears.append(torch.nn.Dropout(p=dropout_structure[i+1])) # Code explanation : [1, 2, 1] -> dropout is missing when last_layer_dropout is false. [1, 2, 2, 1] -> [0, 0.3, 0, 0], when its True, [0, 0.3, 0.3, 0]. - self.linear = torch.nn.Sequential(*linears) - if state_dict is not None: self.fix_old_state_dict(state_dict) self.load_state_dict(state_dict) @@ -102,12 +92,10 @@ class HypernetworkModule(torch.nn.Module): 'linear2.bias': 'linear.1.bias', 'linear2.weight': 'linear.1.weight', } - for fr, to in changes.items(): x = state_dict.get(fr, None) if x is None: continue - del state_dict[fr] state_dict[to] = x @@ -162,7 +150,6 @@ class Hypernetwork: self.optimizer_name = None self.optimizer_state_dict = None self.optional_info = None - for size in enable_sizes or []: self.layers[size] = ( HypernetworkModule(size, None, self.layer_structure, self.activation_func, self.weight_init, @@ -210,10 +197,8 @@ class Hypernetwork: def save(self, filename): state_dict = {} optimizer_saved_dict = {} - for k, v in self.layers.items(): state_dict[k] = (v[0].state_dict(), v[1].state_dict()) - state_dict['step'] = self.step state_dict['name'] = self.name state_dict['layer_structure'] = self.layer_structure @@ -227,10 +212,8 @@ class Hypernetwork: state_dict['dropout_structure'] = self.dropout_structure state_dict['last_layer_dropout'] = (self.dropout_structure[-2] != 0) if self.dropout_structure is not None else self.last_layer_dropout state_dict['optional_info'] = self.optional_info if self.optional_info else None - if self.optimizer_name is not None: optimizer_saved_dict['optimizer_name'] = self.optimizer_name - torch.save(state_dict, filename) if shared.opts.save_optimizer_state and self.optimizer_state_dict: optimizer_saved_dict['hash'] = self.shorthash() @@ -241,10 +224,8 @@ class Hypernetwork: self.filename = filename if self.name is None: self.name = os.path.splitext(os.path.basename(filename))[0] - with progress.open(filename, 'rb', description=f'Loading hypernetwork: [cyan]{filename}', auto_refresh=True) as f: state_dict = torch.load(f, map_location='cpu') - self.layer_structure = state_dict.get('layer_structure', [1, 2, 1]) self.optional_info = state_dict.get('optional_info', None) self.activation_func = state_dict.get('activation_func', None) @@ -257,11 +238,9 @@ class Hypernetwork: # Dropout structure should have same length as layer structure, Every digits should be in [0,1), and last digit must be 0. if self.dropout_structure is None: self.dropout_structure = parse_dropout_structure(self.layer_structure, self.use_dropout, self.last_layer_dropout) - if shared.opts.print_hypernet_extra: if self.optional_info is not None: print(f" INFO:\n {self.optional_info}\n") - print(f" Layer structure: {self.layer_structure}") print(f" Activation function: {self.activation_func}") print(f" Weight initialization: {self.weight_init}") @@ -269,9 +248,7 @@ class Hypernetwork: print(f" Dropout usage: {self.use_dropout}" ) print(f" Activate last layer: {self.activate_output}") print(f" Dropout structure: {self.dropout_structure}") - optimizer_saved_dict = torch.load(self.filename + '.optim', map_location='cpu') if os.path.exists(self.filename + '.optim') else {} - if self.shorthash() == optimizer_saved_dict.get('hash', None): self.optimizer_state_dict = optimizer_saved_dict.get('optimizer_state_dict', None) else: @@ -285,7 +262,6 @@ class Hypernetwork: self.optimizer_name = "AdamW" if shared.opts.print_hypernet_extra: print("No saved optimizer exists in checkpoint") - for size, sd in state_dict.items(): if type(size) == int: self.layers[size] = ( @@ -294,7 +270,6 @@ class Hypernetwork: HypernetworkModule(size, sd[1], self.layer_structure, self.activation_func, self.weight_init, self.add_layer_norm, self.activate_output, self.dropout_structure), ) - self.name = state_dict.get('name', self.name) self.step = state_dict.get('step', 0) self.sd_checkpoint = state_dict.get('sd_checkpoint', None) @@ -303,54 +278,49 @@ class Hypernetwork: def shorthash(self): sha256 = hashes.sha256(self.filename, f'hypernet/{self.name}') - return sha256[0:10] if sha256 else None def list_hypernetworks(path): res = {} - for filename in sorted(glob.iglob(os.path.join(path, '**/*.pt'), recursive=True), key=str.lower): - name = os.path.splitext(os.path.basename(filename))[0] - # Prevent a hypothetical "None.pt" from being listed. - if name != "None": - res[name] = filename + def list_folder(folder): + for filename in os.listdir(folder): + fn = os.path.join(folder, filename) + if os.path.isfile(fn) and fn.lower().endswith(".pt"): + name = os.path.splitext(os.path.basename(fn))[0] + res[name] = filename + elif os.path.isdir(fn) and not fn.startswith('.'): + list_folder(fn) + + list_folder(path) return res def load_hypernetwork(name): path = shared.hypernetworks.get(name, None) - if path is None: return None - hypernetwork = Hypernetwork() - try: hypernetwork.load(path) except Exception as e: errors.display(e, f'hypernetwork load: {path}') return None - return hypernetwork def load_hypernetworks(names, multipliers=None): already_loaded = {} - for hypernetwork in shared.loaded_hypernetworks: if hypernetwork.name in names: already_loaded[hypernetwork.name] = hypernetwork - shared.loaded_hypernetworks.clear() - for i, name in enumerate(names): hypernetwork = already_loaded.get(name, None) if hypernetwork is None: hypernetwork = load_hypernetwork(name) - if hypernetwork is None: continue - hypernetwork.set_multiplier(multipliers[i] if multipliers else 1.0) shared.loaded_hypernetworks.append(hypernetwork) @@ -368,14 +338,11 @@ def find_closest_hypernetwork_name(search: str): def apply_single_hypernetwork(hypernetwork, context_k, context_v, layer=None): hypernetwork_layers = (hypernetwork.layers if hypernetwork is not None else {}).get(context_k.shape[2], None) - if hypernetwork_layers is None: return context_k, context_v - if layer is not None: layer.hyper_k = hypernetwork_layers[0] layer.hyper_v = hypernetwork_layers[1] - context_k = devices.cond_cast_unet(hypernetwork_layers[0](devices.cond_cast_float(context_k))) context_v = devices.cond_cast_unet(hypernetwork_layers[1](devices.cond_cast_float(context_v))) return context_k, context_v @@ -386,33 +353,25 @@ def apply_hypernetworks(hypernetworks, context, layer=None): context_v = context for hypernetwork in hypernetworks: context_k, context_v = apply_single_hypernetwork(hypernetwork, context_k, context_v, layer) - return context_k, context_v def attention_CrossAttention_forward(self, x, context=None, mask=None): h = self.heads - q = self.to_q(x) context = default(context, x) - context_k, context_v = apply_hypernetworks(shared.loaded_hypernetworks, context, self) k = self.to_k(context_k) v = self.to_v(context_v) - q, k, v = (rearrange(t, 'b n (h d) -> (b h) n d', h=h) for t in (q, k, v)) - sim = einsum('b i d, b j d -> b i j', q, k) * self.scale - if mask is not None: mask = rearrange(mask, 'b ... -> b (...)') max_neg_value = -torch.finfo(sim.dtype).max mask = repeat(mask, 'b j -> (b h) () j', h=h) sim.masked_fill_(~mask, max_neg_value) - # attention, what we cannot get enough of attn = sim.softmax(dim=-1) - out = einsum('b i j, b j d -> b i d', attn, v) out = rearrange(out, '(b h) n d -> b n (h d)', h=h) return self.to_out(out) @@ -421,7 +380,6 @@ def attention_CrossAttention_forward(self, x, context=None, mask=None): def stack_conds(conds): if len(conds) == 1: return torch.stack(conds) - # same as in reconstruct_multicond_batch token_count = max([x.shape[0] for x in conds]) for i in range(len(conds)): @@ -429,7 +387,6 @@ def stack_conds(conds): last_vector = conds[i][-1:] last_vector_repeated = last_vector.repeat([token_count - conds[i].shape[0], 1]) conds[i] = torch.vstack([conds[i], last_vector_repeated]) - return torch.stack(conds) @@ -464,19 +421,15 @@ def create_hypernetwork(name, enable_sizes, overwrite_old, layer_structure=None, # Remove illegal characters from name. name = "".join( x for x in name if (x.isalnum() or x in "._- ")) assert name, "Name cannot be empty!" - fn = os.path.join(shared.opts.hypernetwork_dir, f"{name}.pt") if not overwrite_old: assert not os.path.exists(fn), f"file {fn} already exists" - if type(layer_structure) == str: layer_structure = [float(x.strip()) for x in layer_structure.split(",")] - if use_dropout and dropout_structure and type(dropout_structure) == str: dropout_structure = [float(x.strip()) for x in dropout_structure.split(",")] else: dropout_structure = [0] * len(layer_structure) - hypernet = modules.hypernetworks.hypernetwork.Hypernetwork( name=name, enable_sizes=[int(x) for x in enable_sizes],