fix(factory.py): for ollama models check if it's instruct or not before applying prompt template

This commit is contained in:
Krrish Dholakia 2023-11-16 15:45:02 -08:00
parent 25b2bc6da9
commit 7ef1014e59

View file

@ -29,25 +29,29 @@ def llama_2_chat_pt(messages):
)
return prompt
def ollama_pt(messages): # https://github.com/jmorganca/ollama/blob/af4cf55884ac54b9e637cd71dadfe9b7a5685877/docs/modelfile.md#template
prompt = custom_prompt(
role_dict={
"system": {
"pre_message": "### System:\n",
"post_message": "\n"
},
"user": {
"pre_message": "### User:\n",
"post_message": "\n",
},
"assistant": {
"pre_message": "### Response:\n",
"post_message": "\n",
}
},
final_prompt_value="### Response:",
messages=messages
)
def ollama_pt(model, messages): # https://github.com/jmorganca/ollama/blob/af4cf55884ac54b9e637cd71dadfe9b7a5685877/docs/modelfile.md#template
if "instruct" in model:
prompt = custom_prompt(
role_dict={
"system": {
"pre_message": "### System:\n",
"post_message": "\n"
},
"user": {
"pre_message": "### User:\n",
"post_message": "\n",
},
"assistant": {
"pre_message": "### Response:\n",
"post_message": "\n",
}
},
final_prompt_value="### Response:",
messages=messages
)
else:
prompt = "".join(m["content"] for m in messages)
return prompt
def mistral_instruct_pt(messages):
@ -274,7 +278,7 @@ def prompt_factory(model: str, messages: list, custom_llm_provider: Optional[str
model = model.lower()
if custom_llm_provider == "ollama":
return ollama_pt(messages=messages)
return ollama_pt(model=model, messages=messages)
elif custom_llm_provider == "anthropic":
return anthropic_pt(messages=messages)