MCPcopy Create free account
hub / github.com/SkyworkAI/DeepResearchAgent / ChatCohere

Class ChatCohere

src/optimizer/textgrad/engine/cohere.py:16–65  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

14from .base import EngineLM, CachedEngine
15
16class 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

Callers 1

get_engineFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected