| 5 | |
| 6 | |
| 7 | class Planner(LLMNode): |
| 8 | def __init__(self, workers, prefix=DEFAULT_PREFIX, suffix=DEFAULT_SUFFIX, fewshot=DEFAULT_FEWSHOT, |
| 9 | model_name="text-davinci-003", stop=None): |
| 10 | super().__init__("Planner", model_name, stop, input_type=str, output_type=str) |
| 11 | self.workers = workers |
| 12 | self.prefix = prefix |
| 13 | self.worker_prompt = self._generate_worker_prompt() |
| 14 | self.suffix = suffix |
| 15 | self.fewshot = fewshot |
| 16 | |
| 17 | def run(self, input, log=False): |
| 18 | assert isinstance(input, self.input_type) |
| 19 | prompt = self.prefix + self.worker_prompt + self.fewshot + self.suffix + input + '\n' |
| 20 | if self.model_name in LLAMA_WEIGHTS: |
| 21 | prompt = [self.prefix + self.worker_prompt, input] |
| 22 | response = self.call_llm(prompt, self.stop) |
| 23 | completion = response["output"] |
| 24 | if log: |
| 25 | return response |
| 26 | return completion |
| 27 | |
| 28 | def _get_worker(self, name): |
| 29 | if name in WORKER_REGISTRY: |
| 30 | return WORKER_REGISTRY[name] |
| 31 | else: |
| 32 | raise ValueError("Worker not found") |
| 33 | |
| 34 | def _generate_worker_prompt(self): |
| 35 | prompt = "Tools can be one of the following:\n" |
| 36 | for name in self.workers: |
| 37 | worker = self._get_worker(name) |
| 38 | prompt += f"{worker.name}[input]: {worker.description}\n" |
| 39 | return prompt + "\n" |