| 14 | from .base import EngineLM, CachedEngine |
| 15 | |
| 16 | class ChatTogether(EngineLM, CachedEngine): |
| 17 | DEFAULT_SYSTEM_PROMPT = "You are a helpful, creative, and smart assistant." |
| 18 | |
| 19 | def __init__( |
| 20 | self, |
| 21 | model_string="meta-llama/Llama-3-70b-chat-hf", |
| 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_together_{model_string}.db") |
| 29 | super().__init__(cache_path=cache_path) |
| 30 | |
| 31 | self.system_prompt = system_prompt |
| 32 | if os.getenv("TOGETHER_API_KEY") is None: |
| 33 | raise ValueError("Please set the TOGETHER_API_KEY environment variable if you'd like to use OpenAI models.") |
| 34 | |
| 35 | self.client = Together( |
| 36 | api_key=os.getenv("TOGETHER_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.completions.create( |
| 51 | model=self.model_string, |
| 52 | messages=[ |
| 53 | {"role": "system", "content": sys_prompt_arg}, |
| 54 | {"role": "user", "content": prompt}, |
| 55 | ], |
| 56 | frequency_penalty=0, |
| 57 | presence_penalty=0, |
| 58 | stop=None, |
| 59 | temperature=temperature, |
| 60 | max_tokens=max_tokens, |
| 61 | top_p=top_p, |
| 62 | ) |
| 63 | |
| 64 | response = response.choices[0].message.content |
| 65 | self._save_cache(sys_prompt_arg + prompt, response) |
| 66 | return response |
| 67 | |
| 68 | @retry(wait=wait_random_exponential(min=1, max=5), stop=stop_after_attempt(5)) |
| 69 | def __call__(self, prompt, **kwargs): |
| 70 | return self.generate(prompt, **kwargs) |
| 71 | |