Team Ai
Datasetpublic

PerturbReason/PerturbReason_dataset_code

sourceHugging Facemitupdated 5mo agoView on Hugging Face
0likes50downloads
llm_api_batch.py431 linesDownload Raw Back to Baselines
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