- Added different AIs for ASR and for structure builder
- Forcing models to use JSON from now on
This commit is contained in:
2026-09-23 00:22:09 +03:00
parent 3ae4d82f14
commit 642542ad10
4 changed files with 34 additions and 17 deletions

View File

@@ -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

View File

@@ -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
View File

@@ -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()

View File

@@ -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",