modelscope/DiffSynth-Painter
14
1from ..models.model_manager import ModelManager2import torch3 4 5 6def tokenize_long_prompt(tokenizer, prompt, max_length=None):7 # Get model_max_length from self.tokenizer8 length = tokenizer.model_max_length if max_length is None else max_length9 10 # To avoid the warning. set self.tokenizer.model_max_length to +oo.11 tokenizer.model_max_length = 9999999912 13 # Tokenize it!14 input_ids = tokenizer(prompt, return_tensors="pt").input_ids15 16 # Determine the real length.17 max_length = (input_ids.shape[1] + length - 1) // length * length18 19 # Restore tokenizer.model_max_length20 tokenizer.model_max_length = length21 22 # Tokenize it again with fixed length.23 input_ids = tokenizer(24 prompt,25 return_tensors="pt",26 padding="max_length",27 max_length=max_length,28 truncation=True29 ).input_ids30 31 # Reshape input_ids to fit the text encoder.32 num_sentence = input_ids.shape[1] // length33 input_ids = input_ids.reshape((num_sentence, length))34 35 return input_ids36 37 38 39class BasePrompter:40 def __init__(self, refiners=[]):41 self.refiners = refiners42 43 44 def load_prompt_refiners(self, model_nameger: ModelManager, refiner_classes=[]):45 for refiner_class in refiner_classes:46 refiner = refiner_class.from_model_manager(model_nameger)47 self.refiners.append(refiner)48 49 50 @torch.no_grad()51 def process_prompt(self, prompt, positive=True):52 if isinstance(prompt, list):53 prompt = [self.process_prompt(prompt_, positive=positive) for prompt_ in prompt]54 else:55 for refiner in self.refiners:56 prompt = refiner(prompt, positive=positive)57 return prompt58 