MCPcopy Create free account
hub / github.com/OpenImagingLab/FlashVSR / BasePrompter

Class BasePrompter

diffsynth/prompters/base_prompter.py:39–70  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

37
38
39class 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

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected