| 14 | from .base import EngineLM, CachedEngine |
| 15 | |
| 16 | class ChatCohere(EngineLM, CachedEngine): |
| 17 | DEFAULT_SYSTEM_PROMPT = "You are a helpful, creative, and smart assistant." |
| 18 | |
| 19 | def __init__( |
| 20 | self, |
| 21 | model_string="command-r-plus", |
| 22 | system_prompt=DEFAULT_SYSTEM_PROMPT): |
| 23 | """ |
| 24 | :param model_string: |
| 25 | :param system_prompt: |
| 26 | """ |
| 27 | root = platformdirs.user_cache_dir("textgrad") |
| 28 | cache_path = os.path.join(root, f"cache_cohere_{model_string}.db") |
| 29 | super().__init__(cache_path=cache_path) |
| 30 | |
| 31 | self.system_prompt = system_prompt |
| 32 | if os.getenv("COHERE_API_KEY") is None: |
| 33 | raise ValueError("Please set the COHERE_API_KEY environment variable if you'd like to use Cohere models.") |
| 34 | |
| 35 | self.client = cohere.Client( |
| 36 | api_key=os.getenv("COHERE_API_KEY"), |
| 37 | ) |
| 38 | self.model_string = model_string |
| 39 | |
| 40 | def generate( |
| 41 | self, prompt, system_prompt=None, temperature=0, max_tokens=2000, top_p=0.99 |
| 42 | ): |
| 43 | |
| 44 | sys_prompt_arg = system_prompt if system_prompt else self.system_prompt |
| 45 | |
| 46 | cache_or_none = self._check_cache(sys_prompt_arg + prompt) |
| 47 | if cache_or_none is not None: |
| 48 | return cache_or_none |
| 49 | |
| 50 | response = self.client.chat( |
| 51 | model=self.model_string, |
| 52 | message=prompt, |
| 53 | preamble=sys_prompt_arg, |
| 54 | temperature=temperature, |
| 55 | max_tokens=max_tokens, |
| 56 | p=top_p, |
| 57 | ) |
| 58 | |
| 59 | response = response.text |
| 60 | self._save_cache(sys_prompt_arg + prompt, response) |
| 61 | return response |
| 62 | |
| 63 | @retry(wait=wait_random_exponential(min=1, max=5), stop=stop_after_attempt(5)) |
| 64 | def __call__(self, prompt, **kwargs): |
| 65 | return self.generate(prompt, **kwargs) |
| 66 | |