jvamvas's picture
Slim Azure-only llm clients; remove description/preamble
eb93814 verified
Raw
History Blame Contribute Delete
5.14 kB
import json
import os
from dataclasses import dataclass
from pathlib import Path
from typing import Optional
import diskcache
import openai
from dotenv import load_dotenv
from inference_utils import sanitize_path_component
from llm_query.span_marking.utils import is_pass_response
load_dotenv()
CONFIGS_DIR = Path(__file__).parent / "configs"
@dataclass
class SpanMarkingLLMResponse:
model: str
config: str
text_a: str
text_b: str
marked_text_a: Optional[str]
response_obj: dict
class AzureClient:
API_KEY_NAME = "AZURE_API_KEY"
def __init__(
self,
model_name: str = "gpt-5.6-terra",
config_name: str = "config.span_marking.v0.6.json",
cache_directory: Optional[Path] = None,
use_cache: bool = True,
):
self.config_name = config_name
with open(CONFIGS_DIR / config_name) as config_file:
self.config = json.load(config_file)
assert "{text_a}" in self.config["prompt_template"]
assert "{text_b}" in self.config["prompt_template"]
for example in self.config.get("examples", []):
assert "text_a" in example
assert "text_b" in example
assert "marked_text_a" in example
self.model_name = model_name
self.BASE_URL = os.environ.get("AZURE_BASE_URL")
self.cache_directory = cache_directory or Path(__file__).parent / ".llm_cache"
self.client = openai.Client(
api_key=os.environ.get(self.API_KEY_NAME),
base_url=self.BASE_URL,
)
self.cache = None
if use_cache:
self.cache = diskcache.Cache(self.model_cache_dir)
@property
def model_cache_dir(self) -> Path:
sanitized_model_name = sanitize_path_component(self.model_name)
config_stem = sanitize_path_component(self.config_name.removesuffix(".json"))
return self.cache_directory / f"{sanitized_model_name}_{config_stem}"
def _get_cache_key(self, text_a: str, text_b: str) -> tuple[str, str, str, str]:
return (self.model_name, self.config_name, text_a, text_b)
def _completion_extra_kwargs(self) -> dict:
kwargs = {"reasoning_effort": "none"}
for key in ("seed", "temperature"):
if key in self.config:
kwargs[key] = self.config[key]
return kwargs
def _format_prompt(self, text_a: str, text_b: str) -> str:
return (
self.config["prompt_template"]
.replace("{text_a}", text_a)
.replace("{text_b}", text_b)
)
def _build_messages(self, text_a: str, text_b: str) -> list[dict]:
messages = []
if self.config.get("use_examples", False):
for example in self.config.get("examples", []):
messages.append({
"role": "user",
"content": self._format_prompt(example["text_a"], example["text_b"]),
})
messages.append({
"role": "assistant",
"content": example["marked_text_a"],
})
messages.append({
"role": "user",
"content": self._format_prompt(text_a, text_b),
})
return messages
def query(self, text_a: str, text_b: str) -> SpanMarkingLLMResponse:
cache_key = self._get_cache_key(text_a, text_b)
if self.cache is not None and cache_key in self.cache:
return self._build_response(text_a, text_b, self.cache[cache_key])
completion = self.client.chat.completions.create(
model=self.model_name,
messages=self._build_messages(text_a, text_b),
**self._completion_extra_kwargs(),
)
response_obj = completion.to_dict()
if self.cache is not None:
self.cache[cache_key] = response_obj
return self._build_response(text_a, text_b, response_obj)
@staticmethod
def _normalize_marked_text(content: str, text_a: str) -> str:
stripped = content.strip()
if is_pass_response(stripped):
return text_a
candidate = stripped
if len(candidate) >= 2 and candidate[0] == candidate[-1] and candidate[0] in "\"'":
candidate = candidate[1:-1].strip()
return candidate
return stripped
def _build_response(
self,
text_a: str,
text_b: str,
response_obj: dict,
) -> SpanMarkingLLMResponse:
llm_response = SpanMarkingLLMResponse(
model=self.model_name,
config=self.config_name,
text_a=text_a,
text_b=text_b,
marked_text_a=None,
response_obj=response_obj,
)
try:
content = response_obj["choices"][0]["message"]["content"]
except (KeyError, IndexError, TypeError):
return llm_response
if content is None:
return llm_response
llm_response.marked_text_a = self._normalize_marked_text(content, text_a)
return llm_response