| 37 | |
| 38 | |
| 39 | class BasePrompter: |
| 40 | def __init__(self): |
| 41 | self.refiners = [] |
| 42 | self.extenders = [] |
| 43 | |
| 44 | |
| 45 | def load_prompt_refiners(self, model_manager: ModelManager, refiner_classes=[]): |
| 46 | for refiner_class in refiner_classes: |
| 47 | refiner = refiner_class.from_model_manager(model_manager) |
| 48 | self.refiners.append(refiner) |
| 49 | |
| 50 | def load_prompt_extenders(self,model_manager:ModelManager,extender_classes=[]): |
| 51 | for extender_class in extender_classes: |
| 52 | extender = extender_class.from_model_manager(model_manager) |
| 53 | self.extenders.append(extender) |
| 54 | |
| 55 | |
| 56 | @torch.no_grad() |
| 57 | def process_prompt(self, prompt, positive=True): |
| 58 | if isinstance(prompt, list): |
| 59 | prompt = [self.process_prompt(prompt_, positive=positive) for prompt_ in prompt] |
| 60 | else: |
| 61 | for refiner in self.refiners: |
| 62 | prompt = refiner(prompt, positive=positive) |
| 63 | return prompt |
| 64 | |
| 65 | @torch.no_grad() |
| 66 | def extend_prompt(self, prompt:str, positive=True): |
| 67 | extended_prompt = dict(prompt=prompt) |
| 68 | for extender in self.extenders: |
| 69 | extended_prompt = extender(extended_prompt) |
| 70 | return extended_prompt |
nothing calls this directly
no outgoing calls
no test coverage detected