Initial commit
This commit is contained in:
@@ -0,0 +1,210 @@
|
||||
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()
|
||||
Reference in New Issue
Block a user