PerturbReason/PerturbReason_dataset_code
050
1"""2LLM API Batch Inference3Replaces vLLM-based Qwen inference with OpenAI-compatible API calls.4Runs on compute cluster (requires web proxy configuration in the shell script).5 6Usage:7 python llm_api_batch.py \8 --config config.txt \9 --input_dir <dir_with_jsonl_files> \10 --output_dir <output_dir> \11 [--model gpt-4o-mini] \12 [--max_tokens 4096] \13 [--workers 16] \14 [--test_limit 50] # set >0 to run a small test batch15"""16 17import argparse18import asyncio19import json20import os21import glob22import time23from pathlib import Path24 25import httpx26 27 28# ---------------------------------------------------------------------------29# Key pool – rotates to next API key on quota / rate-limit errors30# ---------------------------------------------------------------------------31 32class KeyPool:33 """Thread-safe (asyncio) pool of API keys with automatic rotation.34 35 On a 429 or quota error the pool advances to the next available key.36 When all keys are exhausted every subsequent call immediately returns37 False from rotate_on_quota() so callers can return [KEY_EXHAUSTED].38 """39 40 _QUOTA_STATUS = {429}41 _QUOTA_KEYWORDS = ("quota", "rate_limit", "rate limit", "insufficient_quota",42 "exceeded", "too many requests")43 44 def __init__(self, keys: list):45 self._keys = [k for k in keys if k]46 if not self._keys:47 raise ValueError("No API keys provided.")48 self._idx = 049 self._exhausted: set = set()50 self._lock = asyncio.Lock()51 self._all_done = False52 53 @property54 def current_key(self) -> str:55 return self._keys[self._idx]56 57 @property58 def current_idx(self) -> int:59 return self._idx60 61 @property62 def all_exhausted(self) -> bool:63 return self._all_done64 65 def is_quota_error(self, status_code: int, body: str) -> bool:66 if status_code in self._QUOTA_STATUS:67 return True68 body_lower = body.lower()69 return any(kw in body_lower for kw in self._QUOTA_KEYWORDS)70 71 async def rotate_on_quota(self, failed_idx: int) -> bool:72 """Mark failed_idx as quota-exhausted and rotate to the next key.73 74 Returns True if rotation succeeded (caller should retry with new key),75 False if all keys are now exhausted.76 """77 async with self._lock:78 if self._all_done:79 return False80 self._exhausted.add(failed_idx)81 for i in range(len(self._keys)):82 if i not in self._exhausted:83 if self._idx != i:84 print(f" [KEY_POOL] Key #{failed_idx + 1} quota hit "85 f"→ switching to key #{i + 1}.")86 self._idx = i87 return True88 # All keys exhausted89 self._all_done = True90 print(f"\n [KEY_POOL] ⚠ All {len(self._keys)} API key(s) exhausted! "91 f"Partial results saved; re-run to resume.\n")92 return False93 94 95# ---------------------------------------------------------------------------96# Config helpers97# ---------------------------------------------------------------------------98 99def load_config(config_path: str) -> dict:100 """Parse key=value pairs from config.txt (ignores comment lines)."""101 cfg = {}102 with open(config_path, "r") as f:103 for line in f:104 line = line.strip()105 if not line or line.startswith("#"):106 continue107 if "=" in line:108 key, _, value = line.partition("=")109 cfg[key.strip()] = value.strip()110 return cfg111 112 113def build_client(cfg: dict) -> tuple:114 """Return (httpx.AsyncClient, KeyPool, base_url) for direct REST calls to TAMU AI."""115 # Collect all API keys: primary env var, then TAMU_AI_API_KEY, TAMU_AI_API_KEY_2, …116 keys = []117 env_key = os.environ.get("TAMU_AI_API_KEY")118 if env_key:119 keys.append(env_key)120 # Pick up numbered keys from config: TAMU_AI_API_KEY, TAMU_AI_API_KEY_2, …121 for suffix in ("", "_2", "_3", "_4", "_5"):122 cfg_key = cfg.get(f"TAMU_AI_API_KEY{suffix}", "").strip()123 if cfg_key and cfg_key not in keys:124 keys.append(cfg_key)125 if not keys:126 raise ValueError(127 "No API key found. Set TAMU_AI_API_KEY env var or add it to config.txt."128 )129 130 base_url = cfg.get("BASE_URL", "https://chat-api.tamu.ai/openai")131 key_pool = KeyPool(keys)132 print(f"API key pool : {len(keys)} key(s) loaded")133 134 http_client = httpx.AsyncClient(135 proxy=os.environ.get("https_proxy") or os.environ.get("http_proxy"),136 timeout=httpx.Timeout(120.0),137 )138 139 return http_client, key_pool, base_url140 141 142# ---------------------------------------------------------------------------143# Inference144# ---------------------------------------------------------------------------145 146async def call_api(147 http_client: httpx.AsyncClient,148 key_pool: KeyPool,149 base_url: str,150 prompt: str,151 model: str,152 max_tokens: int,153 semaphore: asyncio.Semaphore,154 max_retries: int = 5,155) -> str:156 """Call the TAMU AI chat completions endpoint.157 158 * Rotates to the next API key automatically on 429 / quota errors.159 * Returns '[KEY_EXHAUSTED]' if the entire KeyPool is drained.160 * Returns '[API_ERROR]' after max_retries non-quota failures.161 """162 payload = {163 "model": model,164 "messages": [{"role": "user", "content": prompt}],165 "max_tokens": max_tokens,166 "temperature": 0.7,167 "stream": True, # TAMU AI always returns SSE168 }169 regular_retries = 0170 while True:171 # Short-circuit if pool already exhausted (set by another coroutine)172 if key_pool.all_exhausted:173 return "[KEY_EXHAUSTED]"174 if regular_retries > max_retries:175 print(" [error] Giving up after max retries.")176 return "[API_ERROR]"177 178 current_idx = key_pool.current_idx179 headers = {180 "accept": "application/json",181 "Content-Type": "application/json",182 "Authorization": f"Bearer {key_pool.current_key}",183 }184 try:185 async with semaphore:186 response = await http_client.post(187 f"{base_url}/chat/completions",188 headers=headers,189 json=payload,190 )191 192 # --- Quota / rate-limit: rotate key and retry immediately ---193 if key_pool.is_quota_error(response.status_code, response.text):194 rotated = await key_pool.rotate_on_quota(current_idx)195 if not rotated:196 return "[KEY_EXHAUSTED]"197 await asyncio.sleep(1) # brief pause before using new key198 continue # don't count as a regular retry199 200 # --- Other HTTP errors: regular back-off retry ---201 if response.status_code >= 400:202 regular_retries += 1203 wait = 2 ** regular_retries204 print(f" [warn] HTTP {response.status_code} "205 f"(retry {regular_retries}/{max_retries}). Wait {wait}s…")206 await asyncio.sleep(wait)207 continue208 209 # --- Parse SSE stream ---210 content_parts = []211 for line in response.text.splitlines():212 line = line.strip()213 if not line.startswith("data:"):214 continue215 data_str = line[5:].strip()216 if data_str == "[DONE]":217 break218 try:219 chunk = json.loads(data_str)220 delta = chunk.get("choices", [{}])[0].get("delta", {})221 content_parts.append(delta.get("content") or "")222 except (json.JSONDecodeError, KeyError, IndexError):223 continue224 return "".join(content_parts)225 226 except Exception as e:227 regular_retries += 1228 wait = 2 ** regular_retries229 print(f" [warn] API error (retry {regular_retries}/{max_retries}): {e}. Wait {wait}s…")230 await asyncio.sleep(wait)231 232 233async def process_file(234 http_client: httpx.AsyncClient,235 key_pool: KeyPool,236 base_url: str,237 file_path: str,238 output_dir: str,239 model: str,240 max_tokens: int,241 workers: int,242 test_limit: int,243) -> dict:244 """Process a single JSONL file; returns timing stats.245 246 Supports partial resume: if the output file already exists with K lines,247 the first K input records are skipped and results are appended.248 """249 file_name = os.path.basename(file_path)250 output_path = os.path.join(output_dir, f"llm_api_pred_{file_name}")251 252 # Load all records253 all_records = []254 with open(file_path, "r", encoding="utf-8") as f:255 for line in f:256 line = line.strip()257 if not line:258 continue259 try:260 all_records.append(json.loads(line))261 except json.JSONDecodeError:262 continue263 264 if test_limit > 0:265 all_records = all_records[:test_limit]266 print(f" [test] Limiting to first {len(all_records)} records.")267 268 records_with_prompt = [r for r in all_records if r.get("prompt")]269 expected_total = len(records_with_prompt)270 271 if not expected_total:272 print(f" [warn] No records found in {file_name}.")273 return {"file": file_name, "skipped": False, "n": 0}274 275 # --- Partial resume: count already-written lines ---276 done_count = 0277 if os.path.exists(output_path):278 done_count = sum(1 for line in open(output_path, encoding="utf-8") if line.strip())279 if done_count >= expected_total:280 print(f" [skip] {file_name} already complete ({done_count}/{expected_total}).")281 return {"file": file_name, "skipped": True, "n": done_count}282 print(f" [resume] {file_name}: {done_count}/{expected_total} done, "283 f"resuming from record {done_count + 1}.")284 285 records_to_process = records_with_prompt[done_count:]286 print(f"Processing {file_name} "287 f"({len(records_to_process)} of {expected_total} records) → {output_path}")288 289 semaphore = asyncio.Semaphore(workers)290 t0 = time.perf_counter()291 292 tasks = [293 call_api(http_client, key_pool, base_url,294 record["prompt"], model, max_tokens, semaphore)295 for record in records_to_process296 ]297 outputs = await asyncio.gather(*tasks)298 299 elapsed = time.perf_counter() - t0300 301 # --- Write results; stop before the first [KEY_EXHAUSTED] so resume works ---302 n_written = 0303 n_exhausted = 0304 write_mode = "a" if done_count > 0 else "w"305 with open(output_path, write_mode, encoding="utf-8") as f_out:306 for record, generated_text in zip(records_to_process, outputs):307 if generated_text == "[KEY_EXHAUSTED]":308 n_exhausted += 1309 break # stop; remaining records will be picked up on re-run310 entry = {311 "source_file": file_name,312 "prompt": record["prompt"],313 "ground_truth_response": record.get("response", ""),314 "model_output": generated_text,315 "id": record.get("id"),316 "label": record.get("label"),317 "pert_type": record.get("pert_type"),318 "sample_idx": record.get("sample_idx"),319 "model": model,320 }321 f_out.write(json.dumps(entry) + "\n")322 n_written += 1323 324 if n_exhausted > 0:325 remaining = expected_total - done_count - n_written326 print(f" [KEY_EXHAUSTED] Wrote {n_written} more records; "327 f"{remaining} pending → re-run this script to resume.")328 else:329 rate = n_written / elapsed if elapsed > 0 else 0330 print(f" Done: {n_written} items in {elapsed:.1f}s ({rate:.2f} items/s)")331 332 return {333 "file": file_name,334 "skipped": False,335 "n": n_written,336 "elapsed": elapsed,337 "rate": n_written / elapsed if elapsed > 0 else 0,338 "exhausted": n_exhausted > 0,339 }340 341 342# ---------------------------------------------------------------------------343# Main344# ---------------------------------------------------------------------------345 346async def main_async(args):347 cfg = load_config(args.config)348 model = args.model or cfg.get("MODEL", "gpt-4o-mini")349 350 http_client, key_pool, base_url = build_client(cfg)351 352 os.makedirs(args.output_dir, exist_ok=True)353 input_files = sorted(glob.glob(os.path.join(args.input_dir, "*.jsonl")))354 if not input_files:355 print(f"No .jsonl files found in {args.input_dir}")356 return357 358 print(f"Model : {model}")359 print(f"Workers : {args.workers}")360 print(f"Max tokens : {args.max_tokens}")361 print(f"Test limit : {args.test_limit if args.test_limit > 0 else 'disabled (full run)'}")362 print(f"Input files : {len(input_files)}")363 print(f"Output dir : {args.output_dir}")364 print()365 366 all_stats = []367 global_t0 = time.perf_counter()368 369 for fp in input_files:370 stats = await process_file(371 http_client=http_client,372 key_pool=key_pool,373 base_url=base_url,374 file_path=fp,375 output_dir=args.output_dir,376 model=model,377 max_tokens=args.max_tokens,378 workers=args.workers,379 test_limit=args.test_limit,380 )381 all_stats.append(stats)382 # Stop submitting new files if all keys are exhausted383 if key_pool.all_exhausted:384 remaining = [f for f in input_files if f != fp and385 not os.path.exists(os.path.join(386 args.output_dir, f"llm_api_pred_{os.path.basename(f)}"))]387 if remaining:388 print(f" [KEY_POOL] Skipping {len(remaining)} remaining file(s); "389 f"re-run to process them.")390 break391 392 global_elapsed = time.perf_counter() - global_t0393 394 # --- Summary & Speed Estimate ---395 total_items = sum(s.get("n", 0) for s in all_stats if not s.get("skipped"))396 if total_items > 0 and global_elapsed > 0:397 overall_rate = total_items / global_elapsed398 target = 50_000399 est_seconds = target / overall_rate400 est_hours = est_seconds / 3600401 print()402 print("=" * 55)403 print(f" Processed : {total_items} items in {global_elapsed:.1f}s")404 print(f" Throughput : {overall_rate:.2f} items/s")405 print(f" ETA 50k pts : {est_hours:.1f} hours ({est_seconds/60:.0f} min)")406 print("=" * 55)407 408 409def parse_args():410 parser = argparse.ArgumentParser(description="LLM API batch inference for PerturbQA")411 parser.add_argument("--config", type=str, default="config.txt",412 help="Path to config.txt with API key")413 parser.add_argument("--input_dir", type=str, required=True,414 help="Directory containing *.jsonl input files")415 parser.add_argument("--output_dir", type=str, required=True,416 help="Directory to write prediction outputs")417 parser.add_argument("--model", type=str, default="",418 help="Model name (overrides config.txt MODEL; default gpt-4o-mini)")419 parser.add_argument("--max_tokens", type=int, default=4096,420 help="Max tokens per response")421 parser.add_argument("--workers", type=int, default=16,422 help="Max concurrent API calls (tune to stay under rate limit)")423 parser.add_argument("--test_limit", type=int, default=0,424 help="If >0, only process first N records per file (for testing)")425 return parser.parse_args()426 427 428if __name__ == "__main__":429 args = parse_args()430 asyncio.run(main_async(args))431 