WIP
- Added different AIs for ASR and for structure builder - Forcing models to use JSON from now on
This commit is contained in:
3
agent.py
3
agent.py
@@ -1,5 +1,6 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Literal
|
||||
from pydantic import BaseModel
|
||||
|
||||
import httpx2
|
||||
from openai import OpenAI
|
||||
@@ -35,7 +36,7 @@ class Agent:
|
||||
"content": m.content
|
||||
}
|
||||
)
|
||||
response = self._client.chat.completions.create(
|
||||
response = self._client.chat.completions.parse(
|
||||
model=self._model,
|
||||
messages=messages_raw,
|
||||
**kwargs
|
||||
|
||||
@@ -52,7 +52,7 @@ class AsrEventizer:
|
||||
retries_left = 5
|
||||
while retries_left > 0:
|
||||
retries_left -= 1
|
||||
response = self._agent.completion(messages=messages)
|
||||
response = self._agent.completion(messages=messages, response_format=_EventizeResult)
|
||||
if debug_dir:
|
||||
with open(f"{debug_dir}/{retries_left}-retries-left.txt", "w") as f:
|
||||
f.write(response)
|
||||
|
||||
37
main.py
37
main.py
@@ -50,6 +50,7 @@ def setup_arguments() -> argparse.Namespace:
|
||||
- argparse namespace
|
||||
"""
|
||||
voice_models = Asr.get_models_list()
|
||||
ai_api_key = os.environ.get("AI_API_KEY", default=...)
|
||||
|
||||
parser = argparse.ArgumentParser(
|
||||
prog="sumka",
|
||||
@@ -66,18 +67,34 @@ def setup_arguments() -> argparse.Namespace:
|
||||
default="ru"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ai-model",
|
||||
"--asr-ai-model",
|
||||
type=str,
|
||||
default="google/gemini-3.1-flash-lite"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ai-base-url",
|
||||
"--asr-ai-base-url",
|
||||
type=str,
|
||||
default="https://api.proxyapi.ru/v1"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ai-api-key",
|
||||
type=str
|
||||
"--asr-ai-api-key",
|
||||
type=str,
|
||||
default=ai_api_key
|
||||
)
|
||||
parser.add_argument(
|
||||
"--structure-ai-model",
|
||||
type=str,
|
||||
default="google/gemini-3.5-flash"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--structure-ai-base-url",
|
||||
type=str,
|
||||
default="https://api.proxyapi.ru/v1"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--structure-ai-api-key",
|
||||
type=str,
|
||||
default=ai_api_key
|
||||
)
|
||||
parser.add_argument("-v", action='store_true')
|
||||
return parser.parse_args()
|
||||
@@ -160,9 +177,9 @@ def on_asr_events(current_step: Step, input_data: dict | None) -> tuple[Step | N
|
||||
return (None, None)
|
||||
logging.info("Creating the agent")
|
||||
agent = Agent(
|
||||
model=ARGS.ai_model,
|
||||
base_url=ARGS.ai_base_url,
|
||||
api_key=ARGS.ai_api_key
|
||||
model=ARGS.asr_ai_model,
|
||||
base_url=ARGS.asr_ai_base_url,
|
||||
api_key=ARGS.asr_ai_api_key
|
||||
)
|
||||
logging.info("Creating the windowizer")
|
||||
windowizer = Windowizer()
|
||||
@@ -208,9 +225,9 @@ def on_structure_builder(current_step: Step, input_data: dict | None) -> tuple[S
|
||||
return (None, None)
|
||||
logging.info("Creating the agent")
|
||||
agent = Agent(
|
||||
model=ARGS.ai_model,
|
||||
base_url=ARGS.ai_base_url,
|
||||
api_key=ARGS.ai_api_key
|
||||
model=ARGS.structure_ai_model,
|
||||
base_url=ARGS.structure_ai_base_url,
|
||||
api_key=ARGS.structure_ai_api_key
|
||||
)
|
||||
logging.info("Creating the windowizer")
|
||||
windowizer = Windowizer()
|
||||
|
||||
@@ -79,16 +79,15 @@ class ImageElement(_StrictModel):
|
||||
event_id: int
|
||||
|
||||
|
||||
StructureElement = Annotated[
|
||||
StructureElement = (
|
||||
HeadingElement
|
||||
| ParagraphElement
|
||||
| UnorderedListElement
|
||||
| OrderedListElement
|
||||
| DefinitionElement
|
||||
| ImportantElement
|
||||
| ImageElement,
|
||||
Field(discriminator="type"),
|
||||
]
|
||||
| ImageElement
|
||||
)
|
||||
|
||||
|
||||
class Structure(_StrictModel):
|
||||
@@ -204,7 +203,7 @@ class StructureBuilder:
|
||||
retries_left = self.MAX_RETRIES
|
||||
while retries_left > 0:
|
||||
retries_left -= 1
|
||||
response = self._agent.completion(messages=messages)
|
||||
response = self._agent.completion(messages=messages, response_format=_BuildResult)
|
||||
if debug_dir:
|
||||
with open(
|
||||
f"{debug_dir}/{retries_left}-retries-left.txt",
|
||||
|
||||
Reference in New Issue
Block a user