File size: 9,570 Bytes
96e6518
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
from __future__ import annotations

import argparse
import asyncio
import fcntl
import json
import os
import random
import re
import sys
import types
from collections import defaultdict
from pathlib import Path
from typing import Any

from langchain_core.messages import HumanMessage, SystemMessage
from langchain_openai import ChatOpenAI

LAB_BENCH_ROOT = Path("/225040511/project/LAB-Bench")
PROJECT_ROOT = Path("/225040511/project/react_code_bioagent_deepseek")
DEFAULT_OUTPUT_ROOT = PROJECT_ROOT / "labbench_runs" / "base_llm"
DEFAULT_RESULT_FILE = DEFAULT_OUTPUT_ROOT / "base_llm_results.jsonl"
DEFAULT_REASONING_LOG = DEFAULT_OUTPUT_ROOT / "base_llm_reasoning.log"
DEFAULT_EVALS = ("DbQA", "SeqQA")
DEFAULT_DEV_SIZE = 45
DEFAULT_TEST_SIZE = 315
DEFAULT_SEED = 20260514

ANSWER_RE = re.compile(r"\[ANSWER\]\s*([A-Z])\s*\[/ANSWER\]", re.IGNORECASE)
LETTER_RE = re.compile(r"\b([A-Z])\b", re.IGNORECASE)

sys.path.insert(0, str(LAB_BENCH_ROOT))


def install_labbench_import_stubs() -> None:
    """Stub optional provider packages needed only while importing LAB-Bench."""
    if "vertexai" not in sys.modules:
        vertexai = types.ModuleType("vertexai")
        vertexai.init = lambda *_args, **_kwargs: None
        sys.modules["vertexai"] = vertexai
    if "google.auth" not in sys.modules:
        google = sys.modules.setdefault("google", types.ModuleType("google"))
        auth = types.ModuleType("google.auth")
        auth.default = lambda *_args, **_kwargs: (types.SimpleNamespace(refresh=lambda *_a, **_k: None, token=""), None)
        transport = types.ModuleType("google.auth.transport")
        requests = types.ModuleType("google.auth.transport.requests")
        requests.Request = lambda *_args, **_kwargs: None
        transport.requests = requests
        auth.transport = transport
        google.auth = auth
        sys.modules["google.auth"] = auth
        sys.modules["google.auth.transport"] = transport
        sys.modules["google.auth.transport.requests"] = requests
    if "chembench" not in sys.modules:
        chembench = types.ModuleType("chembench")
        sys.modules["chembench"] = chembench
        constant = types.ModuleType("chembench.constant")
        constant.COT_PROMPT = "Think step by step."
        constant.MCQ_REGEX_TEMPLATE_1 = r"\[ANSWER\]\s*([A-Z])\s*\[/ANSWER\]"
        sys.modules["chembench.constant"] = constant
        prompter = types.ModuleType("chembench.prompter")
        prompter.prepare_mcq_answer = lambda text, *_args, **_kwargs: text
        sys.modules["chembench.prompter"] = prompter
        utils = types.ModuleType("chembench.utils")
        utils.create_multiple_choice_regex = lambda letters: r"\b(" + "|".join(letters) + r")\b"
        utils.post_process_prompts = lambda text: text
        utils.run_regex = lambda _regex, text, return_first=True: None
        sys.modules["chembench.utils"] = utils


install_labbench_import_stubs()
import labbench  # noqa: E402


def load_dotenv_files(paths: list[Path]) -> None:
    for path in paths:
        if not path.exists():
            continue
        for raw_line in path.read_text(encoding="utf-8", errors="replace").splitlines():
            line = raw_line.strip()
            if not line or line.startswith("#") or "=" not in line:
                continue
            key, value = line.split("=", 1)
            os.environ.setdefault(key.strip(), value.strip().strip('"').strip("'"))


def make_model(args: argparse.Namespace) -> ChatOpenAI:
    api_key = args.api_key or os.getenv("DEEPSEEK_API_KEY") or os.getenv("BIOMNI_CUSTOM_API_KEY") or os.getenv("OPENAI_API_KEY")
    if not api_key:
        raise SystemExit("Missing API key. Set DEEPSEEK_API_KEY, BIOMNI_CUSTOM_API_KEY, or OPENAI_API_KEY.")
    model_name = args.model
    if (
        "deepseek" in args.base_url.lower()
        and not args.allow_deepseek_reasoner
        and ("reasoner" in model_name.lower() or "thinking" in model_name.lower())
    ):
        print(
            f"DeepSeek model {model_name!r} uses thinking/reasoning_content mode; "
            "falling back to 'deepseek-chat' for LAB-Bench base LLM calls.",
            flush=True,
        )
        model_name = "deepseek-chat"
    return ChatOpenAI(
        model=model_name,
        api_key=api_key,
        base_url=args.base_url,
        temperature=args.temperature,
        timeout=args.llm_timeout,
        max_retries=args.max_retries,
    )


