Repository navigation
Expand file tree
/
Copy pathbatch_api.py
More file actions
158 lines (121 loc) · 4.54 KB
/
Copy pathbatch_api.py
File metadata and controls
158 lines (121 loc) · 4.54 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
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
"""
Testing the OpenAI Batch API in a few different scenarios.
"""
import asyncio
import logging
import pydantic
from dotenv import load_dotenv
import agents
from agents import BatchProcessorIterable, agent_callable
from agents.providers.openai import AzureOpenAIBatchProvider, AzureOpenAIProvider
load_dotenv()
logging.basicConfig(filename="batch_api.log", filemode="w", level=logging.INFO)
# A knock-knock joke return model
class KnockKnock(pydantic.BaseModel):
setup: str = pydantic.Field(description="The setup for the knock-knock joke")
punchline: str = pydantic.Field(
description="The punchline for the knock-knock joke"
)
def __str__(self) -> str:
return f"- Knock-Knock.\n- Who's There?\n- {self.setup}\n- {self.setup} who?\n- {self.punchline}"
def __repr__(self) -> str:
return f"KnockKnock(setup={self.setup!r}, punchline={self.punchline!r})"
# Define an agent
class KnockKnockAgent(agents.StructuredOutputAgent):
"An agent that writes good Knock Knock jokes."
BASE_PROMPT = """
You are a top-comedian that specializes in writing knock-knock jokes.
You're in a competition to write the best knock-knock joke, and you have to really think of a unique joke to win.
You'll have two tasks:
1. Plan for how you'd write a knock-knock joke that would win in a competition using the plan() tool
2. Write the joke using the KnockKnock model, which should consist of a setup (the part that follows "who's there?") and a punchline (the part that follows "who?").
"""
def __init__(
self,
model_name: str | None = None,
stopping_condition=None,
provider=None,
tools=None,
callbacks=None,
oai_kwargs=None,
**fmt_kwargs,
):
if oai_kwargs is not None:
# Making temp higher
oai_kwargs.update({"temperature": 0.9, "parallel_tool_calls": False})
else:
oai_kwargs = {"temperature": 0.9, "parallel_tool_calls": False}
super().__init__(
response_model=KnockKnock,
model_name=model_name,
stopping_condition=stopping_condition,
provider=provider,
tools=tools,
callback=callbacks,
oai_kwargs=oai_kwargs,
**fmt_kwargs,
)
@agent_callable(
"Plan for how you'd write a knock-knock joke that would win in a competition.",
{"text": "Your plan for the joke."},
)
def plan(self, text: str) -> str:
return "Good thinking. Now send your joke."
class KnockKnockJudge(agents.PredictionAgent):
"""
The Knock-knock joke contest Judge
"""
BASE_PROMPT = """
You are the judge of a knock-knock joke competition.
You know a good joke when you hear it, and it's time to crown a winner.
Read the following knock-knock jokes and crown a winner using a tool call with the joke number:
{jokes}
""".strip()
def __init__(self, jokes: list[str], provider=None, **fmt_kwargs):
labels = [str(i) for i, _ in enumerate(jokes)]
if fmt_kwargs is None:
fmt_kwargs = {}
fmt_kwargs.update(
{"jokes": "\n".join(f"Joke {i}:\n{joke}" for i, joke in enumerate(jokes))}
)
super().__init__(labels=labels, provider=provider, **fmt_kwargs)
async def agents_example():
async with AzureOpenAIBatchProvider(
"gpt-4o-batch",
batch_size=5,
n_workers=2,
progress_max_items=10,
) as provider:
# Kind of a hacky way to use this, but just for demonstration purposes
proc = BatchProcessorIterable(
[i for i in range(10)],
KnockKnockAgent,
batch_size=1,
provider=provider,
n_retry=1,
)
jokes = await proc.process()
usage = provider.usage
print("Got the following entries:")
print(
"\n".join(
f"Joke {i}:\n{KnockKnock.model_validate(res)!s}"
for i, res in enumerate(jokes)
)
)
# Judge the jokes
judge = KnockKnockJudge(
[str(KnockKnock.model_validate(joke)) for joke in jokes],
# Chat provider since it's a single call
provider=AzureOpenAIProvider(
model_name="gpt-4o-nofilter", interactive=False
),
)
await judge()
best_joke_idx = judge.answer["labels"][0]
print(
f"The judge crowned a winner!:\n\n{KnockKnock.model_validate(jokes[int(best_joke_idx)])!s}"
)
print(f"Token usage: {usage}")
if __name__ == "__main__":
asyncio.run(agents_example())