From 1b73161e4d03933724c77faaf18e20a67d86d3a6 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Mon, 15 May 2023 12:25:41 -0400 Subject: [PATCH] update prompt parser --- modules/prompt_parser.py | 135 +++++++++++++------------------------- modules/sd_hijack_clip.py | 62 +---------------- modules/ui_extensions.py | 10 +-- scripts/xyz_grid.py | 2 +- test.py | 14 ++++ 5 files changed, 68 insertions(+), 155 deletions(-) create mode 100644 test.py diff --git a/modules/prompt_parser.py b/modules/prompt_parser.py index 7006f2822..7617ef23d 100644 --- a/modules/prompt_parser.py +++ b/modules/prompt_parser.py @@ -1,7 +1,10 @@ +# pylint: disable=anomalous-backslash-in-string import re from collections import namedtuple from typing import List import lark +import torch +from installer import log # a prompt like this: "fantasy landscape with a [mountain:lake:0.25] and [an oak:a christmas tree:0.75][ in foreground::0.6][ in background:0.25] [shoddy:masterful:0.5]" # will be represented with prompt_schedule like this (assuming steps=100): @@ -11,6 +14,11 @@ import lark # [75, 'fantasy landscape with a lake and an oak in background masterful'] # [100, 'fantasy landscape with a lake and a christmas tree in background masterful'] +round_bracket_multiplier = 1.1 +square_bracket_multiplier = 0.9 +re_AND = re.compile(r"\bAND\b") +re_weight = re.compile(r"^(.*?)(?:\s*:\s*([-+]?(?:\d+\.?|\d*\.\d+)))?\s*$") +ScheduledPromptConditioning = namedtuple("ScheduledPromptConditioning", ["end_at_step", "cond"]) schedule_parser = lark.Lark(r""" !start: (prompt | /[][():]/+)* prompt: (emphasized | scheduled | alternate | plain | WHITESPACE)* @@ -23,6 +31,16 @@ WHITESPACE: /\s+/ plain: /([^\\\[\]():|]|\\.)+/ %import common.SIGNED_NUMBER -> NUMBER """) +re_clean = re.compile(r"^\W+", re.S) +re_break = re.compile(r"\s*\bBREAK\b\s*", re.S) +re_attention = re.compile(r""" +\(|\[|\\\(|\\\[|\\|\\\\| +:([+-]?[.\d]+)| +\)|\]|\\\)|\\\]| +[^\(\)\[\]:]+| +: +""", re.X) + def get_learned_conditioning_prompt_schedules(prompts, steps): """ @@ -62,7 +80,7 @@ def get_learned_conditioning_prompt_schedules(prompts, steps): tree.children[-1] *= steps tree.children[-1] = min(steps, int(tree.children[-1])) l.append(tree.children[-1]) - def alternate(self, tree): + def alternate(self, tree): # pylint: disable=unused-argument l.extend(range(1, steps+1)) CollectSteps().visit(tree) return sorted(set(l)) @@ -92,7 +110,7 @@ def get_learned_conditioning_prompt_schedules(prompts, steps): def get_schedule(prompt): try: tree = schedule_parser.parse(prompt) - except lark.exceptions.LarkError as e: + except lark.exceptions.LarkError: return [[steps, prompt]] return [[t, at_step(t, tree)] for t in collect_steps(steps, tree)] @@ -100,82 +118,56 @@ def get_learned_conditioning_prompt_schedules(prompts, steps): return [promptdict[prompt] for prompt in prompts] -ScheduledPromptConditioning = namedtuple("ScheduledPromptConditioning", ["end_at_step", "cond"]) - - def get_learned_conditioning(model, prompts, steps): """converts a list of prompts into a list of prompt schedules - each schedule is a list of ScheduledPromptConditioning, specifying the comdition (cond), and the sampling step at which this condition is to be replaced by the next one. - Input: - (model, ['a red crown', 'a [blue:green:5] jeweled crown'], 20) - + (model, ['a red crown', 'a [blue:green:5] jeweled crown'], 20) Output: [ - [ - ScheduledPromptConditioning(end_at_step=20, cond=tensor([[-0.3886, 0.0229, -0.0523, ..., -0.4901, -0.3066, 0.0674], ..., [ 0.3317, -0.5102, -0.4066, ..., 0.4119, -0.7647, -1.0160]], device='cuda:0')) - ], - [ - ScheduledPromptConditioning(end_at_step=5, cond=tensor([[-0.3886, 0.0229, -0.0522, ..., -0.4901, -0.3067, 0.0673], ..., [-0.0192, 0.3867, -0.4644, ..., 0.1135, -0.3696, -0.4625]], device='cuda:0')), - ScheduledPromptConditioning(end_at_step=20, cond=tensor([[-0.3886, 0.0229, -0.0522, ..., -0.4901, -0.3067, 0.0673], ..., [-0.7352, -0.4356, -0.7888, ..., 0.6994, -0.4312, -1.2593]], device='cuda:0')) + [ ScheduledPromptConditioning(end_at_step=20, cond=tensor([[-0.3886, 0.0229, -0.0523, ..., -0.4901, -0.3066, 0.0674], ..., [ 0.3317, -0.5102, -0.4066, ..., 0.4119, -0.7647, -1.0160]], device='cuda:0')) ], + [ ScheduledPromptConditioning(end_at_step=5, cond=tensor([[-0.3886, 0.0229, -0.0522, ..., -0.4901, -0.3067, 0.0673], ..., [-0.0192, 0.3867, -0.4644, ..., 0.1135, -0.3696, -0.4625]], device='cuda:0')), + ScheduledPromptConditioning(end_at_step=20, cond=tensor([[-0.3886, 0.0229, -0.0522, ..., -0.4901, -0.3067, 0.0673], ..., [-0.7352, -0.4356, -0.7888, ..., 0.6994, -0.4312, -1.2593]], device='cuda:0')), ] ] """ res = [] - prompt_schedules = get_learned_conditioning_prompt_schedules(prompts, steps) cache = {} - for prompt, prompt_schedule in zip(prompts, prompt_schedules): - + log.debug(f'Prompt schedule: {prompt_schedule}') cached = cache.get(prompt, None) if cached is not None: res.append(cached) continue - texts = [x[1] for x in prompt_schedule] conds = model.get_learned_conditioning(texts) - cond_schedule = [] - for i, (end_at_step, text) in enumerate(prompt_schedule): + for i, (end_at_step, _text) in enumerate(prompt_schedule): cond_schedule.append(ScheduledPromptConditioning(end_at_step, conds[i])) - cache[prompt] = cond_schedule res.append(cond_schedule) - return res -re_AND = re.compile(r"\bAND\b") -re_weight = re.compile(r"^(.*?)(?:\s*:\s*([-+]?(?:\d+\.?|\d*\.\d+)))?\s*$") - def get_multicond_prompt_list(prompts): res_indexes = [] - prompt_flat_list = [] prompt_indexes = {} - for prompt in prompts: subprompts = re_AND.split(prompt) - indexes = [] for subprompt in subprompts: match = re_weight.search(subprompt) - text, weight = match.groups() if match is not None else (subprompt, 1.0) - weight = float(weight) if weight is not None else 1.0 - index = prompt_indexes.get(text, None) if index is None: index = len(prompt_flat_list) prompt_flat_list.append(text) prompt_indexes[text] = index - indexes.append((index, weight)) - res_indexes.append(indexes) - return res_indexes, prompt_flat_list, prompt_indexes @@ -190,21 +182,17 @@ class MulticondLearnedConditioning: self.shape: tuple = shape # the shape field is needed to send this object to DDIM/PLMS self.batch: List[List[ComposableScheduledPromptConditioning]] = batch + def get_multicond_learned_conditioning(model, prompts, steps) -> MulticondLearnedConditioning: """same as get_learned_conditioning, but returns a list of ScheduledPromptConditioning along with the weight objects for each prompt. For each prompt, the list is obtained by splitting the prompt using the AND separator. - https://energy-based-model.github.io/Compositional-Visual-Generation-with-Composable-Diffusion-Models/ """ - - res_indexes, prompt_flat_list, prompt_indexes = get_multicond_prompt_list(prompts) - + res_indexes, prompt_flat_list, _prompt_indexes = get_multicond_prompt_list(prompts) learned_conditioning = get_learned_conditioning(model, prompt_flat_list, steps) - res = [] for indexes in res_indexes: res.append([ComposableScheduledPromptConditioning(learned_conditioning[i], weight) for i, weight in indexes]) - return MulticondLearnedConditioning(shape=(len(prompts),), batch=res) @@ -213,66 +201,39 @@ def reconstruct_cond_batch(c: List[List[ScheduledPromptConditioning]], current_s res = torch.zeros((len(c),) + param.shape, device=param.device, dtype=param.dtype) for i, cond_schedule in enumerate(c): target_index = 0 - for current, (end_at, cond) in enumerate(cond_schedule): + for current, (end_at, _cond) in enumerate(cond_schedule): if current_step <= end_at: target_index = current break res[i] = cond_schedule[target_index].cond - return res def reconstruct_multicond_batch(c: MulticondLearnedConditioning, current_step): param = c.batch[0][0].schedules[0].cond - tensors = [] conds_list = [] - - for batch_no, composable_prompts in enumerate(c.batch): + for _batch_no, composable_prompts in enumerate(c.batch): conds_for_batch = [] - - for cond_index, composable_prompt in enumerate(composable_prompts): + for _cond_index, composable_prompt in enumerate(composable_prompts): target_index = 0 - for current, (end_at, cond) in enumerate(composable_prompt.schedules): + for current, (end_at, _cond) in enumerate(composable_prompt.schedules): if current_step <= end_at: target_index = current break - conds_for_batch.append((len(tensors), composable_prompt.weight)) tensors.append(composable_prompt.schedules[target_index].cond) - conds_list.append(conds_for_batch) - - # if prompts have wildly different lengths above the limit we'll get tensors fo different shapes - # and won't be able to torch.stack them. So this fixes that. + # if prompts have wildly different lengths above the limit we'll get tensors fo different shapes and won't be able to torch.stack them. So this fixes that. token_count = max([x.shape[0] for x in tensors]) for i in range(len(tensors)): if tensors[i].shape[0] != token_count: last_vector = tensors[i][-1:] last_vector_repeated = last_vector.repeat([token_count - tensors[i].shape[0], 1]) tensors[i] = torch.vstack([tensors[i], last_vector_repeated]) - return conds_list, torch.stack(tensors).to(device=param.device, dtype=param.dtype) -re_attention = re.compile(r""" -\\\(| -\\\)| -\\\[| -\\]| -\\\\| -\\| -\(| -\[| -:([+-]?[.\d]+)\)| -\)| -]| -[^\\()\[\]:]+| -: -""", re.X) - -re_break = re.compile(r"\s*\bBREAK\b\s*", re.S) - def parse_prompt_attention(text): """ Parses a string with attention tokens and returns a list of pairs: text and its associated weight. @@ -286,7 +247,6 @@ def parse_prompt_attention(text): \] - literal character ']' \\ - literal character '\' anything else - just text - >>> parse_prompt_attention('normal text') [['normal text', 1.0]] >>> parse_prompt_attention('an (important) word') @@ -308,22 +268,15 @@ def parse_prompt_attention(text): ['sky', 1.4641000000000006], ['.', 1.1]] """ - res = [] round_brackets = [] square_brackets = [] - - round_bracket_multiplier = 1.1 - square_bracket_multiplier = 1 / 1.1 - def multiply_range(start_position, multiplier): for p in range(start_position, len(res)): res[p][1] *= multiplier - for m in re_attention.finditer(text): text = m.group(0) weight = m.group(1) - if text.startswith('\\'): res.append([text[1:], 1.0]) elif text == '(': @@ -332,6 +285,8 @@ def parse_prompt_attention(text): square_brackets.append(len(res)) elif weight is not None and len(round_brackets) > 0: multiply_range(round_brackets.pop(), float(weight)) + elif weight is not None and len(square_brackets) > 0: + multiply_range(square_brackets.pop(), float(weight)) elif text == ')' and len(round_brackets) > 0: multiply_range(round_brackets.pop(), round_bracket_multiplier) elif text == ']' and len(square_brackets) > 0: @@ -339,19 +294,18 @@ def parse_prompt_attention(text): else: parts = re.split(re_break, text) for i, part in enumerate(parts): + part = re.sub(re_clean, "", part) + if len(part) == 0: + continue if i > 0: res.append(["BREAK", -1]) res.append([part, 1.0]) - for pos in round_brackets: multiply_range(pos, round_bracket_multiplier) - for pos in square_brackets: multiply_range(pos, square_bracket_multiplier) - if len(res) == 0: res = [["", 1.0]] - # merge runs of identical weights i = 0 while i + 1 < len(res): @@ -360,11 +314,14 @@ def parse_prompt_attention(text): res.pop(i + 1) else: i += 1 - + log.debug(f'Prompt parse-attention: {res}') return res if __name__ == "__main__": - import doctest - doctest.testmod(optionflags=doctest.NORMALIZE_WHITESPACE) -else: - import torch # doctest faster + # import os + # import sys + # sys.path.append(os.path.join(os.path.dirname(__file__), '..')) + input_text = "(upzero) (upone:1.1), ((uptwo:1.2)), [downzero], [downone:0.9], [[downtwo:0.8]], this is a test" + output_list = parse_prompt_attention(input_text) + print('INPUT', input_text) + print('OUTPUT', output_list) diff --git a/modules/sd_hijack_clip.py b/modules/sd_hijack_clip.py index 945f7732d..fe59f976c 100644 --- a/modules/sd_hijack_clip.py +++ b/modules/sd_hijack_clip.py @@ -1,8 +1,6 @@ import math from collections import namedtuple - import torch - from modules import prompt_parser, devices, sd_hijack from modules.shared import opts @@ -14,7 +12,6 @@ class PromptChunk: Each PromptChunk contains an exact amount of tokens - 77, which includes one for start and end token, so just 75 tokens from prompt. """ - def __init__(self): self.tokens = [] self.multipliers = [] @@ -31,20 +28,16 @@ class FrozenCLIPEmbedderWithCustomWordsBase(torch.nn.Module): """A pytorch module that is a wrapper for FrozenCLIPEmbedder module. it enhances FrozenCLIPEmbedder, making it possible to have unlimited prompt length and assign weights to tokens in prompt. """ - def __init__(self, wrapped, hijack): super().__init__() - self.wrapped = wrapped """Original FrozenCLIPEmbedder module; can also be FrozenOpenCLIPEmbedder or xlmr.BertSeriesModelWithTransformation, depending on model.""" - self.hijack: sd_hijack.StableDiffusionModelHijack = hijack self.chunk_length = 75 def empty_chunk(self): """creates an empty PromptChunk and returns it""" - chunk = PromptChunk() chunk.tokens = [self.id_start] + [self.id_end] * (self.chunk_length + 1) chunk.multipliers = [1.0] * (self.chunk_length + 2) @@ -52,12 +45,10 @@ class FrozenCLIPEmbedderWithCustomWordsBase(torch.nn.Module): def get_target_prompt_token_count(self, token_count): """returns the maximum number of tokens a prompt of a known length can have before it requires one more PromptChunk to be represented""" - return math.ceil(max(token_count, 1) / self.chunk_length) * self.chunk_length def tokenize(self, texts): """Converts a batch of texts into a batch of token ids""" - raise NotImplementedError def encode_with_transformers(self, tokens): @@ -68,13 +59,11 @@ class FrozenCLIPEmbedderWithCustomWordsBase(torch.nn.Module): model - can be 768 and 1024. Among other things, this call will read self.hijack.fixes, apply it to its inputs, and clear it (setting it to None). """ - raise NotImplementedError def encode_embedding_init_text(self, init_text, nvpt): """Converts text into a tensor with this text's tokens' embeddings. Note that those are embeddings before they are passed through transformers. nvpt is used as a maximum length in tokens. If text produces less teokens than nvpt, only this many is returned.""" - raise NotImplementedError def tokenize_line(self, line): @@ -83,14 +72,8 @@ class FrozenCLIPEmbedderWithCustomWordsBase(torch.nn.Module): represent the prompt. Returns the list and the total number of tokens in the prompt. """ - - if opts.enable_emphasis: - parsed = prompt_parser.parse_prompt_attention(line) - else: - parsed = [[line, 1.0]] - + parsed = prompt_parser.parse_prompt_attention(line) tokenized = self.tokenize([text for text, _ in parsed]) - chunks = [] chunk = PromptChunk() token_count = 0 @@ -102,20 +85,16 @@ class FrozenCLIPEmbedderWithCustomWordsBase(torch.nn.Module): nonlocal token_count nonlocal last_comma nonlocal chunk - if is_last: token_count += len(chunk.tokens) else: token_count += self.chunk_length - to_add = self.chunk_length - len(chunk.tokens) if to_add > 0: chunk.tokens += [self.id_end] * to_add chunk.multipliers += [1.0] * to_add - chunk.tokens = [self.id_start] + chunk.tokens + [self.id_end] chunk.multipliers = [1.0] + chunk.multipliers + [1.0] - last_comma = -1 chunks.append(chunk) chunk = PromptChunk() @@ -124,52 +103,40 @@ class FrozenCLIPEmbedderWithCustomWordsBase(torch.nn.Module): if text == 'BREAK' and weight == -1: next_chunk() continue - position = 0 while position < len(tokens): token = tokens[position] - if token == self.comma_token: last_comma = len(chunk.tokens) - # this is when we are at the end of alloted 75 tokens for the current chunk, and the current token is not a comma. opts.comma_padding_backtrack # is a setting that specifies that if there is a comma nearby, the text after the comma should be moved out of this chunk and into the next. elif opts.comma_padding_backtrack != 0 and len(chunk.tokens) == self.chunk_length and last_comma != -1 and len(chunk.tokens) - last_comma <= opts.comma_padding_backtrack: break_location = last_comma + 1 - reloc_tokens = chunk.tokens[break_location:] reloc_mults = chunk.multipliers[break_location:] - chunk.tokens = chunk.tokens[:break_location] chunk.multipliers = chunk.multipliers[:break_location] - next_chunk() chunk.tokens = reloc_tokens chunk.multipliers = reloc_mults - if len(chunk.tokens) == self.chunk_length: next_chunk() - embedding, embedding_length_in_tokens = self.hijack.embedding_db.find_embedding_at_position(tokens, position) if embedding is None: chunk.tokens.append(token) chunk.multipliers.append(weight) position += 1 continue - emb_len = int(embedding.vec.shape[0]) if len(chunk.tokens) + emb_len > self.chunk_length: next_chunk() - chunk.fixes.append(PromptChunkFix(len(chunk.tokens), embedding)) - chunk.tokens += [0] * emb_len chunk.multipliers += [weight] * emb_len position += embedding_length_in_tokens - if len(chunk.tokens) > 0 or len(chunks) == 0: next_chunk(is_last=True) - + # print('CHUNKS', [vars(c) for c in chunks]) # TODO return chunks, token_count def process_texts(self, texts): @@ -177,9 +144,7 @@ class FrozenCLIPEmbedderWithCustomWordsBase(torch.nn.Module): Accepts a list of texts and calls tokenize_line() on each, with cache. Returns the list of results and maximum length, in tokens, of all texts. """ - token_count = 0 - cache = {} batch_chunks = [] for line in texts: @@ -188,11 +153,8 @@ class FrozenCLIPEmbedderWithCustomWordsBase(torch.nn.Module): else: chunks, current_token_count = self.tokenize_line(line) token_count = max(current_token_count, token_count) - cache[line] = chunks - batch_chunks.append(chunks) - return batch_chunks, token_count def forward(self, texts): @@ -204,31 +166,23 @@ class FrozenCLIPEmbedderWithCustomWordsBase(torch.nn.Module): Webui usually sends just one text at a time through this function - the only time when texts is an array with more than one elemenet is when you do prompt editing: "a picture of a [cat:dog:0.4] eating ice cream" """ - batch_chunks, _token_count = self.process_texts(texts) - used_embeddings = {} chunk_count = max([len(x) for x in batch_chunks]) - zs = [] for i in range(chunk_count): batch_chunk = [chunks[i] if i < len(chunks) else self.empty_chunk() for chunks in batch_chunks] - tokens = [x.tokens for x in batch_chunk] multipliers = [x.multipliers for x in batch_chunk] self.hijack.fixes = [x.fixes for x in batch_chunk] - for fixes in self.hijack.fixes: for _position, embedding in fixes: used_embeddings[embedding.name] = embedding - z = self.process_tokens(tokens, multipliers) zs.append(z) - if len(used_embeddings) > 0: embeddings_list = ", ".join([f'{name} [{embedding.checksum()}]' for name, embedding in used_embeddings.items()]) self.hijack.comments.append(f"Used embeddings: {embeddings_list}") - return torch.hstack(zs) def process_tokens(self, remade_batch_tokens, batch_multipliers): @@ -240,7 +194,6 @@ class FrozenCLIPEmbedderWithCustomWordsBase(torch.nn.Module): corresponds to one token. """ tokens = torch.asarray(remade_batch_tokens).to(devices.device) - # this is for SD2: SD1 uses the same token for padding and end of text, while SD2 uses different ones. if self.id_end != self.id_pad: for batch_pos in range(len(remade_batch_tokens)): @@ -248,14 +201,12 @@ class FrozenCLIPEmbedderWithCustomWordsBase(torch.nn.Module): tokens[batch_pos, index+1:tokens.shape[1]] = self.id_pad z = self.encode_with_transformers(tokens) - # restoring original mean is likely not correct, but it seems to work well to prevent artifacts that happen otherwise batch_multipliers = torch.asarray(batch_multipliers).to(devices.device) original_mean = z.mean() z = z * batch_multipliers.reshape(batch_multipliers.shape + (1,)).expand(z.shape) new_mean = z.mean() z = z * (original_mean / new_mean) - return z @@ -263,11 +214,8 @@ class FrozenCLIPEmbedderWithCustomWords(FrozenCLIPEmbedderWithCustomWordsBase): def __init__(self, wrapped, hijack): super().__init__(wrapped, hijack) self.tokenizer = wrapped.tokenizer - vocab = self.tokenizer.get_vocab() - self.comma_token = vocab.get(',', None) - self.token_mults = {} tokens_with_parens = [(k, v) for k, v in vocab.items() if '(' in k or ')' in k or '[' in k or ']' in k] for text, ident in tokens_with_parens: @@ -281,35 +229,29 @@ class FrozenCLIPEmbedderWithCustomWords(FrozenCLIPEmbedderWithCustomWordsBase): mult *= 1.1 if c == ')': mult /= 1.1 - if mult != 1.0: self.token_mults[ident] = mult - self.id_start = self.wrapped.tokenizer.bos_token_id self.id_end = self.wrapped.tokenizer.eos_token_id self.id_pad = self.id_end def tokenize(self, texts): tokenized = self.wrapped.tokenizer(texts, truncation=False, add_special_tokens=False)["input_ids"] - return tokenized def encode_with_transformers(self, tokens): if opts.CLIP_stop_at_last_layers is None: opts.CLIP_stop_at_last_layers = 1 outputs = self.wrapped.transformer(input_ids=tokens, output_hidden_states=-opts.CLIP_stop_at_last_layers) - if opts.CLIP_stop_at_last_layers > 1: z = outputs.hidden_states[-opts.CLIP_stop_at_last_layers] z = self.wrapped.transformer.text_model.final_layer_norm(z) else: z = outputs.last_hidden_state - return z def encode_embedding_init_text(self, init_text, nvpt): embedding_layer = self.wrapped.transformer.text_model.embeddings ids = self.wrapped.tokenizer(init_text, max_length=nvpt, return_tensors="pt", add_special_tokens=False)["input_ids"] embedded = embedding_layer.token_embedding.wrapped(ids.to(embedding_layer.token_embedding.wrapped.weight.device)).squeeze(0) - return embedded diff --git a/modules/ui_extensions.py b/modules/ui_extensions.py index 916803576..fccc55693 100644 --- a/modules/ui_extensions.py +++ b/modules/ui_extensions.py @@ -91,10 +91,10 @@ def check_updates(_id_task, disable_list, search_text, sort_column): if ext.can_update: ext.fetch_and_reset_hard() ext.read_info_from_repo() - commit_date = ext.get('commit_date', 1577836800) or 1577836800 + commit_date = ext.commit_date or 1577836800 shared.log.info(f'Extensions updated: {ext.name} {ext.commit_hash[:8]} {datetime.utcfromtimestamp(commit_date)}') else: - commit_date = ext.get('commit_date', 1577836800) or 1577836800 + commit_date = ext.commit_date or 1577836800 shared.log.debug(f'Extensions no update available: {ext.name} {ext.commit_hash[:8]} {datetime.utcfromtimestamp(commit_date)}') except FileNotFoundError as e: if 'FETCH_HEAD' not in str(e): @@ -195,10 +195,10 @@ def update_extension(extension_path, search_text, sort_column): if ext.can_update: ext.fetch_and_reset_hard() ext.read_info_from_repo() - commit_date = ext.get('commit_date', 1577836800) or 1577836800 + commit_date = ext.commit_date or 1577836800 shared.log.info(f'Extensions updated: {ext.name} {ext.commit_hash[:8]} {datetime.utcfromtimestamp(commit_date)}') else: - commit_date = ext.get('commit_date', 1577836800) or 1577836800 + commit_date = ext.commit_date or 1577836800 shared.log.info(f'Extensions no update available: {ext.name} {ext.commit_hash[:8]} {datetime.utcfromtimestamp(commit_date)}') except FileNotFoundError as e: if 'FETCH_HEAD' not in str(e): @@ -306,7 +306,7 @@ def refresh_extensions_list_from_data(search_text, sort_column): enabled = ext.get("enabled", False) path = ext.get("path", "") remote = ext.get("remote", None) - commit_date = ext.get('commit_date', 1577836800) or 1577836800 + commit_date = ext.get("commit_date", 1577836800) or 1577836800 update_available = (remote is not None) & (installed) & (datetime.utcfromtimestamp(commit_date + 60 * 60) < datetime.fromisoformat(ext.get('updated', '2000-01-01T00:00:00.000Z')[:-1])) tags = ext.get("tags", []) tags_string = ' '.join(tags) diff --git a/scripts/xyz_grid.py b/scripts/xyz_grid.py index 0fda6b4c5..64ccab852 100644 --- a/scripts/xyz_grid.py +++ b/scripts/xyz_grid.py @@ -350,7 +350,7 @@ class SharedSettingsStackHelper(object): self.token_merging_random = shared.opts.token_merging_random self.sd_model_checkpoint = shared.opts.sd_model_checkpoint self.sd_vae_checkpoint = shared.opts.sd_vae - self.xyz_fallback_sampler = shared.opts.data["xyz_fallback_sampler"] + self.xyz_fallback_sampler = shared.opts.xyz_fallback_sampler def __exit__(self, exc_type, exc_value, tb): #Restore overriden settings after plot generation. diff --git a/test.py b/test.py new file mode 100644 index 000000000..cd30cba8c --- /dev/null +++ b/test.py @@ -0,0 +1,14 @@ +import re + +re_attention = re.compile(r""" +\(|\[|\\\(|\\\[|\\|\\\\| +:([+-]?[.\d]+)| +\)|\]|\\\)|\\\]| +[^\(\)\[\]:]+| +: +""", re.X) + +texts = ["car:2.0", "(car:1.1)", "((car:1.2))", "[car:0.9]", "[[car:0.8]]"] +for text in texts: + for m in re_attention.finditer(text): + print(text, '0:', m.group(0), '1:', m.group(1))