def parse_answer(text: str, n_choices: int) -> str:
    valid = set("ABCDEFGHIJKLMNOPQRSTUVWXYZ"[:n_choices])
    if match := ANSWER_RE.search(text or ""):
        letter = match.group(1).upper()
        if letter in valid:
            return letter
    for match in LETTER_RE.finditer(text or ""):
        letter = match.group(1).upper()
        if letter in valid:
            return letter
    return "A"


def load_eval(eval_name: str) -> labbench.Evaluator:
    return labbench.Evaluator(labbench.Eval(eval_name), debug=False, open_answer=False, use_hf=False)


def split_counts(total: int, eval_names: list[str], seed: int) -> dict[str, int]:
    sizes = {name: len(load_eval(name).eval_set.instances) for name in eval_names}
    total_available = sum(sizes.values())
    if total >= total_available:
        return sizes
    raw = {name: total * sizes[name] / total_available for name in eval_names}
    counts = {name: int(raw[name]) for name in eval_names}
    remaining = total - sum(counts.values())
    rng = random.Random(seed)
    order = sorted(eval_names, key=lambda name: (raw[name] - counts[name], rng.random()), reverse=True)
    for name in order[:remaining]:
        counts[name] += 1
    return counts


def select_instances(
    evaluator: labbench.Evaluator,
    *,
    eval_name: str,
    split: str,
    split_size: int,
    dev_count: int,
    seed: int,
    debug: bool,
    shard_index: int,
    shard_count: int,
) -> list[tuple[str, Any]]:
    instances = list(evaluator.eval_set.instances)
    rng = random.Random(f"{seed}:{eval_name}:question-set")
    rng.shuffle(instances)
    if debug:
        selected = instances[: min(3, len(instances))]
    elif split == "dev":
        selected = instances[: min(split_size, len(instances))]
    elif split == "test":
        start = min(dev_count, len(instances))
        selected = instances[start : min(start + split_size, len(instances))]
    else:
        selected = instances
    if shard_count > 1:
        total = len(selected)
        chunk_size = (total + shard_count - 1) // shard_count
        selected = selected[min(total, shard_index * chunk_size) : min(total, (shard_index + 1) * chunk_size)]
    return selected


def load_completed_questions(path: Path | None) -> set[str]:
    if path is None or not path.exists():
        return set()
    completed: set[str] = set()
    for raw_line in path.read_text(encoding="utf-8", errors="replace").splitlines():
        if not raw_line.strip():
            continue
        try:
            record = json.loads(raw_line)
        except json.JSONDecodeError:
            continue
        question = str(record.get("question") or "").strip()
        if question:
            completed.add(question)
    return completed


def append_text_locked(path: Path, text: str) -> None:
    path.parent.mkdir(parents=True, exist_ok=True)
    with path.open("a", encoding="utf-8") as handle:
        fcntl.flock(handle.fileno(), fcntl.LOCK_EX)
        handle.write(text)
        handle.flush()
        os.fsync(handle.fileno())
        fcntl.flock(handle.fileno(), fcntl.LOCK_UN)


def append_jsonl_locked(path: Path, payload: dict[str, Any]) -> None:
    append_text_locked(path, json.dumps(payload, ensure_ascii=False, default=str) + "\n")


def build_prompt(input_obj: Any, eval_name: str) -> str:
    choices = "\n".join(input_obj.choices)
    return f"""
The following is a multiple-choice LAB-Bench biology question from {eval_name}.
Please answer by responding with the letter of the correct answer.

Question:
{input_obj.question}

Options:
{choices}

You MUST include the letter of the correct answer within the following tags: [ANSWER] and [/ANSWER].
For example: [ANSWER]A[/ANSWER]
Always answer in exactly this format of a single letter between the tags, even if you are unsure.
""".strip()


class BaseLLMLabBenchAgent:
    def __init__(self, model: ChatOpenAI):
        self.model = model

    async def run_task(self, input_obj: Any, eval_name: str) -> tuple[str, str, str]:
        prompt = build_prompt(input_obj, eval_name)
        response = await self.model.ainvoke(
            [
                SystemMessage(content="You are a careful biology benchmark assistant. Return exactly one [ANSWER]X[/ANSWER] tag."),
                HumanMessage(content=prompt),
            ]
        )
        raw_output = str(response.content)
        return parse_answer(raw_output, len(input_obj.choices)), raw_output, prompt


def compute_metrics(results: list[dict[str, Any]]) -> dict[str, float]:
    n_total = len(results)
    n_correct = sum(bool(r["correct"]) for r in results)
    n_sure = sum(bool(r["sure"]) for r in results)
    return {
        "accuracy": n_correct / n_total if n_total else 0.0,
        "precision": n_correct / n_sure if n_sure else 0.0,
        "coverage": n_sure / n_total if n_total else 0.0,
        "n_total": n_total,
    }


def compact_record(result: dict[str, Any]) -> dict[str, str]:
    return {
        "question": str(result.get("question") or ""),
        "answer": str(result.get("target_choice") or ""),
        "agent_answer": str(result.get("agent_output") or ""),
    }