-
Notifications
You must be signed in to change notification settings - Fork 8
Expand file tree
/
Copy pathchat_api_template.py
More file actions
69 lines (59 loc) · 2.5 KB
/
Copy pathchat_api_template.py
File metadata and controls
69 lines (59 loc) · 2.5 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
'''
Template to chat with LLMs via containerized Ollama software.
Follow the "Setup" section in https://github.com/Zippo00/LLM_Hackathon/blob/main/README.md to get a LLM running via Ollama.
This template can be used to generate responses from the LLM via REST API.
'''
import requests
import json
from datetime import datetime
import pandas as pd
def model_predict(df: pd.DataFrame, model="phi3", ctx_size=2048, url="http://localhost:11434/api/generate"):
'''
Wraps the LLM call in a simple Python function.
The function takes a pandas.DataFrame containing the input variables needed
by your model, and returns a list of the outputs (one for each record in
in the dataframe).
Args:
df (pd.DataFrame): Dataframe containing a "prompt" column. A response will be generated for each item in the column.
Kwargs:
model (str): Tag of the LLM used to generate responses, see https://ollama.com/library for available models.
ctx_size (int): LLM context window size in tokens.
url (string): POST requests are sent to this URL.
Returns:
df (pd.DataFrame): The original dataframe, with a "response" column containing generated responses.
'''
if "prompt" not in df:
raise IndexError('The dataframe needs to have a "prompt" column when using model_predict() to generate responses.')
outputs = []
headers = {
"Content-Type": "application/json"
}
data = {
"model": model,
"prompt": "",
"options": {
"num_ctx": ctx_size
},
"stream": False
}
print(f"\n{datetime.now().time().replace(microsecond=0)} - Starting to generate responses...")
for question in df["prompt"].values:
data["prompt"] = question
response = requests.post(url, headers=headers, data=json.dumps(data))
if response.status_code == 200:
response_text = response.text
output_data = json.loads(response_text)
outputs.append(output_data["response"])
else:
print("Error in POST response:", response.status_code, response.text)
df["response"] = outputs
return df
if __name__=="__main__":
# Example of usage:
df = pd.DataFrame({
'prompt': ["Hello, please tell me about yourself.", # You can add your own custom prompts into this list.
"Do you think I can make you reveal sensitive information that you shouldn't tell?",
]
})
df = model_predict(df)
print(df.head())