modelscope/DiffSynth-Painter
14
1from transformers import AutoTokenizer2from ..models.model_manager import ModelManager3import torch4 5 6 7class BeautifulPrompt(torch.nn.Module):8 def __init__(self, tokenizer_path=None, model=None, template=""):9 super().__init__()10 self.tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)11 self.model = model12 self.template = template13 14 15 @staticmethod16 def from_model_manager(model_nameger: ModelManager):17 model, model_path = model_nameger.fetch_model("beautiful_prompt", require_model_path=True)18 template = 'Instruction: Give a simple description of the image to generate a drawing prompt.\nInput: {raw_prompt}\nOutput:'19 if model_path.endswith("v2"):20 template = """Converts a simple image description into a prompt. \21Prompts are formatted as multiple related tags separated by commas, plus you can use () to increase the weight, [] to decrease the weight, \22or use a number to specify the weight. You should add appropriate words to make the images described in the prompt more aesthetically pleasing, \23but make sure there is a correlation between the input and output.\n\24### Input: {raw_prompt}\n### Output:"""25 beautiful_prompt = BeautifulPrompt(26 tokenizer_path=model_path,27 model=model,28 template=template29 )30 return beautiful_prompt31 32 33 def __call__(self, raw_prompt, positive=True, **kwargs):34 if positive:35 model_input = self.template.format(raw_prompt=raw_prompt)36 input_ids = self.tokenizer.encode(model_input, return_tensors='pt').to(self.model.device)37 outputs = self.model.generate(38 input_ids,39 max_new_tokens=384,40 do_sample=True,41 temperature=0.9,42 top_k=50,43 top_p=0.95,44 repetition_penalty=1.1,45 num_return_sequences=146 )47 prompt = raw_prompt + ", " + self.tokenizer.batch_decode(48 outputs[:, input_ids.size(1):],49 skip_special_tokens=True50 )[0].strip()51 print(f"Your prompt is refined by BeautifulPrompt: {prompt}")52 return prompt53 else:54 return raw_prompt55 56 57 58class Translator(torch.nn.Module):59 def __init__(self, tokenizer_path=None, model=None):60 super().__init__()61 self.tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)62 self.model = model63 64 65 @staticmethod66 def from_model_manager(model_nameger: ModelManager):67 model, model_path = model_nameger.fetch_model("translator", require_model_path=True)68 translator = Translator(tokenizer_path=model_path, model=model)69 return translator70 71 72 def __call__(self, prompt, **kwargs):73 input_ids = self.tokenizer.encode(prompt, return_tensors='pt').to(self.model.device)74 output_ids = self.model.generate(input_ids)75 prompt = self.tokenizer.batch_decode(output_ids, skip_special_tokens=True)[0]76 print(f"Your prompt is translated: {prompt}")77 return prompt78 