Team Ai
Apppublic

modelscope/DiffSynth-Painter

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