Team Ai
Apppublic

modelscope/DiffSynth-Painter

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
14likes
prompt_refiners.py78 linesDownload Raw Back to prompters
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