211 lines
7.3 KiB
Python
211 lines
7.3 KiB
Python
import re
|
|
import os
|
|
import json
|
|
from pathlib import Path
|
|
import time
|
|
from dotenv import load_dotenv
|
|
from tqdm import tqdm
|
|
|
|
from .paths import PROMPTS_DIR, RESULTS_DIR
|
|
from .providers import chat_completion_options, client_for, parse_model_reference
|
|
from .retry_policy import (
|
|
MAX_REQUEST_ATTEMPTS,
|
|
REQUEST_INTERVAL_SECONDS,
|
|
REQUEST_TIMEOUT_SECONDS,
|
|
STREAM_HEARTBEAT_SECONDS,
|
|
response_is_retryable_failure,
|
|
retry_delay_seconds,
|
|
)
|
|
|
|
# --- Configuration ---
|
|
load_dotenv()
|
|
|
|
# Target model IDs must match the identifiers available in OpenCode Zen.
|
|
TARGET_MODELS = [
|
|
"opencode/qwen3.6-plus"
|
|
# "deepseek-v4-flash",
|
|
# "openai/gpt-4o",
|
|
# "openai/gpt-5",
|
|
# "meta-llama/llama-3.1-405b-instruct",
|
|
# "meta-llama/llama-3.1-405b",
|
|
# "anthropic/claude-opus-4.1",
|
|
# "google/gemini-2.5-pro",
|
|
# "x-ai/grok-4",
|
|
# "deepseek/deepseek-r1-0528:free"
|
|
# "qwen/qwen3-235b-a22b",
|
|
# "openai/gpt-oss-20b",
|
|
# "qwen/qwen-2.5-14b",
|
|
# "qwen/qwen3-30b-a3b",
|
|
# "meta-llama/llama-3.3-70b-instruct",
|
|
# "deepseek/deepseek-r1-distill-qwen-14b",
|
|
# "deepseek/deepseek-r1-distill-llama-70b",
|
|
# "z-ai/glm-4-32b"
|
|
# "mistralai/mistral-small-3.2-24b-instruct",
|
|
# "pangu/pangu-model-name", # Placeholder for PanGu - needs verification
|
|
]
|
|
|
|
# Allows src/run_profile.py to select a model without editing this file.
|
|
if os.getenv("PROFILE_TARGET_MODEL"):
|
|
TARGET_MODELS = [os.environ["PROFILE_TARGET_MODEL"]]
|
|
|
|
def parse_tex_file(file_path):
|
|
"""
|
|
Parses a LaTeX file to extract prompts and their IDs.
|
|
"""
|
|
try:
|
|
with open(file_path, 'r', encoding='utf-8') as f:
|
|
content = f.read()
|
|
except FileNotFoundError:
|
|
print(f"Error: The file at {file_path} was not found.")
|
|
return []
|
|
prompt_regex = re.compile(
|
|
r"\\item\[Prompt\s+([\d\.]+).*?\]\s*``(.*?)''",
|
|
re.DOTALL
|
|
)
|
|
prompts = []
|
|
matches = prompt_regex.finditer(content)
|
|
for match in matches:
|
|
prompt_id = match.group(1).strip()
|
|
prompt_text = ' '.join(match.group(2).strip().split())
|
|
prompts.append({'id': prompt_id, 'text': prompt_text})
|
|
return prompts
|
|
|
|
def consume_chat_stream(stream, activity_callback=None):
|
|
"""Collect final answer text while exposing incremental stream activity."""
|
|
content_parts = []
|
|
for chunk in stream:
|
|
if activity_callback is not None:
|
|
activity_callback(chunk)
|
|
if not chunk.choices:
|
|
continue
|
|
content = chunk.choices[0].delta.content
|
|
if content:
|
|
content_parts.append(content)
|
|
response = "".join(content_parts)
|
|
if not response:
|
|
raise RuntimeError("API stream completed without answer content")
|
|
return response
|
|
|
|
|
|
def get_model_response(model, prompt_text):
|
|
"""
|
|
Gets a response from a specified model through its selected provider.
|
|
"""
|
|
client = client_for(model, timeout=REQUEST_TIMEOUT_SECONDS)
|
|
if not client:
|
|
time.sleep(0.5)
|
|
return f"This is a simulated response from {model.value} because no provider API key was provided."
|
|
|
|
for attempt in range(1, MAX_REQUEST_ATTEMPTS + 1):
|
|
try:
|
|
stream = client.chat.completions.create(
|
|
model=model.model_id,
|
|
messages=[{"role": "user", "content": prompt_text}],
|
|
stream=True,
|
|
**chat_completion_options(model),
|
|
)
|
|
stream_started = False
|
|
last_heartbeat = time.monotonic()
|
|
|
|
def report_activity(chunk):
|
|
nonlocal stream_started, last_heartbeat
|
|
now = time.monotonic()
|
|
if not stream_started:
|
|
request_id = getattr(chunk, "id", None) or "unknown"
|
|
tqdm.write(f"Target stream connected (request_id={request_id}).")
|
|
stream_started = True
|
|
last_heartbeat = now
|
|
elif now - last_heartbeat >= STREAM_HEARTBEAT_SECONDS:
|
|
tqdm.write("Target stream is still receiving output...")
|
|
last_heartbeat = now
|
|
|
|
return consume_chat_stream(stream, report_activity)
|
|
except Exception as error:
|
|
if attempt == MAX_REQUEST_ATTEMPTS:
|
|
return f"Error: API call failed for {model.value}. Details: {error}"
|
|
delay_seconds = retry_delay_seconds(attempt)
|
|
tqdm.write(
|
|
f"Target API error: {error}. Retrying in {delay_seconds:g}s "
|
|
f"({attempt}/{MAX_REQUEST_ATTEMPTS})..."
|
|
)
|
|
time.sleep(delay_seconds)
|
|
|
|
|
|
def response_needs_retry(output_file_path):
|
|
"""Keep successful cached responses, but retry cached API-failure sentinels."""
|
|
if not output_file_path.exists():
|
|
return True
|
|
try:
|
|
return response_is_retryable_failure(output_file_path.read_text(encoding="utf-8"))
|
|
except OSError:
|
|
return True
|
|
|
|
def main():
|
|
"""
|
|
Main function to execute the script.
|
|
"""
|
|
comm_records_dir = PROMPTS_DIR
|
|
tex_file_path = comm_records_dir / 'prompt_suite.tex'
|
|
prompts_json_path = comm_records_dir / 'prompts.json'
|
|
results_dir = RESULTS_DIR
|
|
|
|
print("Step 1: Loading prompts...")
|
|
if prompts_json_path.exists():
|
|
print(f"Found cached prompts file at {prompts_json_path}. Loading from JSON.")
|
|
with open(prompts_json_path, 'r', encoding='utf-8') as f:
|
|
extracted_prompts = json.load(f)
|
|
else:
|
|
print(f"No cached prompts file found. Parsing from {tex_file_path}.")
|
|
extracted_prompts = parse_tex_file(tex_file_path)
|
|
if extracted_prompts:
|
|
with open(prompts_json_path, 'w', encoding='utf-8') as f:
|
|
json.dump(extracted_prompts, f, indent=4)
|
|
print(f"Saved extracted prompts to {prompts_json_path}.")
|
|
|
|
if not extracted_prompts:
|
|
print("No prompts found. Exiting.")
|
|
return
|
|
print(f"Loaded {len(extracted_prompts)} prompts.\n")
|
|
|
|
print("Step 2: Iterating through models and prompts to get responses...")
|
|
for model_name in TARGET_MODELS:
|
|
model = parse_model_reference(model_name)
|
|
model_results_dir = results_dir / model.value
|
|
model_results_dir.mkdir(parents=True, exist_ok=True)
|
|
|
|
print(f"\nProcessing model: {model.value}")
|
|
|
|
progress = tqdm(
|
|
extracted_prompts,
|
|
desc=f"Responses: {model.value}",
|
|
unit="prompt",
|
|
dynamic_ncols=True,
|
|
)
|
|
for prompt in progress:
|
|
prompt_id = prompt['id']
|
|
prompt_text = prompt['text']
|
|
progress.set_postfix_str(f"current={prompt_id}")
|
|
|
|
output_file_path = model_results_dir / f"{prompt_id}.txt"
|
|
|
|
if not response_needs_retry(output_file_path):
|
|
progress.set_postfix_str(f"current={prompt_id}, cached")
|
|
continue
|
|
|
|
if output_file_path.exists():
|
|
progress.set_postfix_str(f"current={prompt_id}, retrying failed response")
|
|
|
|
progress.set_postfix_str(f"current={prompt_id}, requesting response")
|
|
response = get_model_response(model, prompt_text)
|
|
|
|
with open(output_file_path, 'w', encoding='utf-8') as f:
|
|
f.write(response)
|
|
|
|
progress.set_postfix_str(f"current={prompt_id}, saved")
|
|
time.sleep(REQUEST_INTERVAL_SECONDS)
|
|
|
|
print("\nExperiment complete.")
|
|
|
|
if __name__ == "__main__":
|
|
main()
|