diff --git a/scripts/orchestrator/build_orchestrator_sft.py b/scripts/orchestrator/build_orchestrator_sft.py index ca2a4d34..431bb236 100644 --- a/scripts/orchestrator/build_orchestrator_sft.py +++ b/scripts/orchestrator/build_orchestrator_sft.py @@ -71,35 +71,58 @@ def main(argv: Optional[list[str]] = None) -> int: # data/orchestrator/raw/{label}_{MMDD}[_{tag}]/ and the file inside is always # data.jsonl — see the run_dir block below. Use --out only to distinguish two # runs of the same orchestrator on the same day (e.g. --out balanced50). - p.add_argument("--out", default=None, - help="Optional tag appended to the run dir name (not a path).") + p.add_argument( + "--out", + default=None, + help="Optional tag appended to the run dir name (not a path).", + ) # The orchestrator == the local Qwen3-8B self-sampling over its own rollouts. - p.add_argument("--orchestrator-endpoint", default="http://localhost:8001/v1", - help="OpenAI-compatible vLLM base URL serving the orchestrator.") + p.add_argument( + "--orchestrator-endpoint", + default="http://localhost:8001/v1", + help="OpenAI-compatible vLLM base URL serving the orchestrator.", + ) p.add_argument("--orchestrator-model", default="qwen3-8b") p.add_argument("--orchestrator-api-key", default="EMPTY") # Provenance stamps: when set, every written record carries these fields so # the JSONL is self-identifying once pooled across model families (gemma vs # qwen). Default None -> no stamping (existing qwen runs unaffected). - p.add_argument("--gen-model", default=None, - help="Full HF id of the generating orchestrator, stamped as " - "record['gen_model'] (e.g. google/gemma-4-26B-A4B-it).") - p.add_argument("--orchestrator-label", default=None, - help="Short label stamped as record['orchestrator_model'] " - "(e.g. gemma-4-26b).") + p.add_argument( + "--gen-model", + default=None, + help="Full HF id of the generating orchestrator, stamped as " + "record['gen_model'] (e.g. google/gemma-4-26B-A4B-it).", + ) + p.add_argument( + "--orchestrator-label", + default=None, + help="Short label stamped as record['orchestrator_model'] (e.g. gemma-4-26b).", + ) # Local OSS model endpoints; repeatable "model_id=base_url" (e.g. # "Qwen/Qwen3.5-9B=http://localhost:8001/v1"). Unmapped local models are # still listed but served unconfigured (base_url=None). - p.add_argument("--local-endpoint", action="append", default=[], - metavar="MODEL_ID=URL", - help="Local model id -> vLLM base URL (repeatable).") - p.add_argument("--max-tasks", type=int, default=None, - help="Cap on tasks (default: all of load_sft_tasks()).") - p.add_argument("--skip-task-ids-from", action="append", default=[], - metavar="GLOB", - help="Glob of prior data.jsonl files; skip task_ids already " - "generated there so this run only does unseen prompts " - "(resume). Repeatable; applied before sharding.") + p.add_argument( + "--local-endpoint", + action="append", + default=[], + metavar="MODEL_ID=URL", + help="Local model id -> vLLM base URL (repeatable).", + ) + p.add_argument( + "--max-tasks", + type=int, + default=None, + help="Cap on tasks (default: all of load_sft_tasks()).", + ) + p.add_argument( + "--skip-task-ids-from", + action="append", + default=[], + metavar="GLOB", + help="Glob of prior data.jsonl files; skip task_ids already " + "generated there so this run only does unseen prompts " + "(resume). Repeatable; applied before sharding.", + ) p.add_argument("--samples-per-task", type=int, default=8) p.add_argument("--max-keep-per-task", type=int, default=1) p.add_argument("--max-turns", type=int, default=8) @@ -107,68 +130,107 @@ def main(argv: Optional[list[str]] = None) -> int: # Orchestrator completion cap per turn. The library default (4096) intermittently # truncates the final answer mid-sentence on longer reasoning turns; bump it so # the trace ends cleanly. Raise further (e.g. 16384) if traces still cut off. - p.add_argument("--max-tokens", type=int, default=8192, - help="Per-turn completion cap for the orchestrator (default 8192).") + p.add_argument( + "--max-tokens", + type=int, + default=8192, + help="Per-turn completion cap for the orchestrator (default 8192).", + ) # Number of tasks rolled out concurrently against vLLM. Each in-flight task # issues its samples sequentially, so ~concurrency requests hit the server at # once. ~30-50 is comfortable on 1xL40S (prefix caching + continuous batching). - p.add_argument("--concurrency", type=int, default=32, - help="Tasks rolled out in parallel (default 32).") + p.add_argument( + "--concurrency", + type=int, + default=32, + help="Tasks rolled out in parallel (default 32).", + ) # Balanced-smoke pull: cap//4 from each of GeneralThought + OpenThoughts # code/math/science instead of the GeneralThought-only fast cap path. Use for a # representative smoke; the real run omits --max-tasks for the full 8K balanced set. - p.add_argument("--balanced", action=argparse.BooleanOptionalAction, default=True, - help="Draw an EVEN cross-domain sample (GeneralThought + " - "OpenThoughts code/math/science). Default ON — pass " - "--no-balanced for the old GeneralThought-only skew.") + p.add_argument( + "--balanced", + action=argparse.BooleanOptionalAction, + default=True, + help="Draw an EVEN cross-domain sample (GeneralThought + " + "OpenThoughts code/math/science). Default ON — pass " + "--no-balanced for the old GeneralThought-only skew.", + ) # Data-parallel sharding: split the (deterministic, seed-42) task list across # N independent driver processes, each pointed at its own orchestrator vLLM # replica. Shard i takes tasks[i::N] (a strided slice, so every shard stays # domain-balanced). Each writes its own --out file; concatenate the shard # JSONLs afterward. Lets the GPU-bound orchestrator scale across idle GPUs. - p.add_argument("--shard-index", type=int, default=0, - help="This shard's index in [0, shard-count).") - p.add_argument("--shard-count", type=int, default=1, - help="Total number of shards (default 1 = no sharding).") + p.add_argument( + "--shard-index", + type=int, + default=0, + help="This shard's index in [0, shard-count).", + ) + p.add_argument( + "--shard-count", + type=int, + default=1, + help="Total number of shards (default 1 = no sharding).", + ) # Throughput knob: stop sampling a task as soon as max-keep-per-task passing # trajectories are found, instead of always running all --samples-per-task. # ~3-4x faster when most tasks solve early, but keeps the *first* passers # rather than the *cheapest* of N (drops the cost-optimisation signal). - p.add_argument("--stop-at-keep", action="store_true", - help="Short-circuit a task once max-keep passing samples found.") + p.add_argument( + "--stop-at-keep", + action="store_true", + help="Short-circuit a task once max-keep passing samples found.", + ) # Anonymize model experts (opaque random labels, uniform description, no cost # line, shuffled order) so the policy can't route on a model's name/position/ # cost. The anon->real map is saved per record (metrics.anon_map) for analysis. - p.add_argument("--anonymize-experts", action="store_true", - help="Hide expert identity (random labels + shuffle) when routing.") + p.add_argument( + "--anonymize-experts", + action="store_true", + help="Hide expert identity (random labels + shuffle) when routing.", + ) # By default we keep EVERY rolled-out trajectory (correct + incorrect), each # tagged with ``correct`` (verifier verdict) and ``kept`` (the cheapest-correct # sample the rejection sampler would pick). Lets you compute accuracy / inspect # failures from the JSONL directly. --rejection-only restores the old # drop-the-failures behaviour (only cheapest-correct written). - p.add_argument("--rejection-only", action="store_true", - help="Drop incorrect rollouts; write only the cheapest-correct " - "sample per task (the original behaviour).") + p.add_argument( + "--rejection-only", + action="store_true", + help="Drop incorrect rollouts; write only the cheapest-correct " + "sample per task (the original behaviour).", + ) # Rollouts call shell_exec / file_write with model-chosen relative paths # (e.g. ``solution.py``), which otherwise land in the repo root. Run them # from a throwaway scratch dir so generated files never dirty the tree. # gitignored via the existing ``scratch/`` rule. - p.add_argument("--scratch-dir", default="scratch/sft-rollouts", - help="CWD for rollouts; stray tool-written files go here.") + p.add_argument( + "--scratch-dir", + default="scratch/sft-rollouts", + help="CWD for rollouts; stray tool-written files go here.", + ) args = p.parse_args(argv) logging.basicConfig(level=logging.INFO, format="%(message)s") # Quiet the per-request HTTP / dataset-stream spam so the run log stays # readable (rollout calls and dataset shards otherwise flood it). - for _noisy in ("httpx", "httpcore", "urllib3", "datasets", "fsspec", - "huggingface_hub", "openai"): + for _noisy in ( + "httpx", + "httpcore", + "urllib3", + "datasets", + "fsspec", + "huggingface_hub", + "openai", + ): logging.getLogger(_noisy).setLevel(logging.WARNING) - # Every run lands in its OWN folder under data/orchestrator/raw/, named to - # MATCH the sft/ file it will eventually produce — only the stage word differs: + # Every run lands in its OWN folder under /raw/, named to MATCH the + # sft/ file it will eventually produce — only the stage word differs: # - # data/orchestrator/raw/qwen_0707/data.jsonl <- every rollout, incl. failures - # data/orchestrator/sft/qwen_clean_0707.jsonl <- reject-sampled from it + # raw/qwen_0707/data.jsonl <- every rollout, incl. failures + # sft/qwen_clean_0707.jsonl <- reject-sampled from it # # so any curated file traces back to its generation run by eye. raw/ sits ABOVE # the sft/rl fork on purpose: SFT keeps only correct+clean rows, but GRPO needs @@ -177,9 +239,13 @@ def main(argv: Optional[list[str]] = None) -> int: # suffix (creation order is preserved, and a MM-DD-HHMMpm stamp is redundant # with mtime). --out is an optional extra tag, NOT a path; the data file inside # is always data.jsonl. + # + # OJ_DATA_ROOT keeps the data OUT of the git checkout (repo-relative default so + # a fresh clone still works); this workspace points it at the experiments tree. + data_root = Path(os.getenv("OJ_DATA_ROOT", "data/orchestrator")) prefix = args.orchestrator_label or "orch" tag = f"_{Path(args.out).stem}" if args.out else "" - base = Path("data/orchestrator/raw") / f"{prefix}_{time.strftime('%m%d')}{tag}" + base = data_root / "raw" / f"{prefix}_{time.strftime('%m%d')}{tag}" run_dir = base.resolve() n = 1 while run_dir.exists(): # same name+day (incl. parallel shards) -> disambiguate @@ -192,6 +258,7 @@ def main(argv: Optional[list[str]] = None) -> int: lock_p.write_text(f"{time.strftime('%m-%d-%I%M%p').lower()} pid={os.getpid()}") args.out = str(out_p) import atexit + atexit.register(lambda: lock_p.exists() and lock_p.unlink()) logging.info("Run dir: %s", run_dir) @@ -199,8 +266,7 @@ def main(argv: Optional[list[str]] = None) -> int: # tracing.run_context() reads these). setdefault so an explicit env wins. os.environ.setdefault("OJ_RUN_LABEL", label) os.environ.setdefault("OJ_GEN_MODEL", args.gen_model or args.orchestrator_model) - os.environ.setdefault("OJ_RUN_STAGE", - "smoke" if args.max_tasks else "prod") + os.environ.setdefault("OJ_RUN_STAGE", "smoke" if args.max_tasks else "prod") os.environ["OJ_CFG_TEMPERATURE"] = str(args.temperature) os.environ["OJ_CFG_MAX_TURNS"] = str(args.max_turns) os.environ["OJ_CFG_ANONYMIZE"] = str(bool(args.anonymize_experts)) @@ -247,8 +313,10 @@ def main(argv: Optional[list[str]] = None) -> int: def rollout_fn(task): try: return run_unified_rollout( - task.instruction, tools, - call_orchestrator=call_orch, dispatch=dispatch, + task.instruction, + tools, + call_orchestrator=call_orch, + dispatch=dispatch, max_turns=args.max_turns, anonymize=args.anonymize_experts, ) @@ -265,6 +333,7 @@ def main(argv: Optional[list[str]] = None) -> int: # remaining unseen tasks stride evenly and never re-cover finished prompts. if args.skip_task_ids_from: import glob as _glob + skip_ids = set() for pattern in args.skip_task_ids_from: for fp in _glob.glob(pattern): @@ -285,18 +354,31 @@ def main(argv: Optional[list[str]] = None) -> int: if skip_ids: n_before = len(tasks) tasks = [t for t in tasks if t.task_id not in skip_ids] - logging.info("Resume: %d done task_ids -> skipping; %d of %d tasks remain", - len(skip_ids), len(tasks), n_before) + logging.info( + "Resume: %d done task_ids -> skipping; %d of %d tasks remain", + len(skip_ids), + len(tasks), + n_before, + ) if args.shard_count > 1: n_all = len(tasks) - tasks = tasks[args.shard_index::args.shard_count] - logging.info("Shard %d/%d -> %d of %d tasks", - args.shard_index, args.shard_count, len(tasks), n_all) + tasks = tasks[args.shard_index :: args.shard_count] + logging.info( + "Shard %d/%d -> %d of %d tasks", + args.shard_index, + args.shard_count, + len(tasks), + n_all, + ) from collections import Counter as _Counter + _dom = _Counter(getattr(t, "domain", "unknown") for t in tasks) - logging.info("Loaded %d SFT tasks | balanced=%s | domain split: %s", - len(tasks), args.balanced, - ", ".join(f"{d}={n}" for d, n in sorted(_dom.items()))) + logging.info( + "Loaded %d SFT tasks | balanced=%s | domain split: %s", + len(tasks), + args.balanced, + ", ".join(f"{d}={n}" for d, n in sorted(_dom.items())), + ) stats = generate_sft_dataset( args.out, @@ -321,8 +403,7 @@ def main(argv: Optional[list[str]] = None) -> int: from render_sft_data import render if stats.get("records_written", 0) > 0: rinfo = render(args.out) - logging.info("Rendered viewable traces: %s + %s", - rinfo["pretty"], rinfo["txt"]) + logging.info("Rendered viewable traces: %s + %s", rinfo["pretty"], rinfo["txt"]) print(json.dumps(stats, indent=2)) return 0 diff --git a/scripts/orchestrator/chat_orchestrator.py b/scripts/orchestrator/chat_orchestrator.py index b7a1ad6a..42dbf83d 100644 --- a/scripts/orchestrator/chat_orchestrator.py +++ b/scripts/orchestrator/chat_orchestrator.py @@ -23,25 +23,43 @@ Usage: # point local expert tiers at their own vLLMs if you have them up (optional): # --local-endpoint Qwen/Qwen3.6-27B=http://localhost:8002/v1 """ + import argparse import sys import time def _parse_args(argv=None): - p = argparse.ArgumentParser(description=__doc__, - formatter_class=argparse.RawDescriptionHelpFormatter) - p.add_argument("--endpoint", default="http://localhost:8020/v1", - help="vLLM endpoint serving the orchestrator checkpoint.") - p.add_argument("--model", default="sft-qwen-8k", - help="Served-model-name of the checkpoint.") - p.add_argument("--api-key", default="EMPTY", help="API key for the endpoint (local vLLM = EMPTY).") + p = argparse.ArgumentParser( + description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter + ) + p.add_argument( + "--endpoint", + default="http://localhost:8020/v1", + help="vLLM endpoint serving the orchestrator checkpoint.", + ) + p.add_argument( + "--model", default="sft-qwen-8k", help="Served-model-name of the checkpoint." + ) + p.add_argument( + "--api-key", + default="EMPTY", + help="API key for the endpoint (local vLLM = EMPTY).", + ) p.add_argument("--temperature", type=float, default=0.0) p.add_argument("--max-turns", type=int, default=8, help="Same default as the eval.") - p.add_argument("--max-tokens", type=int, default=3000, - help="Keep high: the model reasons a lot before it emits the delegation.") - p.add_argument("--local-endpoint", action="append", default=[], - help="MODEL_ID=URL for a local expert tier. Repeatable.") + p.add_argument( + "--max-tokens", + type=int, + default=3000, + help="Keep high: the model reasons a lot before it emits the delegation.", + ) + p.add_argument( + "--local-endpoint", + action="append", + default=[], + help="MODEL_ID=URL for a local expert tier. Repeatable.", + ) return p.parse_args(argv) @@ -68,8 +86,10 @@ def main(argv=None): temperature=args.temperature, ) - print(f"orchestrator: {args.model} @ {args.endpoint} " - f"(temp={args.temperature}, max_turns={args.max_turns})") + print( + f"orchestrator: {args.model} @ {args.endpoint} " + f"(temp={args.temperature}, max_turns={args.max_turns})" + ) print("type a question and hit enter. Ctrl-D or 'quit' to exit.\n") try: @@ -86,7 +106,9 @@ def main(argv=None): started = time.time() full = backend.generate_full( - prompt, model=args.model, temperature=args.temperature, + prompt, + model=args.model, + temperature=args.temperature, max_tokens=args.max_tokens, ) elapsed = time.time() - started @@ -96,10 +118,12 @@ def main(argv=None): continue print(f"\norch> {full.get('content', '').strip()}\n") - print(f" [turns={full.get('turn_count', '?')} " - f"tool_calls={full.get('tool_calls', '?')} " - f"cost=${full.get('cost_usd', 0.0):.4f} " - f"{elapsed:.1f}s]\n") + print( + f" [turns={full.get('turn_count', '?')} " + f"tool_calls={full.get('tool_calls', '?')} " + f"cost=${full.get('cost_usd', 0.0):.4f} " + f"{elapsed:.1f}s]\n" + ) finally: backend.close() return 0 diff --git a/scripts/orchestrator/eval_orchestrator.py b/scripts/orchestrator/eval_orchestrator.py index 86673d30..af86ae39 100644 --- a/scripts/orchestrator/eval_orchestrator.py +++ b/scripts/orchestrator/eval_orchestrator.py @@ -258,9 +258,7 @@ def main(argv: Optional[List[str]] = None) -> int: local_endpoints: Dict[str, str] = {} for pair in args.local_endpoint: if "=" not in pair: - raise SystemExit( - f"--local-endpoint expects MODEL_ID=URL, got: {pair!r}" - ) + raise SystemExit(f"--local-endpoint expects MODEL_ID=URL, got: {pair!r}") model_id, url = pair.split("=", 1) local_endpoints[model_id.strip()] = url.strip() @@ -307,6 +305,7 @@ def main(argv: Optional[List[str]] = None) -> int: served_id, served_path = args.orchestrator_model, None try: import urllib.request + req = urllib.request.Request( args.orchestrator_endpoint.rstrip("/") + "/models", headers={"Authorization": f"Bearer {args.orchestrator_api_key}"}, diff --git a/scripts/orchestrator/format_eval_sample.py b/scripts/orchestrator/format_eval_sample.py index e81c9718..39f9ac17 100644 --- a/scripts/orchestrator/format_eval_sample.py +++ b/scripts/orchestrator/format_eval_sample.py @@ -30,6 +30,7 @@ Usage: ... --input --all # one .txt per record ... --input --all --only-wrong # only incorrect samples """ + from __future__ import annotations import argparse @@ -81,8 +82,10 @@ def build_question_map(benchmark: str, n: int, seed: int = 42) -> dict: ds.load(max_samples=n, seed=seed) return {r.record_id: r.problem for r in ds.iter_records()} except Exception as exc: # noqa: BLE001 - never fail the render over this - print(f" [warn] could not load dataset for {benchmark!r}: " - f"{type(exc).__name__}: {exc}") + print( + f" [warn] could not load dataset for {benchmark!r}: " + f"{type(exc).__name__}: {exc}" + ) return {} @@ -134,7 +137,11 @@ def parse_scoring(meta: dict) -> dict: def format_record(rec: dict, question: str | None) -> str: parsed = parse_scoring(rec.get("scoring_metadata")) is_correct = rec.get("is_correct") - verdict = "CORRECT" if is_correct else ("INCORRECT" if is_correct is False else "UNSCORED") + verdict = ( + "CORRECT" + if is_correct + else ("INCORRECT" if is_correct is False else "UNSCORED") + ) def _fmt(v, fmt): try: @@ -142,14 +149,16 @@ def format_record(rec: dict, question: str | None) -> str: except (ValueError, TypeError): return str(v) - top = " · ".join([ - str(rec.get("record_id", "?")), - str(rec.get("benchmark", "?")), - verdict, - f"score {_fmt(rec.get('score'), '{:.3f}')}", - f"lat {_fmt(rec.get('latency_seconds'), '{:.1f}')}s", - f"cost ${_fmt(rec.get('cost_usd'), '{:.4f}')}", - ]) + top = " · ".join( + [ + str(rec.get("record_id", "?")), + str(rec.get("benchmark", "?")), + verdict, + f"score {_fmt(rec.get('score'), '{:.3f}')}", + f"lat {_fmt(rec.get('latency_seconds'), '{:.1f}')}s", + f"cost ${_fmt(rec.get('cost_usd'), '{:.4f}')}", + ] + ) parts = [top] err = rec.get("error") @@ -164,7 +173,13 @@ def format_record(rec: dict, question: str | None) -> str: parts.append(_banner("GOLD")) gold = parsed["gold"] - parts.append(_indent(str(gold) if gold not in (None, "") else "(gold unavailable — see SCORING below)")) + parts.append( + _indent( + str(gold) + if gold not in (None, "") + else "(gold unavailable — see SCORING below)" + ) + ) if parsed["kind"] == "judge" and parsed["judge"]: parts.append(_banner("JUDGE")) @@ -172,7 +187,9 @@ def format_record(rec: dict, question: str | None) -> str: if not parsed["clean"]: parts.append(_banner("SCORING (raw)")) - parts.append(_indent(json.dumps(rec.get("scoring_metadata"), indent=2, default=str))) + parts.append( + _indent(json.dumps(rec.get("scoring_metadata"), indent=2, default=str)) + ) return "\n".join(parts).rstrip() + "\n" @@ -182,10 +199,17 @@ def main(argv=None): description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter ) p.add_argument("--input", required=True, help="Graded *_orchestrator.jsonl file.") - p.add_argument("--out-dir", default="results/formatted", help="Where the .txt files go.") - p.add_argument("--seed", type=int, default=42, help="Subset seed used at eval time.") - p.add_argument("--only-wrong", action="store_true", - help="Render only incorrect/unscored samples.") + p.add_argument( + "--out-dir", default="results/formatted", help="Where the .txt files go." + ) + p.add_argument( + "--seed", type=int, default=42, help="Subset seed used at eval time." + ) + p.add_argument( + "--only-wrong", + action="store_true", + help="Render only incorrect/unscored samples.", + ) g = p.add_mutually_exclusive_group() g.add_argument("--n", type=int, default=1, help="Format the first N records.") g.add_argument("--lines", help="Comma-separated 1-indexed line numbers.") @@ -213,7 +237,9 @@ def main(argv=None): if args.only_wrong: idxs = [i for i in idxs if not records[i].get("is_correct")] - benchmark = _normalize_benchmark(records[0].get("benchmark", src.stem.split("_")[0])) + benchmark = _normalize_benchmark( + records[0].get("benchmark", src.stem.split("_")[0]) + ) # Load the whole subset so any selected id resolves (the eval-time subset is # the first len(records) of seed=42). qmap = build_question_map(benchmark, n=len(records), seed=args.seed) @@ -230,8 +256,10 @@ def main(argv=None): for w in written: print(w) - print(f"\nwrote {len(written)} file(s) to {out_dir}/" - + (" (dataset unavailable — questions omitted)" if not qmap else "")) + print( + f"\nwrote {len(written)} file(s) to {out_dir}/" + + (" (dataset unavailable — questions omitted)" if not qmap else "") + ) return 0 diff --git a/scripts/orchestrator/make_splits.py b/scripts/orchestrator/make_splits.py index a37fa120..be641c41 100644 --- a/scripts/orchestrator/make_splits.py +++ b/scripts/orchestrator/make_splits.py @@ -27,6 +27,7 @@ Deterministic (seed 42). --pool data/orchestrator/sft/qwen_clean_0711.jsonl \ --pool data/orchestrator/sft/gemma_clean_0711.jsonl """ + import argparse import json import os @@ -36,7 +37,12 @@ from collections import defaultdict from datetime import datetime from pathlib import Path -OUT = Path("data/orchestrator/sft") +# Where the orchestrator data tree lives. Repo-relative by default so a fresh +# clone works out of the box; set OJ_DATA_ROOT to keep the data OUT of the git +# checkout (this workspace points it at ~/experiments/orchestrator/data, so a +# stray `git reset` can't touch hundreds of GB of generations). +DATA_ROOT = Path(os.getenv("OJ_DATA_ROOT", "data/orchestrator")) +OUT = DATA_ROOT / "sft" SEED = 42 OVERFIT_N = 100 # Orchestrator that generated each pool — stamped onto every row so provenance @@ -46,16 +52,32 @@ ORCH_MODEL = {"qwen": "qwen3.5-9b", "gemma": "gemma-4-26b"} def main() -> int: - ap = argparse.ArgumentParser(description=__doc__, - formatter_class=argparse.RawDescriptionHelpFormatter) - ap.add_argument("--name", required=True, - help="split family: qwen | gemma | pooled (drives the filename)") - ap.add_argument("--pool", action="append", required=True, metavar="PATH", - help="clean pool jsonl; repeat to merge (use with --name pooled)") - ap.add_argument("--holdout-frac", type=float, default=0.15, - help="fraction held out, domain-stratified (default 0.15)") - ap.add_argument("--date", default=datetime.now().strftime("%m%d"), - help="date tag in the filename (default: today, MMDD)") + ap = argparse.ArgumentParser( + description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter + ) + ap.add_argument( + "--name", + required=True, + help="split family: qwen | gemma | pooled (drives the filename)", + ) + ap.add_argument( + "--pool", + action="append", + required=True, + metavar="PATH", + help="clean pool jsonl; repeat to merge (use with --name pooled)", + ) + ap.add_argument( + "--holdout-frac", + type=float, + default=0.15, + help="fraction held out, domain-stratified (default 0.15)", + ) + ap.add_argument( + "--date", + default=datetime.now().strftime("%m%d"), + help="date tag in the filename (default: today, MMDD)", + ) ap.add_argument("--out", type=Path, default=OUT) args = ap.parse_args() @@ -94,8 +116,11 @@ def main() -> int: args.out.mkdir(parents=True, exist_ok=True) written = {} - for split, data in (("train", train), ("holdout", holdout), - (f"overfit{OVERFIT_N}", overfit)): + for split, data in ( + ("train", train), + ("holdout", holdout), + (f"overfit{OVERFIT_N}", overfit), + ): f = args.out / f"{args.name}_{split}_{args.date}.jsonl" f.write_text("".join(json.dumps(r) + "\n" for r in data)) written[split] = f @@ -109,8 +134,10 @@ def main() -> int: from upload_to_braintrust import autoupload autoupload( - [f"{written['train']}={args.name}_train_{args.date}", - f"{written['holdout']}={args.name}_holdout_{args.date}"], + [ + f"{written['train']}={args.name}_train_{args.date}", + f"{written['holdout']}={args.name}_holdout_{args.date}", + ], run_label=os.getenv("OJ_RUN_LABEL", f"{args.name}_splits_{args.date}"), description=f"{args.name} orchestrator-SFT splits (make_splits.py)", ) diff --git a/scripts/orchestrator/render_sft_data.py b/scripts/orchestrator/render_sft_data.py index c3670fd1..e4309417 100644 --- a/scripts/orchestrator/render_sft_data.py +++ b/scripts/orchestrator/render_sft_data.py @@ -64,7 +64,11 @@ def _fmt_tool_calls(content: str, anon_map=None) -> str: lines = [f"-> {_call_label(name, anon_map)}"] if isinstance(args, dict): for k, v in args.items(): - val = v if isinstance(v, str) else json.dumps(v, ensure_ascii=False, indent=2) + val = ( + v + if isinstance(v, str) + else json.dumps(v, ensure_ascii=False, indent=2) + ) indented = "\n".join(" " + ln for ln in val.splitlines()) lines.append(f" {k}:\n{indented}") blocks.append("\n".join(lines)) @@ -136,7 +140,8 @@ def _system_header(records: List[dict]) -> str: if str(t.get("role", "")).lower() == "system": return ( "=" * 80 + "\nSHARED SYSTEM PROMPT + TOOL CATALOG " - "(identical for every record below)\n" + "=" * 80 + "(identical for every record below)\n" + + "=" * 80 + f"\n{t.get('content', '').strip()}\n" ) return "" @@ -186,27 +191,42 @@ def _html_turns(record: dict) -> str: if role == "system": continue if role == "user": - out.append(f'
user
{_esc(content.strip())}
') + out.append( + f'
user
{_esc(content.strip())}
' + ) elif role in ("tool", "function"): name = _esc(t.get("name", "tool")) out.append( f'
tool result · {name} ' - f'({len(content)} chars)
{_esc(content.strip())}
' + f"({len(content)} chars)
{_esc(content.strip())}
" ) elif role == "assistant": prose = _TOOL_CALL_RE.sub("", content).strip() if prose: - out.append(f'
assistant
{_esc(prose)}
') + out.append( + f'
assistant
{_esc(prose)}
' + ) for name, args in _parse_tool_calls(content): real = anon_map.get(name) - label = (f'{_esc(str(name))} → {_esc(str(real))}' - if real else _esc(str(name))) + label = ( + f'{_esc(str(name))} → {_esc(str(real))}' + if real + else _esc(str(name)) + ) rows = [f'
→ {label}
'] if isinstance(args, dict): for k, v in args.items(): - val = v if isinstance(v, str) else json.dumps(v, ensure_ascii=False, indent=2) - rows.append(f'
{_esc(str(k))}
{_esc(val)}
') - out.append(f'
tool call
{"".join(rows)}
') + val = ( + v + if isinstance(v, str) + else json.dumps(v, ensure_ascii=False, indent=2) + ) + rows.append( + f'
{_esc(str(k))}
{_esc(val)}
' + ) + out.append( + f'
tool call
{"".join(rows)}
' + ) return "".join(out) @@ -215,16 +235,26 @@ def _html_record(record: dict, index: int) -> str: pills = [f'{_esc(str(record.get("domain")))}'] correct = record.get("correct") if correct is not None: - pills.append('correct' if correct - else 'wrong') + pills.append( + 'correct' + if correct + else 'wrong' + ) if record.get("kept"): pills.append('kept') - meta = (f'cost ${m.get("cost_usd")} · {m.get("tokens")} tok · ' - f'{m.get("num_tool_calls")} calls · {m.get("num_turns")} turns') - summary = (f'#{index}{_esc(str(record.get("task_id")))}' - + "".join(pills) + meta) - return (f'
{summary}' - f'
{_html_turns(record)}
') + meta = ( + f'cost ${m.get("cost_usd")} · {m.get("tokens")} tok · ' + f"{m.get('num_tool_calls')} calls · {m.get('num_turns')} turns" + ) + summary = ( + f"#{index}{_esc(str(record.get('task_id')))}" + + "".join(pills) + + meta + ) + return ( + f'
{summary}' + f'
{_html_turns(record)}
' + ) def _html_doc(records: List[dict], title: str) -> str: @@ -262,9 +292,7 @@ def render(jsonl_path: str) -> dict: pretty.write_text(json.dumps(records, indent=2, ensure_ascii=False)) txt = src.with_suffix(".txt") - banner = ( - f"ORCHESTRATOR SFT TRANSCRIPT — {src.name}\n{len(records)} records\n\n" - ) + banner = f"ORCHESTRATOR SFT TRANSCRIPT — {src.name}\n{len(records)} records\n\n" sep = f"\n\n{_RULE}\n\n" parts = [banner + _system_header(records)] parts.append(sep.join(_transcript(r, i + 1) for i, r in enumerate(records))) @@ -273,8 +301,12 @@ def render(jsonl_path: str) -> dict: htmlp = src.with_suffix(".html") htmlp.write_text(_html_doc(records, src.name)) - return {"records": len(records), "pretty": str(pretty), - "txt": str(txt), "html": str(htmlp)} + return { + "records": len(records), + "pretty": str(pretty), + "txt": str(txt), + "html": str(htmlp), + } def main(argv: List[str]) -> int: diff --git a/scripts/orchestrator/run_sft_fsdp.py b/scripts/orchestrator/run_sft_fsdp.py index 384c016f..af63e791 100644 --- a/scripts/orchestrator/run_sft_fsdp.py +++ b/scripts/orchestrator/run_sft_fsdp.py @@ -24,6 +24,7 @@ Reuses the conversation->tokens + assistant-only masking from sft_tokenize.py. On these no-NVLink L40S, the accelerate/FSDP NCCL env (NCCL_P2P_DISABLE=1 etc.) must be set by the launcher. """ + from __future__ import annotations import argparse @@ -39,8 +40,14 @@ from pathlib import Path sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) from sft_tokenize import build_example # noqa: E402 -MODEL_EXPERTS = {"gpt_5_5", "claude_opus_4_8", "qwen3_5_9b", "qwen3_6_27b_fp8", - "qwen3_5_122b_a10b_fp8", "qwen3_5_397b_a17b_fp8"} +MODEL_EXPERTS = { + "gpt_5_5", + "claude_opus_4_8", + "qwen3_5_9b", + "qwen3_6_27b_fp8", + "qwen3_5_122b_a10b_fp8", + "qwen3_5_397b_a17b_fp8", +} def record_is_correct(r: dict) -> bool: @@ -53,7 +60,9 @@ def record_is_correct(r: dict) -> bool: def record_routed_to_expert(r: dict) -> bool: """True if the trajectory called a model expert (resolving anon labels).""" amap = r.get("metrics", {}).get("anon_map", {}) - names = re.findall(r'"name"\s*:\s*"([a-z0-9_]+)"', json.dumps(r.get("conversations", []))) + names = re.findall( + r'"name"\s*:\s*"([a-z0-9_]+)"', json.dumps(r.get("conversations", [])) + ) for n in names: real = amap.get(n, n) if real in MODEL_EXPERTS: @@ -84,7 +93,9 @@ def select(records, variant: str, require_clean: bool = True): elif variant == "correct": base = [r for r in records if record_is_correct(r)] elif variant == "correct_routed": - base = [r for r in records if record_is_correct(r) and record_routed_to_expert(r)] + base = [ + r for r in records if record_is_correct(r) and record_routed_to_expert(r) + ] else: raise ValueError(f"unknown variant {variant}") if not require_clean: @@ -94,14 +105,25 @@ def select(records, variant: str, require_clean: bool = True): def main() -> int: p = argparse.ArgumentParser() - p.add_argument("--data", action="append", required=True, help="JSONL(s) (repeatable)") - p.add_argument("--val-data", default=None, - help="Held-out JSONL for val-loss (excluded from --data). Eval'd per epoch.") - p.add_argument("--variant", choices=["all", "correct", "correct_routed"], default="correct") - p.add_argument("--require-clean", action=argparse.BooleanOptionalAction, default=True, - help="Drop records whose `clean` flag is False (bloated / garbled / " - "unrouted). Rows missing the flag (legacy data) are kept. " - "Use --no-require-clean to disable.") + p.add_argument( + "--data", action="append", required=True, help="JSONL(s) (repeatable)" + ) + p.add_argument( + "--val-data", + default=None, + help="Held-out JSONL for val-loss (excluded from --data). Eval'd per epoch.", + ) + p.add_argument( + "--variant", choices=["all", "correct", "correct_routed"], default="correct" + ) + p.add_argument( + "--require-clean", + action=argparse.BooleanOptionalAction, + default=True, + help="Drop records whose `clean` flag is False (bloated / garbled / " + "unrouted). Rows missing the flag (legacy data) are kept. " + "Use --no-require-clean to disable.", + ) p.add_argument("--model", default="Qwen/Qwen3.5-9B") p.add_argument("--out", required=True) p.add_argument("--epochs", type=float, default=3.0) @@ -109,9 +131,12 @@ def main() -> int: p.add_argument("--grad-accum", type=int, default=8) p.add_argument("--lr", type=float, default=1e-5) p.add_argument("--max-seq", type=int, default=8192) - p.add_argument("--supervise-last-only", action="store_true", - help="Legacy: supervise ONLY the final assistant turn. Default " - "(off) supervises every assistant turn incl. routing.") + p.add_argument( + "--supervise-last-only", + action="store_true", + help="Legacy: supervise ONLY the final assistant turn. Default " + "(off) supervises every assistant turn incl. routing.", + ) p.add_argument("--warmup-ratio", type=float, default=0.03) p.add_argument("--seed", type=int, default=42) p.add_argument("--wandb-project", default="orchestrator-sft") @@ -120,10 +145,11 @@ def main() -> int: import random from datetime import timedelta + import torch - from torch.utils.data import DataLoader from accelerate import Accelerator - from accelerate.utils import set_seed, InitProcessGroupKwargs + from accelerate.utils import InitProcessGroupKwargs, set_seed + from torch.utils.data import DataLoader from transformers import AutoModelForCausalLM, AutoTokenizer set_seed(args.seed) @@ -133,7 +159,10 @@ def main() -> int: pg_timeout_min = int(os.environ.get("SFT_PG_TIMEOUT_MIN", "30")) accelerator = Accelerator( gradient_accumulation_steps=args.grad_accum, - kwargs_handlers=[InitProcessGroupKwargs(timeout=timedelta(minutes=pg_timeout_min))]) + kwargs_handlers=[ + InitProcessGroupKwargs(timeout=timedelta(minutes=pg_timeout_min)) + ], + ) is_main = accelerator.is_main_process def log(m): @@ -153,14 +182,20 @@ def main() -> int: if args.require_clean: pre_clean = select(records, args.variant, require_clean=False) n_unclean = len(pre_clean) - len(sel) - log(f"require_clean=True: dropped {n_unclean} unclean rows " - f"({len(pre_clean)}->{len(sel)}); use --no-require-clean to keep them") + log( + f"require_clean=True: dropped {n_unclean} unclean rows " + f"({len(pre_clean)}->{len(sel)}); use --no-require-clean to keep them" + ) else: log("require_clean=False: NOT filtering on `clean` flag") examples = [] for r in sel: - ex = build_example(tok, r.get("conversations", []), args.max_seq, - supervise_all_turns=not args.supervise_last_only) + ex = build_example( + tok, + r.get("conversations", []), + args.max_seq, + supervise_all_turns=not args.supervise_last_only, + ) if ex: examples.append(ex) log(f"built {len(examples)} training examples") @@ -179,7 +214,9 @@ def main() -> int: am.append([1] * len(b["input_ids"]) + [0] * pad) return (torch.tensor(ii), torch.tensor(lab), torch.tensor(am)) - dl = DataLoader(examples, batch_size=args.batch_size, shuffle=True, collate_fn=collate) + dl = DataLoader( + examples, batch_size=args.batch_size, shuffle=True, collate_fn=collate + ) # ---- held-out val set (leak-free; excluded from --data upstream) ---- val_dl = None @@ -187,21 +224,32 @@ def main() -> int: vrecs = [json.loads(l) for l in open(args.val_data) if l.strip()] vex = [] for r in vrecs: - ex = build_example(tok, r.get("conversations", []), args.max_seq, - supervise_all_turns=not args.supervise_last_only) + ex = build_example( + tok, + r.get("conversations", []), + args.max_seq, + supervise_all_turns=not args.supervise_last_only, + ) if ex: vex.append(ex) log(f"val: {len(vex)} examples from {args.val_data}") if vex: - val_dl = DataLoader(vex, batch_size=args.batch_size, shuffle=False, collate_fn=collate) + val_dl = DataLoader( + vex, batch_size=args.batch_size, shuffle=False, collate_fn=collate + ) model = AutoModelForCausalLM.from_pretrained( - args.model, dtype=torch.bfloat16, trust_remote_code=True, - attn_implementation="sdpa") + args.model, + dtype=torch.bfloat16, + trust_remote_code=True, + attn_implementation="sdpa", + ) if os.environ.get("SFT_NO_GRAD_CKPT") == "1": log("gradient checkpointing DISABLED (SFT_NO_GRAD_CKPT=1)") else: - model.gradient_checkpointing_enable(gradient_checkpointing_kwargs={"use_reentrant": False}) + model.gradient_checkpointing_enable( + gradient_checkpointing_kwargs={"use_reentrant": False} + ) model.config.use_cache = False opt = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=0.0) @@ -234,17 +282,25 @@ def main() -> int: def lr_at(s): if s < warmup: return s / warmup - return 0.5 * (1 + math.cos(math.pi * (s - warmup) / max(1, total_steps - warmup))) + return 0.5 * ( + 1 + math.cos(math.pi * (s - warmup) / max(1, total_steps - warmup)) + ) use_wandb = False if is_main: try: import wandb - base_name = args.wandb_name or f"{Path(args.model).name}-{args.variant}-fsdp" + + base_name = ( + args.wandb_name or f"{Path(args.model).name}-{args.variant}-fsdp" + ) run_name = f"{base_name}-{time.strftime('%m%d-%H%M', time.gmtime())}" - wandb.init(project=args.wandb_project, - name=run_name, - config=vars(args) | {"n_examples": len(examples), "total_steps": total_steps}) + wandb.init( + project=args.wandb_project, + name=run_name, + config=vars(args) + | {"n_examples": len(examples), "total_steps": total_steps}, + ) use_wandb = True except Exception as e: log(f"wandb off ({e})") @@ -262,9 +318,11 @@ def main() -> int: p.mkdir(parents=True, exist_ok=True) tok.save_pretrained(str(p)) unwrapped.save_pretrained( - str(p), is_main_process=is_main, + str(p), + is_main_process=is_main, save_function=accelerator.save, - state_dict=accelerator.get_state_dict(model)) + state_dict=accelerator.get_state_dict(model), + ) if is_main: log(f"checkpoint saved -> {p}") @@ -319,9 +377,11 @@ def main() -> int: out_dir.mkdir(parents=True, exist_ok=True) tok.save_pretrained(str(out_dir)) unwrapped.save_pretrained( - str(out_dir), is_main_process=is_main, + str(out_dir), + is_main_process=is_main, save_function=accelerator.save, - state_dict=accelerator.get_state_dict(model)) + state_dict=accelerator.get_state_dict(model), + ) if is_main: log(f"saved -> {out_dir}") if use_wandb: @@ -329,8 +389,17 @@ def main() -> int: wandb.finish() except Exception: pass - print("FSDP_SFT_DONE " + json.dumps({"variant": args.variant, "steps": gstep, - "wall_s": round(time.time() - t0, 1), "out": str(out_dir)})) + print( + "FSDP_SFT_DONE " + + json.dumps( + { + "variant": args.variant, + "steps": gstep, + "wall_s": round(time.time() - t0, 1), + "out": str(out_dir), + } + ) + ) return 0 diff --git a/scripts/orchestrator/sft_tokenize.py b/scripts/orchestrator/sft_tokenize.py index 100687d1..8b5e6666 100644 --- a/scripts/orchestrator/sft_tokenize.py +++ b/scripts/orchestrator/sft_tokenize.py @@ -6,6 +6,7 @@ the path we actually use) doesn't depend on it. ``build_example`` tokenizes one ``conversations`` record and returns ``input_ids`` + ``labels`` with everything but the supervised assistant turns masked to -100. """ + from __future__ import annotations from typing import Any, Dict, List, Optional @@ -102,8 +103,9 @@ def build_example( if len(msgs) < 2 or msgs[-1]["role"] != "assistant": return None try: - full_text = tok.apply_chat_template(msgs, tokenize=False, - add_generation_prompt=False) + full_text = tok.apply_chat_template( + msgs, tokenize=False, add_generation_prompt=False + ) full = tok(full_text, add_special_tokens=False)["input_ids"] except Exception: return None @@ -114,7 +116,7 @@ def build_example( if spans: selected = spans if supervise_all_turns else spans[-1:] labels = [-100] * len(full) - for (s, e) in selected: + for s, e in selected: for k in range(s, e): labels[k] = full[k] last_start = selected[-1][0] @@ -122,8 +124,9 @@ def build_example( # Non-ChatML template: proven last-turn boundary via longest-common-prefix # (robust to templates whose add_generation_prompt emits extra preamble). try: - prompt_text = tok.apply_chat_template(msgs[:-1], tokenize=False, - add_generation_prompt=True) + prompt_text = tok.apply_chat_template( + msgs[:-1], tokenize=False, add_generation_prompt=True + ) prompt = tok(prompt_text, add_special_tokens=False)["input_ids"] except Exception: return None diff --git a/scripts/orchestrator/upload_to_braintrust.py b/scripts/orchestrator/upload_to_braintrust.py index e4d06d20..899915cb 100644 --- a/scripts/orchestrator/upload_to_braintrust.py +++ b/scripts/orchestrator/upload_to_braintrust.py @@ -34,6 +34,7 @@ Programmatic (used by the pipeline auto-upload hook in make_splits.py): autoupload(["data/orchestrator/sft/qwen_train_0707.jsonl=qwen_train_0707"], run_label="...") """ + import argparse import json import logging @@ -82,6 +83,7 @@ def _dataset_desc(dataset): from openjarvis.learning.intelligence.orchestrator.sft_data.reject_sample import ( dataset_description, ) + return dataset_description(dataset) except Exception: return "" @@ -89,10 +91,16 @@ def _dataset_desc(dataset): def _git_sha(): try: - return subprocess.check_output( - ["git", "rev-parse", "--short", "HEAD"], - cwd=Path(__file__).resolve().parent, stderr=subprocess.DEVNULL, - ).decode().strip() or None + return ( + subprocess.check_output( + ["git", "rev-parse", "--short", "HEAD"], + cwd=Path(__file__).resolve().parent, + stderr=subprocess.DEVNULL, + ) + .decode() + .strip() + or None + ) except Exception: return None @@ -113,7 +121,17 @@ def _model_short(rows, path): def _split_of(path): stem = Path(path).stem.lower() - for s in ("holdout", "train", "clean", "overfit", "partial", "8k", "4k", "2k", "1k"): + for s in ( + "holdout", + "train", + "clean", + "overfit", + "partial", + "8k", + "4k", + "2k", + "1k", + ): if s in stem: return s return "data" @@ -129,8 +147,16 @@ def _config_from_env(): for env_key, meta_key, cast in [ ("OJ_CFG_TEMPERATURE", "temperature", float), ("OJ_CFG_MAX_TURNS", "max_turns", int), - ("OJ_CFG_ANONYMIZE", "anonymize", lambda v: v.strip().lower() in ("1", "true", "yes")), - ("OJ_CFG_REJECTION_ONLY", "rejection_only", lambda v: v.strip().lower() in ("1", "true", "yes")), + ( + "OJ_CFG_ANONYMIZE", + "anonymize", + lambda v: v.strip().lower() in ("1", "true", "yes"), + ), + ( + "OJ_CFG_REJECTION_ONLY", + "rejection_only", + lambda v: v.strip().lower() in ("1", "true", "yes"), + ), ]: v = os.getenv(env_key) if v not in (None, ""): @@ -148,7 +174,7 @@ def _row(rec): asst = [c["content"] for c in convs if c["role"] == "assistant"] fa = asst[-1] if asst else "" m = list(_FA.finditer(fa)) - model_answer = fa[m[-1].end():].strip() if m else fa.strip() + model_answer = fa[m[-1].end() :].strip() if m else fa.strip() # Gold reference the verifier graded against (stamped at generation time as # `gold_answer`). `expected` in Braintrust is the GOLD so the UI shows # gold-vs-model; the model's own answer goes to metadata.model_answer. @@ -158,7 +184,9 @@ def _row(rec): routed = [] for c in convs: if c["role"] == "assistant": - for t in re.findall(r"(\{.*?\})", c["content"], re.DOTALL): + for t in re.findall( + r"(\{.*?\})", c["content"], re.DOTALL + ): try: real = am.get(json.loads(t).get("name")) if real: @@ -171,14 +199,18 @@ def _row(rec): correct = rec.get("correct") clean = rec.get("clean") kept = rec.get("kept") - tags = [t for t in [ - gen_model, - f"domain:{domain}" if domain else None, - "correct" if correct else "incorrect", - "clean" if clean else "dirty", - "kept" if kept else "dropped", - "gold" if has_gold else "no_gold", - ] if t] + tags = [ + t + for t in [ + gen_model, + f"domain:{domain}" if domain else None, + "correct" if correct else "incorrect", + "clean" if clean else "dirty", + "kept" if kept else "dropped", + "gold" if has_gold else "no_gold", + ] + if t + ] return { "input": question, "expected": gold, @@ -191,7 +223,8 @@ def _row(rec): "area": rec.get("area"), "difficulty": rec.get("difficulty") or None, "dataset": rec.get("dataset"), - "dataset_description": rec.get("dataset_description") or _dataset_desc(rec.get("dataset")), + "dataset_description": rec.get("dataset_description") + or _dataset_desc(rec.get("dataset")), "subsector": rec.get("subsector"), "task_id": rec.get("task_id"), "correct": correct, @@ -222,7 +255,9 @@ def _routed_dist(rows): for cc in r.get("conversations", []): if cc.get("role") != "assistant": continue - for t in re.findall(r"(\{.*?\})", cc.get("content", ""), re.DOTALL): + for t in re.findall( + r"(\{.*?\})", cc.get("content", ""), re.DOTALL + ): try: real = am.get(json.loads(t).get("name")) if real: @@ -257,7 +292,7 @@ def _dataset_metadata(rows, path, run_label, gen_model): "source_datasets": ["GeneralThought-430K-filtered", "OpenThoughts3-1.2M"], "task_mix": "balanced (GeneralThought + OpenThoughts code/math/science)", "filter": "correct + clean (rejects file-write echoes, garble, " - "reasoning-degeneration, truncated tails, essays/markdown dumps)", + "reasoning-degeneration, truncated tails, essays/markdown dumps)", "leak_free": "train/holdout disjoint by task_id", "n_total": n, "n_correct": sum(1 for r in rows if r.get("correct")), @@ -277,8 +312,16 @@ def _dataset_metadata(rows, path, run_label, gen_model): return {k: v for k, v in meta.items() if v is not None} -def upload_dataset(path, name=None, *, project_id=None, project=None, - run_label=None, gen_model=None, description=""): +def upload_dataset( + path, + name=None, + *, + project_id=None, + project=None, + run_label=None, + gen_model=None, + description="", +): """Upload one JSONL file as a Braintrust dataset. Returns (name, n_rows, url). Target project: ``project`` (name) if given, else ``project_id`` / @@ -288,24 +331,33 @@ def upload_dataset(path, name=None, *, project_id=None, project=None, rows = [json.loads(l) for l in open(path) if l.strip()] name = name or _default_name(rows, path) - gen_model = gen_model or os.getenv("OJ_GEN_MODEL") or \ - next((r.get("gen_model") for r in rows if r.get("gen_model")), None) + gen_model = ( + gen_model + or os.getenv("OJ_GEN_MODEL") + or next((r.get("gen_model") for r in rows if r.get("gen_model")), None) + ) run_label = run_label or os.getenv("OJ_RUN_LABEL") or name - init_kwargs = {"name": name, "description": description or None, - "metadata": _dataset_metadata(rows, path, run_label, gen_model)} + init_kwargs = { + "name": name, + "description": description or None, + "metadata": _dataset_metadata(rows, path, run_label, gen_model), + } if project: init_kwargs["project"] = project else: init_kwargs["project_id"] = project_id or os.getenv( - "OJ_BRAINTRUST_PROJECT_ID", DEFAULT_PROJECT_ID) + "OJ_BRAINTRUST_PROJECT_ID", DEFAULT_PROJECT_ID + ) ds = braintrust.init_dataset(**init_kwargs) for r in rows: ds.insert(**_row(r)) ds.flush() summ = ds.summarize() - url = getattr(summ, "dataset_url", None) or (project or init_kwargs.get("project_id")) + url = getattr(summ, "dataset_url", None) or ( + project or init_kwargs.get("project_id") + ) return name, len(rows), url @@ -313,7 +365,13 @@ def autoupload(specs, *, run_label=None, gen_model=None, description=""): """No-op-safe wrapper for pipeline hooks. Honors OJ_BRAINTRUST_AUTOUPLOAD (default ON) and never raises: any failure (missing key/pkg, network) is logged and swallowed so the data pipeline is never broken by telemetry.""" - if os.getenv("OJ_BRAINTRUST_AUTOUPLOAD", "1").strip().lower() in ("0", "false", "no", "off", ""): + if os.getenv("OJ_BRAINTRUST_AUTOUPLOAD", "1").strip().lower() in ( + "0", + "false", + "no", + "off", + "", + ): logger.info("[braintrust] autoupload disabled (OJ_BRAINTRUST_AUTOUPLOAD)") return if not os.getenv("BRAINTRUST_API_KEY"): @@ -327,22 +385,37 @@ def autoupload(specs, *, run_label=None, gen_model=None, description=""): for spec in specs: path, _, name = spec.partition("=") try: - nm, n, url = upload_dataset(path, name or None, run_label=run_label, - gen_model=gen_model, description=description) + nm, n, url = upload_dataset( + path, + name or None, + run_label=run_label, + gen_model=gen_model, + description=description, + ) logger.info("[braintrust] uploaded %s: %d rows -> %s", nm, n, url) print(f"[braintrust] uploaded {nm}: {n} rows -> {url}") except Exception as exc: - logger.warning("[braintrust] autoupload FAILED for %s (%s) — continuing", path, exc) - print(f"[braintrust] autoupload FAILED for {path} ({exc}) — pipeline unaffected") + logger.warning( + "[braintrust] autoupload FAILED for %s (%s) — continuing", path, exc + ) + print( + f"[braintrust] autoupload FAILED for {path} ({exc}) — pipeline unaffected" + ) def main(): logging.basicConfig(level=logging.INFO, format="%(message)s") p = argparse.ArgumentParser() - p.add_argument("--project", default=None, - help="Target project by NAME (overrides OJ_BRAINTRUST_PROJECT_ID).") - p.add_argument("--project-id", default=None, - help="Target project by id (default: OJ_BRAINTRUST_PROJECT_ID or research).") + p.add_argument( + "--project", + default=None, + help="Target project by NAME (overrides OJ_BRAINTRUST_PROJECT_ID).", + ) + p.add_argument( + "--project-id", + default=None, + help="Target project by id (default: OJ_BRAINTRUST_PROJECT_ID or research).", + ) p.add_argument("--run-label", default=None) p.add_argument("--gen-model", default=None, help="Specific gen model id override.") p.add_argument("--description", default="") @@ -351,9 +424,14 @@ def main(): for spec in args.specs: path, _, name = spec.partition("=") nm, n, url = upload_dataset( - path, name or None, project_id=args.project_id, project=args.project, - run_label=args.run_label, gen_model=args.gen_model, - description=args.description) + path, + name or None, + project_id=args.project_id, + project=args.project, + run_label=args.run_label, + gen_model=args.gen_model, + description=args.description, + ) print(f"[braintrust] {nm}: {n} rows -> {url}") diff --git a/src/openjarvis/agents/hybrid/_prices.py b/src/openjarvis/agents/hybrid/_prices.py index a37172c9..9cf61a8d 100644 --- a/src/openjarvis/agents/hybrid/_prices.py +++ b/src/openjarvis/agents/hybrid/_prices.py @@ -38,8 +38,8 @@ PRICES: dict[str, tuple[float, float]] = { # OpenRouter instead of self-hosted vLLM. ESTIMATES — no public list price # exists yet for these Qwen3.5/3.6 builds; scaled by active-param size. # VERIFY against openrouter.ai before trusting the cost-aware reward numbers. - "qwen/qwen3.5-9b": (0.05, 0.10), - "qwen/qwen3.6-27b": (0.10, 0.30), + "qwen/qwen3.5-9b": (0.05, 0.10), + "qwen/qwen3.6-27b": (0.10, 0.30), "qwen/qwen3.5-122b-a10b": (0.20, 0.60), "qwen/qwen3.5-397b-a17b": (0.40, 1.20), } diff --git a/src/openjarvis/agents/hybrid/expert_registry.py b/src/openjarvis/agents/hybrid/expert_registry.py index 65b806f1..6ad87030 100644 --- a/src/openjarvis/agents/hybrid/expert_registry.py +++ b/src/openjarvis/agents/hybrid/expert_registry.py @@ -26,7 +26,7 @@ from __future__ import annotations import os import random import re -from dataclasses import dataclass, field +from dataclasses import dataclass from typing import Dict, List, Optional from openjarvis.agents.hybrid._prices import PRICES @@ -57,8 +57,14 @@ CATEGORY_SPECIALIZED = "specialized_model" # ``openjarvis-tool`` bridge, dispatched in ``unified.make_dispatch`` via the # OpenJarvis ToolExecutor rather than ``_call_worker``). VALID_BACKENDS = ( - "vllm", "openai", "anthropic", "gemini", "openrouter", - "anthropic-web-search", "tavily-search", "modal-python", + "vllm", + "openai", + "anthropic", + "gemini", + "openrouter", + "anthropic-web-search", + "tavily-search", + "modal-python", "openjarvis-tool", ) @@ -268,12 +274,22 @@ def anonymize_tools(tools, rng): while tag in anon_to_real: tag = "model_" + "".join(rng.choice(_ANON_ALPHABET) for _ in range(4)) anon_to_real[tag] = t.name - bits = [b for b in (_class_hint(t), _size_hint(t.model), _cost_tier_hint(t)) if b] - experts.append(replace( - t, name=tag, - summary="Another model — " + ", ".join(bits) + ". Send it a sub-question.", - category="model", hide_cost=True, - )) + bits = [ + b + for b in (_class_hint(t), _size_hint(t.model), _cost_tier_hint(t)) + if b + ] + experts.append( + replace( + t, + name=tag, + summary="Another model — " + + ", ".join(bits) + + ". Send it a sub-question.", + category="model", + hide_cost=True, + ) + ) else: basics.append(t) # Shuffle WITHIN the experts to kill per-expert position bias, but keep all @@ -300,70 +316,135 @@ def default_catalog( # ---- generalist / frontier models (one named tool per model VERSION) ---- for model, summary, lat in [ - ("gpt-5", - "Frontier generalist (GPT-5). Strongest reasoning across domains.", 30.0), - ("gpt-5-mini", - "Mid-tier generalist (GPT-5-mini). Solid reasoning, much cheaper.", 15.0), - ("gpt-4o", - "Fast generalist (GPT-4o). Good for simple steps and formatting.", 8.0), - ("claude-opus-4-7", - "Frontier generalist (Claude Opus 4.7). Strong long-horizon reasoning.", 26.0), - ("claude-sonnet-4-6", - "Strong generalist (Claude Sonnet 4.6). Balanced cost/capability.", 15.0), - ("gemini-2.5-pro", - "Frontier generalist (Gemini 2.5 Pro). Strong multimodal reasoning.", 20.0), - ("gemini-2.5-flash", - "Cheap fast generalist (Gemini 2.5 Flash).", 8.0), - ("meta-llama/llama-3.3-70b-instruct", - "Open generalist (Llama-3.3-70B). Decent general knowledge, low cost.", 10.0), - ("qwen/qwen3-32b", - "Open generalist (Qwen3-32B). Strong math/science reasoning, low cost.", 9.0), + ( + "gpt-5", + "Frontier generalist (GPT-5). Strongest reasoning across domains.", + 30.0, + ), + ( + "gpt-5-mini", + "Mid-tier generalist (GPT-5-mini). Solid reasoning, much cheaper.", + 15.0, + ), + ( + "gpt-4o", + "Fast generalist (GPT-4o). Good for simple steps and formatting.", + 8.0, + ), + ( + "claude-opus-4-7", + "Frontier generalist (Claude Opus 4.7). Strong long-horizon reasoning.", + 26.0, + ), + ( + "claude-sonnet-4-6", + "Strong generalist (Claude Sonnet 4.6). Balanced cost/capability.", + 15.0, + ), + ( + "gemini-2.5-pro", + "Frontier generalist (Gemini 2.5 Pro). Strong multimodal reasoning.", + 20.0, + ), + ("gemini-2.5-flash", "Cheap fast generalist (Gemini 2.5 Flash).", 8.0), + ( + "meta-llama/llama-3.3-70b-instruct", + "Open generalist (Llama-3.3-70B). Decent general knowledge, low cost.", + 10.0, + ), + ( + "qwen/qwen3-32b", + "Open generalist (Qwen3-32B). Strong math/science reasoning, low cost.", + 9.0, + ), ]: - ep = ("openai" if model.startswith("gpt") else - "anthropic" if model.startswith("claude") else - "gemini" if model.startswith("gemini") else "openrouter") + ep = ( + "openai" + if model.startswith("gpt") + else "anthropic" + if model.startswith("claude") + else "gemini" + if model.startswith("gemini") + else "openrouter" + ) pi, po = _price(model) - cat.append(ExpertTool( - name=_tool_name(model), kind=KIND_MODEL, backend_type=ep, summary=summary, - model=model, price_in=pi, price_out=po, latency_s=lat, - category=CATEGORY_GENERALIST, - )) + cat.append( + ExpertTool( + name=_tool_name(model), + kind=KIND_MODEL, + backend_type=ep, + summary=summary, + model=model, + price_in=pi, + price_out=po, + latency_s=lat, + category=CATEGORY_GENERALIST, + ) + ) # ---- specialized: code ---- coder = "qwen/qwen-2.5-coder-32b-instruct" pi, po = _price(coder) - cat.append(ExpertTool( - name=_tool_name(coder), kind=KIND_MODEL, backend_type="openrouter", - summary="Specialized code model (Qwen2.5-Coder-32B). Writes/debugs code.", - model=coder, price_in=pi, price_out=po, latency_s=9.0, - category=CATEGORY_SPECIALIZED, - )) + cat.append( + ExpertTool( + name=_tool_name(coder), + kind=KIND_MODEL, + backend_type="openrouter", + summary="Specialized code model (Qwen2.5-Coder-32B). Writes/debugs code.", + model=coder, + price_in=pi, + price_out=po, + latency_s=9.0, + category=CATEGORY_SPECIALIZED, + ) + ) # ---- local backbone as a tool (on-device vLLM), if served ---- # Named after the actual served model (faithful "one named tool per model"), # not a generic "local_model" — e.g. "qwen3-8b" -> tool "qwen3_8b". if local_model and local_endpoint: - cat.append(ExpertTool( - name=_tool_name(local_model), kind=KIND_MODEL, backend_type="vllm", - summary=(f"On-device open model ({local_model}) served locally. Cheap " - "and private; good for extraction, formatting, arithmetic on " - "given data."), - model=local_model, base_url=local_endpoint, - price_in=0.0, price_out=0.0, latency_s=2.0, - category=CATEGORY_GENERALIST, - )) + cat.append( + ExpertTool( + name=_tool_name(local_model), + kind=KIND_MODEL, + backend_type="vllm", + summary=( + f"On-device open model ({local_model}) served locally. Cheap " + "and private; good for extraction, formatting, arithmetic on " + "given data." + ), + model=local_model, + base_url=local_endpoint, + price_in=0.0, + price_out=0.0, + latency_s=2.0, + category=CATEGORY_GENERALIST, + ) + ) # ---- basic tools ---- - cat.append(ExpertTool( - name="web_search", kind=KIND_WEB_SEARCH, backend_type="tavily-search", - summary="Web search (Tavily). Use for facts that need a live lookup.", - model="tavily", latency_s=8.0, category=CATEGORY_BASIC, - )) - cat.append(ExpertTool( - name="code_interpreter", kind=KIND_CODE, backend_type="modal-python", - summary="Python sandbox. Execute code and return stdout/stderr.", - model="modal-python", latency_s=6.0, category=CATEGORY_BASIC, - )) + cat.append( + ExpertTool( + name="web_search", + kind=KIND_WEB_SEARCH, + backend_type="tavily-search", + summary="Web search (Tavily). Use for facts that need a live lookup.", + model="tavily", + latency_s=8.0, + category=CATEGORY_BASIC, + ) + ) + cat.append( + ExpertTool( + name="code_interpreter", + kind=KIND_CODE, + backend_type="modal-python", + summary="Python sandbox. Execute code and return stdout/stderr.", + model="modal-python", + latency_s=6.0, + category=CATEGORY_BASIC, + ) + ) return cat @@ -407,19 +488,29 @@ def _openjarvis_basic_tools() -> List[ExpertTool]: "calculator", summary="Evaluate an arithmetic / math expression and return the result.", params=obj( - {"expression": {"type": "string", - "description": "Math expression to evaluate."}}, + { + "expression": { + "type": "string", + "description": "Math expression to evaluate.", + } + }, ["expression"], ), latency_s=1.0, ), openjarvis_tool( "shell_exec", - summary=("Run a shell command and return its stdout/stderr. Critical " - "for terminal / TerminalBench-style tasks."), + summary=( + "Run a shell command and return its stdout/stderr. Critical " + "for terminal / TerminalBench-style tasks." + ), params=obj( - {"command": {"type": "string", - "description": "Shell command to execute."}}, + { + "command": { + "type": "string", + "description": "Shell command to execute.", + } + }, ["command"], ), latency_s=4.0, @@ -428,7 +519,12 @@ def _openjarvis_basic_tools() -> List[ExpertTool]: "file_read", summary="Read the contents of a file at the given path.", params=obj( - {"path": {"type": "string", "description": "Path of the file to read."}}, + { + "path": { + "type": "string", + "description": "Path of the file to read.", + } + }, ["path"], ), latency_s=1.0, @@ -437,8 +533,13 @@ def _openjarvis_basic_tools() -> List[ExpertTool]: "file_write", summary="Write content to a file at the given path.", params=obj( - {"path": {"type": "string", "description": "Path of the file to write."}, - "content": {"type": "string", "description": "Content to write."}}, + { + "path": { + "type": "string", + "description": "Path of the file to write.", + }, + "content": {"type": "string", "description": "Content to write."}, + }, ["path", "content"], ), latency_s=1.0, @@ -450,8 +551,10 @@ def _openjarvis_basic_tools() -> List[ExpertTool]: "type": "object", "properties": { "url": {"type": "string", "description": "Request URL."}, - "method": {"type": "string", - "description": "HTTP method (GET, POST, ...). Default GET."}, + "method": { + "type": "string", + "description": "HTTP method (GET, POST, ...). Default GET.", + }, }, "required": ["url"], }, @@ -459,56 +562,85 @@ def _openjarvis_basic_tools() -> List[ExpertTool]: ), openjarvis_tool( "think", - summary=("Record a private reasoning step (scratchpad). No external " - "effect; use to plan before acting on hard reasoning tasks."), + summary=( + "Record a private reasoning step (scratchpad). No external " + "effect; use to plan before acting on hard reasoning tasks." + ), params=obj( - {"thought": {"type": "string", - "description": "Your reasoning or thought process."}}, + { + "thought": { + "type": "string", + "description": "Your reasoning or thought process.", + } + }, ["thought"], ), latency_s=0.5, ), openjarvis_tool( "apply_patch", - summary=("Apply a unified-diff patch to a file. Use to edit code for " - "terminal / SWE-style tasks."), + summary=( + "Apply a unified-diff patch to a file. Use to edit code for " + "terminal / SWE-style tasks." + ), params=obj( - {"patch": {"type": "string", - "description": "The unified diff patch text to apply."}, - "path": {"type": "string", - "description": "Target file path (auto-detected from the " - "patch header if omitted)."}}, + { + "patch": { + "type": "string", + "description": "The unified diff patch text to apply.", + }, + "path": { + "type": "string", + "description": "Target file path (auto-detected from the " + "patch header if omitted).", + }, + }, ["patch"], ), latency_s=2.0, ), openjarvis_tool( "pdf_extract", - summary=("Extract text from a PDF file. Use for GAIA-style tasks with " - "PDF attachments."), + summary=( + "Extract text from a PDF file. Use for GAIA-style tasks with " + "PDF attachments." + ), params=obj( - {"file_path": {"type": "string", - "description": "Path to the PDF file."}, - "pages": {"type": "string", - "description": "Page range, e.g. '1-5' or '1,3,5'. " - "Omit for all pages."}}, + { + "file_path": { + "type": "string", + "description": "Path to the PDF file.", + }, + "pages": { + "type": "string", + "description": "Page range, e.g. '1-5' or '1,3,5'. " + "Omit for all pages.", + }, + }, ["file_path"], ), latency_s=3.0, ), openjarvis_tool( "db_query", - summary=("Run a SQL query against a SQLite/Postgres database and return " - "rows. Read-only by default."), + summary=( + "Run a SQL query against a SQLite/Postgres database and return " + "rows. Read-only by default." + ), params=obj( - {"query": {"type": "string", - "description": "SQL query to execute."}, - "db_path": {"type": "string", - "description": "Path to a SQLite DB file. Defaults to " - "in-memory."}, - "read_only": {"type": "boolean", - "description": "Restrict to SELECT/EXPLAIN/PRAGMA. " - "Default: true."}}, + { + "query": {"type": "string", "description": "SQL query to execute."}, + "db_path": { + "type": "string", + "description": "Path to a SQLite DB file. Defaults to " + "in-memory.", + }, + "read_only": { + "type": "boolean", + "description": "Restrict to SELECT/EXPLAIN/PRAGMA. " + "Default: true.", + }, + }, ["query"], ), latency_s=3.0, @@ -540,10 +672,14 @@ _CLOUD_FRONTIER_MODELS = ( # Neutral, uniform summaries (no capability ranking) so the orchestrator # doesn't just pick whichever model is labelled "strongest" — routing should # be learned from the reward, not hand-labelled here. - ("gpt-5.5", "openai", "openai/gpt-5.5", - "Expert model (GPT-5.5).", 30.0), - ("claude-opus-4-8", "anthropic", "anthropic/claude-opus-4.8", - "Expert model (Claude Opus 4.8).", 26.0), + ("gpt-5.5", "openai", "openai/gpt-5.5", "Expert model (GPT-5.5).", 30.0), + ( + "claude-opus-4-8", + "anthropic", + "anthropic/claude-opus-4.8", + "Expert model (Claude Opus 4.8).", + 26.0, + ), ) @@ -570,16 +706,26 @@ def _model_tool( backend = ( model_backends.get(canonical) or ("vllm" if canonical in local_endpoints else None) - or (native_backend if native_backend in ("openai", "anthropic", "gemini") - else "openrouter") + or ( + native_backend + if native_backend in ("openai", "anthropic", "gemini") + else "openrouter" + ) ) name = _tool_name(canonical) if backend == "vllm": # Self-hosted: free per the cost model. return ExpertTool( - name=name, kind=KIND_MODEL, backend_type="vllm", summary=summary, - model=canonical, base_url=local_endpoints.get(canonical), - price_in=0.0, price_out=0.0, latency_s=lat, category=category, + name=name, + kind=KIND_MODEL, + backend_type="vllm", + summary=summary, + model=canonical, + base_url=local_endpoints.get(canonical), + price_in=0.0, + price_out=0.0, + latency_s=lat, + category=category, ) if backend == "openrouter": slug = openrouter_slugs.get(canonical, or_slug) @@ -587,15 +733,29 @@ def _model_tool( if (pi, po) == (0.0, 0.0): # fall back to the canonical id's price pi, po = _price(canonical) return ExpertTool( - name=name, kind=KIND_MODEL, backend_type="openrouter", summary=summary, - model=slug, base_url=None, price_in=pi, price_out=po, - latency_s=lat, category=category, + name=name, + kind=KIND_MODEL, + backend_type="openrouter", + summary=summary, + model=slug, + base_url=None, + price_in=pi, + price_out=po, + latency_s=lat, + category=category, ) # native provider API (openai / anthropic / gemini) pi, po = _price(canonical) return ExpertTool( - name=name, kind=KIND_MODEL, backend_type=backend, summary=summary, - model=canonical, price_in=pi, price_out=po, latency_s=lat, category=category, + name=name, + kind=KIND_MODEL, + backend_type=backend, + summary=summary, + model=canonical, + price_in=pi, + price_out=po, + latency_s=lat, + category=category, ) @@ -633,7 +793,12 @@ def orchestrator_catalog( # match is skipped from the catalog. Used to temporarily drop unreliable # experts (e.g. OpenRouter giants during a provider outage) without editing # the registry — unset the var to restore them. - _excl = {s.strip().lower() for s in os.environ.get("OJ_EXCLUDE_EXPERTS", "").split(",") if s.strip()} + _excl = { + s.strip().lower() + for s in os.environ.get("OJ_EXCLUDE_EXPERTS", "").split(",") + if s.strip() + } + def _excluded(canonical: str) -> bool: c = canonical.lower() return any(x in c for x in _excl) @@ -642,38 +807,63 @@ def orchestrator_catalog( for canonical, native_backend, or_slug, summary, lat in _CLOUD_FRONTIER_MODELS: if _excluded(canonical): continue - cat.append(_model_tool( - canonical, native_backend=native_backend, or_slug=or_slug, - summary=summary, lat=lat, category=CATEGORY_CLOUD_FRONTIER, - local_endpoints=local_endpoints, model_backends=model_backends, - openrouter_slugs=openrouter_slugs, - )) + cat.append( + _model_tool( + canonical, + native_backend=native_backend, + or_slug=or_slug, + summary=summary, + lat=lat, + category=CATEGORY_CLOUD_FRONTIER, + local_endpoints=local_endpoints, + model_backends=model_backends, + openrouter_slugs=openrouter_slugs, + ) + ) # ---- open-source models (OpenRouter by default; vLLM when an endpoint or # a model_backends override is supplied) ---- for canonical, or_slug in _LOCAL_OSS_MODELS: if _excluded(canonical): continue - cat.append(_model_tool( - canonical, native_backend="vllm", or_slug=or_slug, - summary=f"Expert model ({canonical}).", - lat=4.0, category=CATEGORY_LOCAL_OSS, - local_endpoints=local_endpoints, model_backends=model_backends, - openrouter_slugs=openrouter_slugs, - )) + cat.append( + _model_tool( + canonical, + native_backend="vllm", + or_slug=or_slug, + summary=f"Expert model ({canonical}).", + lat=4.0, + category=CATEGORY_LOCAL_OSS, + local_endpoints=local_endpoints, + model_backends=model_backends, + openrouter_slugs=openrouter_slugs, + ) + ) if include_tools: # ---- basic tools ---- - cat.append(ExpertTool( - name="web_search", kind=KIND_WEB_SEARCH, backend_type="tavily-search", - summary="Web search (Tavily). Use for facts that need a live lookup.", - model="tavily", latency_s=8.0, category=CATEGORY_BASIC, - )) - cat.append(ExpertTool( - name="code_interpreter", kind=KIND_CODE, backend_type="modal-python", - summary="Python sandbox. Execute code and return stdout/stderr.", - model="modal-python", latency_s=6.0, category=CATEGORY_BASIC, - )) + cat.append( + ExpertTool( + name="web_search", + kind=KIND_WEB_SEARCH, + backend_type="tavily-search", + summary="Web search (Tavily). Use for facts that need a live lookup.", + model="tavily", + latency_s=8.0, + category=CATEGORY_BASIC, + ) + ) + cat.append( + ExpertTool( + name="code_interpreter", + kind=KIND_CODE, + backend_type="modal-python", + summary="Python sandbox. Execute code and return stdout/stderr.", + model="modal-python", + latency_s=6.0, + category=CATEGORY_BASIC, + ) + ) cat.extend(_openjarvis_basic_tools()) return cat @@ -736,13 +926,19 @@ def sample_tool_config( for t in chosen: if t.kind == KIND_MODEL and (t.price_in or t.price_out): f = rng.uniform(1.0 - price_jitter, 1.0 + price_jitter) - jittered.append(ExpertTool( - name=t.name, kind=t.kind, backend_type=t.backend_type, - summary=t.summary, model=t.model, base_url=t.base_url, - price_in=round(t.price_in * f, 4), - price_out=round(t.price_out * f, 4), - latency_s=t.latency_s, - )) + jittered.append( + ExpertTool( + name=t.name, + kind=t.kind, + backend_type=t.backend_type, + summary=t.summary, + model=t.model, + base_url=t.base_url, + price_in=round(t.price_in * f, 4), + price_out=round(t.price_out * f, 4), + latency_s=t.latency_s, + ) + ) else: jittered.append(t) return jittered diff --git a/src/openjarvis/agents/hybrid/toolorchestra/agent.py b/src/openjarvis/agents/hybrid/toolorchestra/agent.py index b2fc963a..d63fd427 100644 --- a/src/openjarvis/agents/hybrid/toolorchestra/agent.py +++ b/src/openjarvis/agents/hybrid/toolorchestra/agent.py @@ -47,10 +47,8 @@ Qwen if vLLM up, plus a web-search tool via Anthropic, Opus 4.7, gpt-5-mini). """ - from __future__ import annotations -import json import shutil import tempfile from pathlib import Path @@ -62,41 +60,40 @@ from openjarvis.agents.hybrid._prices import PRICES from openjarvis.agents.hybrid.mini_swe_agent import ( _clone_repo, _extract_diff, - run_swe_agent_loop, ) -from openjarvis.core.registry import AgentRegistry - -from openjarvis.agents.hybrid.toolorchestra.prompts import ( - FORCE_FINAL_PROMPT, - ORCHESTRATOR_SYS, - RL_ALL_TOOLS, - RL_ORCHESTRATOR_SYS, - RL_TOOLS_SPEC, +from openjarvis.agents.hybrid.toolorchestra.clients import ( + _call_orchestrator_with_tool_calls, ) from openjarvis.agents.hybrid.toolorchestra.experts import ( _PAPER_CODER_OPENROUTER, _expert_for, _paper_expert_for, ) -from openjarvis.agents.hybrid.toolorchestra.sandbox import ( - _call_modal_python, - _extract_first_python_block, -) -from openjarvis.agents.hybrid.toolorchestra.clients import ( - _call_orchestrator_with_tool_calls, -) from openjarvis.agents.hybrid.toolorchestra.parsing import ( _build_user_prompt, _extract_final_answer_text, _parse_action, _parse_rl_tool_call, ) +from openjarvis.agents.hybrid.toolorchestra.prompts import ( + FORCE_FINAL_PROMPT, + ORCHESTRATOR_SYS, + RL_ALL_TOOLS, + RL_ORCHESTRATOR_SYS, + RL_TOOLS_SPEC, +) +from openjarvis.agents.hybrid.toolorchestra.sandbox import ( + _call_modal_python, + _extract_first_python_block, +) from openjarvis.agents.hybrid.toolorchestra.workers import ( + _TOOLORCH_SEARCH_TYPES, _call_worker, - _default_pool, _resolve_worker_pool, _swe_call_worker, ) +from openjarvis.core.registry import AgentRegistry + @AgentRegistry.register("toolorchestra") class ToolOrchestraAgent(LocalCloudAgent): @@ -182,20 +179,24 @@ class ToolOrchestraAgent(LocalCloudAgent): ) shared_workdir: Optional[Path] = None if swe_mode: - shared_workdir = Path(tempfile.mkdtemp( - prefix=f"toolorch-swe-{task_meta.get('task_id','x')}-" - )) + shared_workdir = Path( + tempfile.mkdtemp( + prefix=f"toolorch-swe-{task_meta.get('task_id', 'x')}-" + ) + ) try: _clone_repo(task_meta["repo"], task_meta["base_commit"], shared_workdir) except Exception: shutil.rmtree(shared_workdir, ignore_errors=True) raise - self.record_trace_event({ - "kind": "toolorchestra_swe_workdir", - "workdir": str(shared_workdir), - "repo": task_meta["repo"], - "base_commit": task_meta["base_commit"], - }) + self.record_trace_event( + { + "kind": "toolorchestra_swe_workdir", + "workdir": str(shared_workdir), + "repo": task_meta["repo"], + "base_commit": task_meta["base_commit"], + } + ) # try/finally guards ``shared_workdir`` against exceptions raised # anywhere in the turn loop, the worker calls, the fallback, or @@ -233,15 +234,22 @@ class ToolOrchestraAgent(LocalCloudAgent): cost += self.cost_usd(self._cloud_model, o_in, o_out) action = _parse_action(text) - history.append({ - "role": "orchestrator", "turn": turn, "raw": text, "action": action, - }) - self.record_trace_event({ - "kind": "toolorchestra_action", - "turn": turn, - "action": action, - "raw": text, - }) + history.append( + { + "role": "orchestrator", + "turn": turn, + "raw": text, + "action": action, + } + ) + self.record_trace_event( + { + "kind": "toolorchestra_action", + "turn": turn, + "action": action, + "raw": text, + } + ) if action is None: parse_failures += 1 @@ -265,12 +273,21 @@ class ToolOrchestraAgent(LocalCloudAgent): continue worker = workers[wid] if swe_mode and shared_workdir is not None: - (w_text, w_in, w_out, is_local, extra_cost, - n_searches, bash_turns) = ( - _swe_call_worker( - worker, str(w_input), cfg, task_meta, - shared_workdir, turn, - ) + ( + w_text, + w_in, + w_out, + is_local, + extra_cost, + n_searches, + bash_turns, + ) = _swe_call_worker( + worker, + str(w_input), + cfg, + task_meta, + shared_workdir, + turn, ) tool_calls += bash_turns else: @@ -284,17 +301,19 @@ class ToolOrchestraAgent(LocalCloudAgent): cost += self.cost_usd(worker["model"], w_in, w_out) + extra_cost n_web_searches_total += n_searches tool_calls += n_searches - history.append({ - "role": "worker", - "turn": turn, - "worker_id": wid, - "worker_name": worker["name"], - "worker_model": worker["model"], - "output": w_text, - "tokens_in": w_in, - "tokens_out": w_out, - "n_web_searches": n_searches, - }) + history.append( + { + "role": "worker", + "turn": turn, + "worker_id": wid, + "worker_name": worker["name"], + "worker_model": worker["model"], + "output": w_text, + "tokens_in": w_in, + "tokens_out": w_out, + "n_web_searches": n_searches, + } + ) continue # Unknown action kind — treat as parse failure. parse_failures += 1 @@ -306,17 +325,22 @@ class ToolOrchestraAgent(LocalCloudAgent): # Search workers are excluded — they answer fact-lookup # questions, not synthesis. non_search = [ - w for w in workers if w.get("type") != "anthropic-web-search" + w for w in workers if w.get("type") not in _TOOLORCH_SEARCH_TYPES ] or workers worker = max( non_search, key=lambda w: PRICES.get(w.get("model", ""), (0.0, 0.0))[1], ) if swe_mode and shared_workdir is not None: - (ans, w_in, w_out, is_local, extra_cost, _, - bash_turns) = _swe_call_worker( - worker, question, cfg, task_meta, - shared_workdir, max_turns + 1, + (ans, w_in, w_out, is_local, extra_cost, _, bash_turns) = ( + _swe_call_worker( + worker, + question, + cfg, + task_meta, + shared_workdir, + max_turns + 1, + ) ) tool_calls += bash_turns else: @@ -328,17 +352,19 @@ class ToolOrchestraAgent(LocalCloudAgent): else: tokens_cloud += w_in + w_out cost += self.cost_usd(worker["model"], w_in, w_out) + extra_cost - history.append({ - "role": "worker", - "turn": max_turns + 1, - "worker_id": worker["id"], - "worker_name": worker["name"], - "worker_model": worker["model"], - "output": ans, - "tokens_in": w_in, - "tokens_out": w_out, - "fallback": True, - }) + history.append( + { + "role": "worker", + "turn": max_turns + 1, + "worker_id": worker["id"], + "worker_name": worker["name"], + "worker_model": worker["model"], + "output": ans, + "tokens_in": w_in, + "tokens_out": w_out, + "fallback": True, + } + ) final_answer = ans # In SWE mode, the authoritative output is the working-tree diff — @@ -348,7 +374,8 @@ class ToolOrchestraAgent(LocalCloudAgent): if patch.strip(): final_answer = ( f"{final_answer}\n\n```diff\n{patch}```" - if final_answer else f"```diff\n{patch}```" + if final_answer + else f"```diff\n{patch}```" ) meta = { @@ -418,20 +445,24 @@ class ToolOrchestraAgent(LocalCloudAgent): ) shared_workdir: Optional[Path] = None if swe_mode: - shared_workdir = Path(tempfile.mkdtemp( - prefix=f"toolorch-rl-swe-{task_meta.get('task_id','x')}-" - )) + shared_workdir = Path( + tempfile.mkdtemp( + prefix=f"toolorch-rl-swe-{task_meta.get('task_id', 'x')}-" + ) + ) try: _clone_repo(task_meta["repo"], task_meta["base_commit"], shared_workdir) except Exception: shutil.rmtree(shared_workdir, ignore_errors=True) raise - self.record_trace_event({ - "kind": "toolorchestra_rl_swe_workdir", - "workdir": str(shared_workdir), - "repo": task_meta["repo"], - "base_commit": task_meta["base_commit"], - }) + self.record_trace_event( + { + "kind": "toolorchestra_rl_swe_workdir", + "workdir": str(shared_workdir), + "repo": task_meta["repo"], + "base_commit": task_meta["base_commit"], + } + ) # ``context_str`` mirrors the upstream's running context — accumulates # search documents and code/exec snippets across turns. We keep this @@ -480,40 +511,53 @@ class ToolOrchestraAgent(LocalCloudAgent): temperature=orch_temp, tools=RL_TOOLS_SPEC, ) - self.record_trace_event({ - "kind": "vllm", - "role": "orchestrator", - "model": orch_model, - "endpoint": orch_endpoint, - "system": RL_ORCHESTRATOR_SYS, - "user": user, - "response": text, - "tool_calls": [ - { - "id": getattr(tc, "id", None), - "type": getattr(tc, "type", None), - "function": { - "name": getattr(getattr(tc, "function", None), "name", None), - "arguments": getattr(getattr(tc, "function", None), "arguments", None), - }, - } - for tc in (sdk_tool_calls or []) - ], - "tokens_in": o_in, - "tokens_out": o_out, - }) + self.record_trace_event( + { + "kind": "vllm", + "role": "orchestrator", + "model": orch_model, + "endpoint": orch_endpoint, + "system": RL_ORCHESTRATOR_SYS, + "user": user, + "response": text, + "tool_calls": [ + { + "id": getattr(tc, "id", None), + "type": getattr(tc, "type", None), + "function": { + "name": getattr( + getattr(tc, "function", None), "name", None + ), + "arguments": getattr( + getattr(tc, "function", None), "arguments", None + ), + }, + } + for tc in (sdk_tool_calls or []) + ], + "tokens_in": o_in, + "tokens_out": o_out, + } + ) tokens_local += o_in + o_out action = _parse_rl_tool_call(text, sdk_tool_calls) - history.append({ - "role": "orchestrator", "turn": turn, "raw": text, "action": action, - }) - self.record_trace_event({ - "kind": "toolorchestra_rl_action", - "turn": turn, - "action": action, - "raw": text, - }) + history.append( + { + "role": "orchestrator", + "turn": turn, + "raw": text, + "action": action, + } + ) + self.record_trace_event( + { + "kind": "toolorchestra_rl_action", + "turn": turn, + "action": action, + "raw": text, + } + ) if action is None: parse_failures += 1 @@ -526,8 +570,10 @@ class ToolOrchestraAgent(LocalCloudAgent): slot = args.get("model", "") # Validate against the upstream tool/arg schema. - valid = name in RL_ALL_TOOLS and isinstance(slot, str) and ( - slot in RL_ALL_TOOLS[name]["model"] + valid = ( + name in RL_ALL_TOOLS + and isinstance(slot, str) + and (slot in RL_ALL_TOOLS[name]["model"]) ) if not valid: parse_failures += 1 @@ -548,8 +594,11 @@ class ToolOrchestraAgent(LocalCloudAgent): # framing). if paper_mode: worker = _paper_expert_for( - slot, self._local_model, self._local_endpoint, - self._cloud_model, self._cloud_endpoint, + slot, + self._local_model, + self._local_endpoint, + self._cloud_model, + self._cloud_endpoint, ) # In paper mode, `enhance_reasoning` is always the coder # specialist regardless of the orchestrator's chosen tier. @@ -563,7 +612,10 @@ class ToolOrchestraAgent(LocalCloudAgent): } else: worker = _expert_for( - slot, self._local_model, self._local_endpoint, self._cloud_model, + slot, + self._local_model, + self._local_endpoint, + self._cloud_model, self._cloud_endpoint, ) @@ -619,13 +671,25 @@ class ToolOrchestraAgent(LocalCloudAgent): # bash_turns=0; vllm/anthropic-typed workers run the loop. bash_turns = 0 if swe_mode and shared_workdir is not None and name != "search": - (w_text, w_in, w_out, is_local, extra_cost, - n_searches, bash_turns) = _swe_call_worker( - worker, w_input, cfg, task_meta, shared_workdir, turn, + ( + w_text, + w_in, + w_out, + is_local, + extra_cost, + n_searches, + bash_turns, + ) = _swe_call_worker( + worker, + w_input, + cfg, + task_meta, + shared_workdir, + turn, ) else: - w_text, w_in, w_out, is_local, extra_cost, n_searches = _call_worker( - worker, w_input, cfg + w_text, w_in, w_out, is_local, extra_cost, n_searches = ( + _call_worker(worker, w_input, cfg) ) if is_local: tokens_local += w_in + w_out @@ -644,13 +708,13 @@ class ToolOrchestraAgent(LocalCloudAgent): # when no python block is found. modal_exec_output: Optional[str] = None modal_exec_rc: Optional[int] = None - if (paper_mode and name == "enhance_reasoning" - and not swe_mode): + if paper_mode and name == "enhance_reasoning" and not swe_mode: code = _extract_first_python_block(w_text) if code: timeout_s = int(cfg.get("modal_python_timeout_s", 60)) modal_exec_output, modal_exec_rc = _call_modal_python( - code, timeout_s=timeout_s, + code, + timeout_s=timeout_s, ) tool_calls += 1 w_text = ( @@ -658,27 +722,29 @@ class ToolOrchestraAgent(LocalCloudAgent): f"(rc={modal_exec_rc})]\n{modal_exec_output}" ) - history.append({ - "role": "worker", - "turn": turn, - "tool": name, - "slot": slot, - "worker_model": worker["model"], - "worker_type": worker["type"], - "output": w_text, - "tokens_in": w_in, - "tokens_out": w_out, - "n_web_searches": n_searches, - "bash_turns": bash_turns, - "modal_exec_rc": modal_exec_rc, - }) + history.append( + { + "role": "worker", + "turn": turn, + "tool": name, + "slot": slot, + "worker_model": worker["model"], + "worker_type": worker["type"], + "output": w_text, + "tokens_in": w_in, + "tokens_out": w_out, + "n_web_searches": n_searches, + "bash_turns": bash_turns, + "modal_exec_rc": modal_exec_rc, + } + ) # Update accumulated context for the next turn. if name == "search": # Treat the search worker's response as a document. doc_list.append(w_text) ctx_docs = "\n\n".join( - f"Doc {i+1}: {d}" for i, d in enumerate(doc_list) + f"Doc {i + 1}: {d}" for i, d in enumerate(doc_list) ) # Crude char-level cap mirrors the upstream's ~24k token cap. context_str = ("Documents:\n" + ctx_docs)[-24000:] @@ -695,15 +761,23 @@ class ToolOrchestraAgent(LocalCloudAgent): # it can still touch the workdir and emit a diff. expert_fn = _paper_expert_for if paper_mode else _expert_for worker = expert_fn( - "answer-1", self._local_model, self._local_endpoint, - self._cloud_model, self._cloud_endpoint, + "answer-1", + self._local_model, + self._local_endpoint, + self._cloud_model, + self._cloud_endpoint, ) fb_bash_turns = 0 if swe_mode and shared_workdir is not None: - (ans, w_in, w_out, is_local, extra_cost, - _, fb_bash_turns) = _swe_call_worker( - worker, question, cfg, task_meta, - shared_workdir, max_turns + 1, + (ans, w_in, w_out, is_local, extra_cost, _, fb_bash_turns) = ( + _swe_call_worker( + worker, + question, + cfg, + task_meta, + shared_workdir, + max_turns + 1, + ) ) tool_calls += fb_bash_turns else: @@ -715,19 +789,21 @@ class ToolOrchestraAgent(LocalCloudAgent): else: tokens_cloud += w_in + w_out cost += self.cost_usd(worker["model"], w_in, w_out) + extra_cost - history.append({ - "role": "worker", - "turn": max_turns + 1, - "tool": "answer", - "slot": "answer-1", - "worker_model": worker["model"], - "worker_type": worker["type"], - "output": ans, - "tokens_in": w_in, - "tokens_out": w_out, - "bash_turns": fb_bash_turns, - "fallback": True, - }) + history.append( + { + "role": "worker", + "turn": max_turns + 1, + "tool": "answer", + "slot": "answer-1", + "worker_model": worker["model"], + "worker_type": worker["type"], + "output": ans, + "tokens_in": w_in, + "tokens_out": w_out, + "bash_turns": fb_bash_turns, + "fallback": True, + } + ) final_answer = ans # In SWE mode, the authoritative output is the working-tree diff — @@ -737,7 +813,8 @@ class ToolOrchestraAgent(LocalCloudAgent): if patch.strip(): final_answer = ( f"{final_answer}\n\n```diff\n{patch}```" - if final_answer else f"```diff\n{patch}```" + if final_answer + else f"```diff\n{patch}```" ) meta = { diff --git a/src/openjarvis/agents/hybrid/toolorchestra/clients.py b/src/openjarvis/agents/hybrid/toolorchestra/clients.py index d941862b..ec390ac4 100644 --- a/src/openjarvis/agents/hybrid/toolorchestra/clients.py +++ b/src/openjarvis/agents/hybrid/toolorchestra/clients.py @@ -4,6 +4,7 @@ from __future__ import annotations from typing import Any, Dict, List, Tuple + def _call_orchestrator_with_tool_calls( model: str, endpoint: str, diff --git a/src/openjarvis/agents/hybrid/toolorchestra/experts.py b/src/openjarvis/agents/hybrid/toolorchestra/experts.py index a4646a63..da6ff6d8 100644 --- a/src/openjarvis/agents/hybrid/toolorchestra/experts.py +++ b/src/openjarvis/agents/hybrid/toolorchestra/experts.py @@ -16,10 +16,14 @@ _DEFAULT_WEB_SEARCH_MODEL = "claude-haiku-4-5" # so the substitution is deferred until we know the cell's resolved local/cloud # pair. Worker dicts share the schema validated by `_resolve_worker_pool`. -def _expert_for(slot: str, local_model: Optional[str], - local_endpoint: Optional[str], - cloud_model: str, - cloud_endpoint: str = "anthropic") -> Dict[str, Any]: + +def _expert_for( + slot: str, + local_model: Optional[str], + local_endpoint: Optional[str], + cloud_model: str, + cloud_endpoint: str = "anthropic", +) -> Dict[str, Any]: """Map an upstream model slot (`answer-1`, `search-3`, …) to a worker spec. Routing policy: @@ -29,10 +33,23 @@ def _expert_for(slot: str, local_model: Optional[str], cost tier for mid OpenAI calls) - `*-3` (local tier) -> local vLLM (`local_model`) - `answer-math-*` -> same tiers as the numeric suffix - - `search-*` -> always the Anthropic web_search tool (the - upstream uses Tavily; we have web_search) + - `search-*` -> provider-native web search when the cloud + endpoint supports it; otherwise Anthropic """ if slot.startswith("search"): + ep = (cloud_endpoint or "anthropic").lower() + if ep == "openai": + return { + "name": f"search:{slot}", + "type": "openai-web-search", + "model": cloud_model, + } + if ep == "gemini": + return { + "name": f"search:{slot}", + "type": "gemini-web-search", + "model": cloud_model, + } return { "name": f"search:{slot}", "type": "anthropic-web-search", diff --git a/src/openjarvis/agents/hybrid/toolorchestra/parsing.py b/src/openjarvis/agents/hybrid/toolorchestra/parsing.py index f2aee9b4..b36feaa3 100644 --- a/src/openjarvis/agents/hybrid/toolorchestra/parsing.py +++ b/src/openjarvis/agents/hybrid/toolorchestra/parsing.py @@ -9,9 +9,7 @@ from typing import Any, Dict, List, Optional # Regex for ``{...}`` blocks emitted by Orchestrator-8B # when the vLLM tool parser doesn't catch them (e.g. `qwen3_xml` parser on a # hermes-style template). Captures the JSON payload. -_TOOL_CALL_TAG_RE = re.compile( - r"\s*(\{.*?\})\s*", re.DOTALL -) +_TOOL_CALL_TAG_RE = re.compile(r"\s*(\{.*?\})\s*", re.DOTALL) # Fallback for a known failure mode: the SFT'd orchestrator often emits its # delegation as ``\boxed{expert_ab12: }`` (math-data habit bleeding @@ -28,8 +26,8 @@ _TOOL_CALL_TAG_RE = re.compile( # alone. ``group("key")`` is the explicit arg key when the model wrote ``, key:``. _BOXED_DELEGATION_RE = re.compile( r"\\boxed\{\s*([A-Za-z_][\w-]*)\s*" - r"(?:,\s*(?P\w+)\s*)?" # optional ", key" hint (e.g. file_read, path:) - r"(?:query\s*)?" # optional literal "query" word before the colon + r"(?:,\s*(?P\w+)\s*)?" # optional ", key" hint (e.g. file_read, path:) + r"(?:query\s*)?" # optional literal "query" word before the colon r":\s*(.+?)\s*\}\s*$", re.DOTALL, ) @@ -113,7 +111,7 @@ def _strip_fences(s: str) -> str: if s.startswith("```"): first_nl = s.find("\n") if first_nl != -1: - s = s[first_nl + 1:] + s = s[first_nl + 1 :] if s.endswith("```"): s = s[:-3] s = s.strip() diff --git a/src/openjarvis/agents/hybrid/toolorchestra/prompts.py b/src/openjarvis/agents/hybrid/toolorchestra/prompts.py index 752e22c7..6c2e9038 100644 --- a/src/openjarvis/agents/hybrid/toolorchestra/prompts.py +++ b/src/openjarvis/agents/hybrid/toolorchestra/prompts.py @@ -98,8 +98,12 @@ RL_ALL_TOOLS: Dict[str, Dict[str, List[str]]] = { "enhance_reasoning": {"model": ["reasoner-1", "reasoner-2", "reasoner-3"]}, "answer": { "model": [ - "answer-1", "answer-2", "answer-3", "answer-4", - "answer-math-1", "answer-math-2", + "answer-1", + "answer-2", + "answer-3", + "answer-4", + "answer-math-1", + "answer-math-2", ], }, "search": {"model": ["search-1", "search-2", "search-3"]}, diff --git a/src/openjarvis/agents/hybrid/toolorchestra/rollout.py b/src/openjarvis/agents/hybrid/toolorchestra/rollout.py index aa69483e..3bfa312a 100644 --- a/src/openjarvis/agents/hybrid/toolorchestra/rollout.py +++ b/src/openjarvis/agents/hybrid/toolorchestra/rollout.py @@ -80,6 +80,7 @@ def build_system_prompt(specs: List[Dict[str, object]]) -> str: "\n" ) + # Char-level cap on the accumulated conversation (mirrors the paper's ~24k-token # cap) and a per-observation cap so one giant tool dump can't blow it out. _CONTEXT_CAP = 24000 @@ -92,8 +93,10 @@ _SOFT_CALL_CAP = 6 def _trim_history(messages: List[Dict[str, str]]) -> None: """Keep the message history under ``_CONTEXT_CAP`` chars by dropping the OLDEST assistant/tool exchange, never the system prompt or the problem.""" + def total() -> int: return sum(len(m.get("content") or "") for m in messages) + # messages[0]=system, messages[1]=user problem — always keep those two. while total() > _CONTEXT_CAP and len(messages) > 4: del messages[2:4] @@ -145,7 +148,9 @@ def _run_unified_rollout_inner( question: str, tools: List[ExpertTool], *, - call_orchestrator: Callable[..., Tuple[str, List[Tuple[str, Dict[str, object]]], int, int]], + call_orchestrator: Callable[ + ..., Tuple[str, List[Tuple[str, Dict[str, object]]], int, int] + ], dispatch: Callable[[ExpertTool, Dict[str, object]], Tuple[str, float, int, bool]], max_turns: int = 50, system: str = RL_ORCHESTRATOR_SYS, @@ -192,9 +197,15 @@ def _run_unified_rollout_inner( if n_tool_calls == 0 and nudges < 2: nudges += 1 messages.append({"role": "assistant", "content": text or ""}) - messages.append({"role": "user", "content": ( - "Make progress by delegating a concrete sub-question to one of the " - "models now via a .")}) + messages.append( + { + "role": "user", + "content": ( + "Make progress by delegating a concrete sub-question to one of the " + "models now via a ." + ), + } + ) continue # No tool call -> the orchestrator is answering. Terminate. final_answer = (text or "").strip() @@ -206,8 +217,12 @@ def _run_unified_rollout_inner( if name not in by_name: parse_failures += 1 messages.append({"role": "assistant", "content": text or ""}) - messages.append({"role": "user", - "content": f"[invalid tool {name!r} — choose one from the provided tool list]"}) + messages.append( + { + "role": "user", + "content": f"[invalid tool {name!r} — choose one from the provided tool list]", + } + ) if parse_failures >= 2: final_answer = (text or "").strip() break @@ -228,43 +243,68 @@ def _run_unified_rollout_inner( if not _has_input and empty_input_nudges < 3: empty_input_nudges += 1 messages.append({"role": "assistant", "content": text or ""}) - messages.append({"role": "user", "content": ( - f"Your call to {name} had an empty 'input'. Resend the " - " with a non-empty 'input' field containing the " - "concrete sub-question to delegate.")}) + messages.append( + { + "role": "user", + "content": ( + f"Your call to {name} had an empty 'input'. Resend the " + " with a non-empty 'input' field containing the " + "concrete sub-question to delegate." + ), + } + ) continue obs, dcost, dtok, _is_local = dispatch(tool, arguments) cost += float(dcost) tokens += int(dtok) n_tool_calls += 1 - turns.append(UnifiedTurn( - reasoning=text or "", tool_name=name, arguments=dict(arguments), - observation=obs, - )) + turns.append( + UnifiedTurn( + reasoning=text or "", + tool_name=name, + arguments=dict(arguments), + observation=obs, + ) + ) # The model's own action as an assistant turn (the tag in # content, matching the SFT serialization), then the observation as a # distinct `tool` turn the model reads as a tool response. call_content = text or "" if "" not in call_content: - call_content = (call_content + "\n" + tool_call_tag(name, arguments)).strip() + call_content = ( + call_content + "\n" + tool_call_tag(name, arguments) + ).strip() messages.append({"role": "assistant", "content": call_content}) obs_text = obs or "" if len(obs_text) > _OBS_CAP: obs_text = obs_text[:_OBS_CAP] + "\n…[truncated]" messages.append({"role": "tool", "name": name, "content": obs_text}) if n_tool_calls >= _SOFT_CALL_CAP: - messages.append({"role": "user", "content": ( - "You now have enough information from the models. Do NOT call any " - "more tools — reply with your FINAL_ANSWER line only.")}) + messages.append( + { + "role": "user", + "content": ( + "You now have enough information from the models. Do NOT call any " + "more tools — reply with your FINAL_ANSWER line only." + ), + } + ) _trim_history(messages) else: # Hit max_turns with no explicit answer: use the last observation/text. - final_answer = (turns[-1].observation or turns[-1].reasoning).strip() if turns else "" + final_answer = ( + (turns[-1].observation or turns[-1].reasoning).strip() if turns else "" + ) return UnifiedRollout( - turns=turns, final_answer=final_answer, cost_usd=cost, tokens=tokens, - num_tool_calls=n_tool_calls, parse_failures=parse_failures, - anon_map=anon_map, tool_specs=specs, + turns=turns, + final_answer=final_answer, + cost_usd=cost, + tokens=tokens, + num_tool_calls=n_tool_calls, + parse_failures=parse_failures, + anon_map=anon_map, + tool_specs=specs, ) @@ -272,7 +312,9 @@ def run_unified_rollout( question: str, tools: List[ExpertTool], *, - call_orchestrator: Callable[..., Tuple[str, List[Tuple[str, Dict[str, object]]], int, int]], + call_orchestrator: Callable[ + ..., Tuple[str, List[Tuple[str, Dict[str, object]]], int, int] + ], dispatch: Callable[[ExpertTool, Dict[str, object]], Tuple[str, float, int, bool]], max_turns: int = 50, system: str = RL_ORCHESTRATOR_SYS, @@ -285,33 +327,49 @@ def run_unified_rollout( _run_meta, _run_tags = run_context() _fields = dict( input={"question": question}, - metadata={"max_turns": max_turns, "anonymize": anonymize, - "n_tools": len(tools), **_run_meta}, + metadata={ + "max_turns": max_turns, + "anonymize": anonymize, + "n_tools": len(tools), + **_run_meta, + }, ) if _run_tags: _fields["tags"] = _run_tags with span("toolorchestra.rollout", span_type="task", **_fields) as _s: roll = _run_unified_rollout_inner( - question, tools, call_orchestrator=call_orchestrator, dispatch=dispatch, - max_turns=max_turns, system=system, anonymize=anonymize, + question, + tools, + call_orchestrator=call_orchestrator, + dispatch=dispatch, + max_turns=max_turns, + system=system, + anonymize=anonymize, ) _s.log( output=roll.final_answer, - metrics={"cost_usd": roll.cost_usd, "tokens": roll.tokens, - "num_tool_calls": roll.num_tool_calls, - "parse_failures": roll.parse_failures}, - metadata={"n_experts_available": len(roll.anon_map or {}), - "answered": bool(roll.final_answer), - # label -> real model, so every anonymized route span in this - # trace can be decoded back to the model that actually ran. - "anon_map": roll.anon_map or {}}, + metrics={ + "cost_usd": roll.cost_usd, + "tokens": roll.tokens, + "num_tool_calls": roll.num_tool_calls, + "parse_failures": roll.parse_failures, + }, + metadata={ + "n_experts_available": len(roll.anon_map or {}), + "answered": bool(roll.final_answer), + # label -> real model, so every anonymized route span in this + # trace can be decoded back to the model that actually ran. + "anon_map": roll.anon_map or {}, + }, ) return roll def tool_call_tag(name: str, arguments: Dict[str, object]) -> str: """Render a tool call as the ``{...}`` text the model emits.""" - return f"{json.dumps({'name': name, 'arguments': arguments})}" + return ( + f"{json.dumps({'name': name, 'arguments': arguments})}" + ) __all__ = [ diff --git a/src/openjarvis/agents/hybrid/toolorchestra/sandbox.py b/src/openjarvis/agents/hybrid/toolorchestra/sandbox.py index 30956d14..bbe28c6d 100644 --- a/src/openjarvis/agents/hybrid/toolorchestra/sandbox.py +++ b/src/openjarvis/agents/hybrid/toolorchestra/sandbox.py @@ -7,6 +7,7 @@ from typing import Optional, Tuple # ---- Tavily + Modal helpers ------------------------------------------------- + def _call_tavily_search(query: str, max_results: int = 5) -> Tuple[str, int, int]: """One-shot Tavily search. Returns (text, p_tok=0, c_tok=0). @@ -43,7 +44,9 @@ def _call_modal_python(code: str, timeout_s: int = 60) -> Tuple[str, int]: # Python image too. We rely on stdlib only — no extra pip installs. image = modal.Image.debian_slim(python_version="3.12") sb = modal.Sandbox.create( - "python", "-c", code, + "python", + "-c", + code, app=app, image=image, timeout=int(timeout_s), diff --git a/src/openjarvis/agents/hybrid/toolorchestra/tracing.py b/src/openjarvis/agents/hybrid/toolorchestra/tracing.py index f612c44f..86e5935f 100644 --- a/src/openjarvis/agents/hybrid/toolorchestra/tracing.py +++ b/src/openjarvis/agents/hybrid/toolorchestra/tracing.py @@ -17,7 +17,7 @@ from __future__ import annotations import contextlib import logging import os -from typing import Any, Optional +from typing import Any logger = logging.getLogger(__name__) @@ -36,7 +36,9 @@ def _resolve() -> bool: if not _truthy(os.getenv("OJ_BRAINTRUST", "1")): # on by default return False if not os.getenv("BRAINTRUST_API_KEY"): - logger.info("braintrust on-by-default but BRAINTRUST_API_KEY unset — tracing disabled") + logger.info( + "braintrust on-by-default but BRAINTRUST_API_KEY unset — tracing disabled" + ) return False try: import braintrust as _bt @@ -48,9 +50,12 @@ def _resolve() -> bool: _bt.init_logger(project=os.getenv("OJ_BRAINTRUST_PROJECT", "toolorchestra")) _STATE["bt"] = _bt _STATE["enabled"] = True - logger.info("braintrust tracing ENABLED (%s)", - f"project_id={proj_id}" if proj_id - else f"project={os.getenv('OJ_BRAINTRUST_PROJECT', 'toolorchestra')}") + logger.info( + "braintrust tracing ENABLED (%s)", + f"project_id={proj_id}" + if proj_id + else f"project={os.getenv('OJ_BRAINTRUST_PROJECT', 'toolorchestra')}", + ) except Exception as exc: # missing pkg / bad key / init failure — never crash logger.warning("braintrust init failed (%s) — tracing disabled", exc) return _STATE["enabled"] @@ -87,10 +92,12 @@ def run_context() -> tuple[dict, list]: if stage: meta["stage"] = stage cfg = {} - for env_key, key in (("OJ_CFG_TEMPERATURE", "temperature"), - ("OJ_CFG_MAX_TURNS", "max_turns"), - ("OJ_CFG_ANONYMIZE", "anonymize"), - ("OJ_CFG_REJECTION_ONLY", "rejection_only")): + for env_key, key in ( + ("OJ_CFG_TEMPERATURE", "temperature"), + ("OJ_CFG_MAX_TURNS", "max_turns"), + ("OJ_CFG_ANONYMIZE", "anonymize"), + ("OJ_CFG_REJECTION_ONLY", "rejection_only"), + ): v = os.getenv(env_key) if v not in (None, ""): cfg[key] = v @@ -154,7 +161,9 @@ def span(name: str, *, span_type: str = "task", **fields: Any): try: cm = _STATE["bt"].start_span(name=name, type=span_type, **fields) except Exception as exc: # span creation failed — run trace-less, never break - logger.warning("braintrust start_span(%s) failed (%s) — continuing untraced", name, exc) + logger.warning( + "braintrust start_span(%s) failed (%s) — continuing untraced", name, exc + ) yield _NullSpan() return with cm as s: diff --git a/src/openjarvis/agents/hybrid/toolorchestra/unified.py b/src/openjarvis/agents/hybrid/toolorchestra/unified.py index e052a4ff..8b2982ad 100644 --- a/src/openjarvis/agents/hybrid/toolorchestra/unified.py +++ b/src/openjarvis/agents/hybrid/toolorchestra/unified.py @@ -14,12 +14,6 @@ import json import threading from typing import Any, Callable, Dict, List, Optional, Tuple -# Guards the lazy, one-time build of the shared ToolExecutor in -# ``_dispatch_openjarvis_tool``: with concurrent rollouts (parallel rejection -# sampling) several threads can reach the build at once and would each instantiate -# the full tool registry. The lock makes it build-once. -_EXECUTOR_BUILD_LOCK = threading.Lock() - from openjarvis.agents.hybrid._prices import cost as _model_cost from openjarvis.agents.hybrid.expert_registry import ExpertTool, to_worker_dict from openjarvis.agents.hybrid.toolorchestra.parsing import _parse_rl_tool_call @@ -28,8 +22,14 @@ from openjarvis.agents.hybrid.toolorchestra.rollout import ( build_system_prompt, run_unified_rollout, ) -from openjarvis.agents.hybrid.toolorchestra.workers import _call_worker from openjarvis.agents.hybrid.toolorchestra.tracing import span +from openjarvis.agents.hybrid.toolorchestra.workers import _call_worker + +# Guards the lazy, one-time build of the shared ToolExecutor in +# ``_dispatch_openjarvis_tool``: with concurrent rollouts (parallel rejection +# sampling) several threads can reach the build at once and would each instantiate +# the full tool registry. The lock makes it build-once. +_EXECUTOR_BUILD_LOCK = threading.Lock() def make_call_orchestrator( @@ -73,7 +73,9 @@ def make_call_orchestrator( from openjarvis.agents.hybrid.toolorchestra.tracing import wrap_client - client = wrap_client(OpenAI(base_url=base_url, api_key=api_key or "EMPTY", timeout=timeout)) + client = wrap_client( + OpenAI(base_url=base_url, api_key=api_key or "EMPTY", timeout=timeout) + ) def call_orchestrator(messages: List[Dict[str, Any]], specs: List[Dict[str, Any]]): # ``messages`` is the full running conversation (system/user/assistant/tool) @@ -93,7 +95,9 @@ def make_call_orchestrator( # repetition_penalty is a vLLM extra (not OpenAI-native); only send it to a # local vLLM endpoint (base_url set), never to the cloud frontier APIs. if base_url and repetition_penalty and repetition_penalty != 1.0: - kwargs.setdefault("extra_body", {})["repetition_penalty"] = repetition_penalty + kwargs.setdefault("extra_body", {})["repetition_penalty"] = ( + repetition_penalty + ) resp = client.chat.completions.create( model=model, messages=send, @@ -204,10 +208,17 @@ def make_dispatch( # The orchestrator emitted a tool call with no/empty input. Don't hit # the API (it 400s on empty content) — return a usable error so the # rollout keeps going instead of dropping. - return (f"[{tool.name}: no input provided — supply a non-empty " - "'input' to delegate]", 0.0, 0, True) + return ( + f"[{tool.name}: no input provided — supply a non-empty " + "'input' to delegate]", + 0.0, + 0, + True, + ) text, p, c, is_local, extra_cost, _n = _call_worker(worker, prompt, cfg) - usd = (0.0 if is_local else _model_cost(str(tool.model), p, c)) + float(extra_cost) + usd = (0.0 if is_local else _model_cost(str(tool.model), p, c)) + float( + extra_cost + ) return text, usd, int(p) + int(c), bool(is_local) def dispatch(tool: ExpertTool, arguments: Dict[str, Any]): @@ -218,12 +229,22 @@ def make_dispatch( # the trace reads clearly even under anonymization; ``anon_label`` in the # metadata records the opaque label the orchestrator actually saw/chose. real_model = str(tool.model) - with span(f"route:{real_model}", span_type="tool", input=arguments, - metadata={"real_model": real_model, "anon_label": tool.name, - "backend": tool.backend_type}) as s: + with span( + f"route:{real_model}", + span_type="tool", + input=arguments, + metadata={ + "real_model": real_model, + "anon_label": tool.name, + "backend": tool.backend_type, + }, + ) as s: obs, usd, toks, is_local = _dispatch_inner(tool, arguments) - s.log(output=obs, metrics={"cost_usd": usd, "tokens": toks}, - metadata={"is_local": is_local}) + s.log( + output=obs, + metrics={"cost_usd": usd, "tokens": toks}, + metadata={"is_local": is_local}, + ) return obs, usd, toks, is_local return dispatch @@ -245,7 +266,10 @@ def teacher_rollout( question, tools, call_orchestrator=make_call_orchestrator( - teacher_model, base_url=base_url, api_key=api_key, temperature=temperature, + teacher_model, + base_url=base_url, + api_key=api_key, + temperature=temperature, ), dispatch=make_dispatch(cfg), max_turns=max_turns, diff --git a/src/openjarvis/agents/hybrid/toolorchestra/workers.py b/src/openjarvis/agents/hybrid/toolorchestra/workers.py index ee524bdd..4375fb9e 100644 --- a/src/openjarvis/agents/hybrid/toolorchestra/workers.py +++ b/src/openjarvis/agents/hybrid/toolorchestra/workers.py @@ -8,6 +8,8 @@ from typing import Any, Dict, List, Optional, Tuple from openjarvis.agents.hybrid._base import ( ANTHROPIC_WEB_SEARCH_TOOL, + GEMINI_SEARCH_COST_PER_CALL, + OPENAI_WEB_SEARCH_COST_PER_CALL, WEB_SEARCH_COST_PER_CALL, LocalCloudAgent, ) @@ -26,6 +28,7 @@ from openjarvis.agents.hybrid.toolorchestra.sandbox import ( _call_tavily_search, ) + def _paper_pool( local_model: Optional[str], local_endpoint: Optional[str], @@ -39,46 +42,73 @@ def _paper_pool( """ pool: List[Dict[str, Any]] = [] if local_model and local_endpoint: - pool.append({ + pool.append( + { + "id": len(pool), + "name": "local-qwen", + "type": "vllm", + "model": local_model, + "base_url": local_endpoint, + "description": "Local Qwen vLLM (paper uses Qwen3-32B).", + } + ) + pool.append( + { "id": len(pool), - "name": "local-qwen", - "type": "vllm", - "model": local_model, - "base_url": local_endpoint, - "description": "Local Qwen vLLM (paper uses Qwen3-32B).", - }) - pool.append({ - "id": len(pool), "name": "tavily-search", - "type": "tavily-search", "model": "tavily", - "description": "Tavily web search.", - }) - pool.append({ - "id": len(pool), "name": "modal-python", - "type": "modal-python", "model": "modal-python", - "description": "Modal Sandbox for one-shot Python exec.", - }) - pool.append({ - "id": len(pool), "name": "code-specialist", - "type": "openrouter", "model": _PAPER_CODER_OPENROUTER, - "description": "Qwen-2.5-Coder-32B via OpenRouter (paper).", - }) - pool.append({ - "id": len(pool), "name": "generalist-llama", - "type": "openrouter", "model": _PAPER_GENERALIST_TIER3_OPENROUTER, - "description": "Llama-3.3-70B-Instruct via OpenRouter (paper tier-3).", - }) - pool.append({ - "id": len(pool), "name": "generalist-gpt5", - "type": "openai", "model": "gpt-5", - "description": "GPT-5 frontier generalist.", - }) - pool.append({ - "id": len(pool), "name": "generalist-gpt5-mini", - "type": "openai", "model": "gpt-5-mini", - "description": "GPT-5-mini mid generalist.", - }) + "name": "tavily-search", + "type": "tavily-search", + "model": "tavily", + "description": "Tavily web search.", + } + ) + pool.append( + { + "id": len(pool), + "name": "modal-python", + "type": "modal-python", + "model": "modal-python", + "description": "Modal Sandbox for one-shot Python exec.", + } + ) + pool.append( + { + "id": len(pool), + "name": "code-specialist", + "type": "openrouter", + "model": _PAPER_CODER_OPENROUTER, + "description": "Qwen-2.5-Coder-32B via OpenRouter (paper).", + } + ) + pool.append( + { + "id": len(pool), + "name": "generalist-llama", + "type": "openrouter", + "model": _PAPER_GENERALIST_TIER3_OPENROUTER, + "description": "Llama-3.3-70B-Instruct via OpenRouter (paper tier-3).", + } + ) + pool.append( + { + "id": len(pool), + "name": "generalist-gpt5", + "type": "openai", + "model": "gpt-5", + "description": "GPT-5 frontier generalist.", + } + ) + pool.append( + { + "id": len(pool), + "name": "generalist-gpt5-mini", + "type": "openai", + "model": "gpt-5-mini", + "description": "GPT-5-mini mid generalist.", + } + ) return pool + def _default_pool( local_model: Optional[str], local_endpoint: Optional[str], @@ -96,47 +126,67 @@ def _default_pool( ep = "anthropic" pool: List[Dict[str, Any]] = [] if local_model and local_endpoint: - pool.append({ + pool.append( + { + "id": len(pool), + "name": "local-qwen", + "type": "vllm", + "model": local_model, + "base_url": local_endpoint, + "description": ( + "Open-weights Qwen3.5 served locally. Cheap and fast. Good at " + "concise extraction, formatting, arithmetic on given data." + ), + } + ) + if ep == "openai": + search_type = "openai-web-search" + search_model = cloud_model + search_desc = "OpenAI hosted web search on the configured frontier model." + elif ep == "gemini": + search_type = "gemini-web-search" + search_model = cloud_model + search_desc = "Gemini Google Search grounding on the configured frontier model." + else: + search_type = "anthropic-web-search" + search_model = _DEFAULT_WEB_SEARCH_MODEL + search_desc = "Anthropic server-side web_search." + pool.append( + { "id": len(pool), - "name": "local-qwen", - "type": "vllm", - "model": local_model, - "base_url": local_endpoint, + "name": "web-search", + "type": search_type, + "model": search_model, "description": ( - "Open-weights Qwen3.5 served locally. Cheap and fast. Good at " - "concise extraction, formatting, arithmetic on given data." + f"{search_desc} Use for facts that need a lookup " + "(recent events, rare names/dates, niche sources). Returns a digest." ), - }) - pool.append({ - "id": len(pool), - "name": "web-search", - "type": "anthropic-web-search", - "model": "claude-haiku-4-5", - "description": ( - "Anthropic server-side web_search. Use for facts that need a lookup " - "(recent events, rare names/dates, niche sources). Returns a digest." - ), - }) - pool.append({ - "id": len(pool), - "name": f"frontier-{ep}", - "type": ep, - "model": cloud_model, - "description": ( - "Frontier reasoning model. Use for hard multi-step reasoning, " - "code review, or a final synthesis pass. Expensive — use sparingly." - ), - }) - pool.append({ - "id": len(pool), - "name": "frontier-openai-mini", - "type": "openai", - "model": "gpt-5-mini", - "description": ( - "Mid-tier OpenAI model. Solid general knowledge and reasoning at a " - "fraction of frontier cost." - ), - }) + } + ) + pool.append( + { + "id": len(pool), + "name": f"frontier-{ep}", + "type": ep, + "model": cloud_model, + "description": ( + "Frontier reasoning model. Use for hard multi-step reasoning, " + "code review, or a final synthesis pass. Expensive — use sparingly." + ), + } + ) + pool.append( + { + "id": len(pool), + "name": "frontier-openai-mini", + "type": "openai", + "model": "gpt-5-mini", + "description": ( + "Mid-tier OpenAI model. Solid general knowledge and reasoning at a " + "fraction of frontier cost." + ), + } + ) return pool @@ -150,8 +200,22 @@ def _default_pool( # `modal-python` — One-shot Python exec in a fresh Modal Sandbox (the # paper's "Python sandbox" inside `enhance_reasoning`). _TOOLORCH_VALID_TYPES = ( - "vllm", "openai", "anthropic", "anthropic-web-search", "gemini", - "tavily-search", "openrouter", "modal-python", + "vllm", + "openai", + "anthropic", + "anthropic-web-search", + "openai-web-search", + "gemini", + "gemini-web-search", + "tavily-search", + "openrouter", + "modal-python", +) +_TOOLORCH_SEARCH_TYPES = ( + "anthropic-web-search", + "openai-web-search", + "gemini-web-search", + "tavily-search", ) # Default model used when an `anthropic-web-search` entry omits `model`. @@ -172,10 +236,12 @@ def _resolve_worker_pool( the override is absent. Each user-supplied entry must be a dict with keys ``id``, ``name``, - ``type``, and (for non-search types) ``model``. ``type`` must be one - of ``vllm`` / ``openai`` / ``anthropic`` / ``anthropic-web-search``. - ``anthropic-web-search`` entries may omit ``model`` — it defaults to - ``claude-haiku-4-5``. + ``type``, and (for non-search types) ``model``. Search worker types are + ``anthropic-web-search``, ``openai-web-search``, ``gemini-web-search``, + and ``tavily-search``. ``anthropic-web-search`` entries may omit + ``model`` — it defaults to ``claude-haiku-4-5``. OpenAI and Gemini + search workers default to the configured cloud model. Tavily does not + require a model. Substitution: ``model = "$local"`` (or ``""``) resolves to ``local_model``; ``model = "$cloud"`` / ``""`` to ``cloud_model``. @@ -208,9 +274,7 @@ def _resolve_worker_pool( f"Invalid worker_pool entry [{wid_repr}]: 'id' must be an int" ) if wid in seen_ids: - raise ValueError( - f"Invalid worker_pool entry [{wid}]: duplicate id" - ) + raise ValueError(f"Invalid worker_pool entry [{wid}]: duplicate id") seen_ids.add(wid) if not entry.get("name") or not isinstance(entry["name"], str): raise ValueError( @@ -237,14 +301,27 @@ def _resolve_worker_pool( elif isinstance(model, str) and model in ("$cloud", ""): model = cloud_model entry["model"] = model - if wtype == "anthropic-web-search": + if wtype in _TOOLORCH_SEARCH_TYPES: if model in (None, ""): - model = _DEFAULT_WEB_SEARCH_MODEL + if wtype == "anthropic-web-search": + model = _DEFAULT_WEB_SEARCH_MODEL + elif wtype in ("openai-web-search", "gemini-web-search"): + model = cloud_model + else: + model = wtype entry["model"] = model elif not isinstance(model, str): raise ValueError( f"Invalid worker_pool entry [{wid}]: 'model' must be a string when set" ) + if ( + wtype in ("openai-web-search", "gemini-web-search") + and model not in PRICES + ): + raise ValueError( + f"Invalid worker_pool entry [{wid}]: model {model!r} " + f"is not in PRICES (known: {sorted(PRICES)})" + ) # Search workers don't satisfy the "needs a solver" requirement. else: if not isinstance(model, str) or not model: @@ -276,7 +353,7 @@ def _resolve_worker_pool( if not has_non_search: raise ValueError( "Invalid worker_pool entry [-]: worker_pool must contain at least " - "one non-search worker (vllm / openai / anthropic)" + "one non-search worker (vllm / openai / anthropic / gemini)" ) return resolved @@ -286,7 +363,9 @@ def _call_worker( ) -> Tuple[str, int, int, bool, float, int]: """Returns (text, p_tok, c_tok, is_local, extra_cost, n_web_searches).""" wtype = worker.get("type", "openai") - max_tok = int(cfg.get("worker_max_tokens") or os.environ.get("OJ_WORKER_MAX_TOKENS", "4096")) + max_tok = int( + cfg.get("worker_max_tokens") or os.environ.get("OJ_WORKER_MAX_TOKENS", "4096") + ) temp = float(cfg.get("worker_temperature", 0.2)) if wtype == "vllm": @@ -343,6 +422,27 @@ def _call_worker( ) extra = n_searches * WEB_SEARCH_COST_PER_CALL return text, p, c, False, extra, n_searches + if wtype == "openai-web-search": + eff_temp = 1.0 if is_gpt5_family(worker["model"]) else temp + text, p, c, n_searches, _ = LocalCloudAgent._call_openai_agent( + worker["model"], + user=prompt, + max_tokens=max(max_tok, 16384) + if is_gpt5_family(worker["model"]) + else max_tok, + temperature=eff_temp, + ) + extra = n_searches * OPENAI_WEB_SEARCH_COST_PER_CALL + return text, p, c, False, extra, n_searches + if wtype == "gemini-web-search": + text, p, c, n_searches, _ = LocalCloudAgent._call_gemini_agent( + worker["model"], + user=prompt, + max_tokens=max_tok, + temperature=temp, + ) + extra = n_searches * GEMINI_SEARCH_COST_PER_CALL + return text, p, c, False, extra, n_searches if wtype == "tavily-search": # Tavily costs are flat per call; charge `WEB_SEARCH_COST_PER_CALL` # for parity with the Anthropic web-search worker. One call = one @@ -384,7 +484,7 @@ def _swe_call_worker( caller can surface ``tool_calls`` per row. Fallbacks to one-shot workers return 0 bash turns (no agent loop ran).""" wtype = worker.get("type", "openai") - if wtype == "anthropic-web-search": + if wtype in _TOOLORCH_SEARCH_TYPES: # Search workers stay one-shot. text, p, c, is_local, extra, n_searches = _call_worker(worker, prompt, cfg) return text, p, c, is_local, extra, n_searches, 0 @@ -417,6 +517,10 @@ def _swe_call_worker( is_local = backbone == "local" return ( out["final_summary"] or out["answer"], - out["tokens_in"], out["tokens_out"], - is_local, 0.0, 0, int(out["turns"]), + out["tokens_in"], + out["tokens_out"], + is_local, + 0.0, + 0, + int(out["turns"]), ) diff --git a/src/openjarvis/evals/core/scorer.py b/src/openjarvis/evals/core/scorer.py index 59700cf3..bc02bd96 100644 --- a/src/openjarvis/evals/core/scorer.py +++ b/src/openjarvis/evals/core/scorer.py @@ -93,13 +93,18 @@ class LLMJudgeScorer(Scorer): ) except Exception as exc: # noqa: BLE001 - re-raised below last_exc = exc - if attempt == _JUDGE_MAX_RETRIES - 1 or not _is_retryable_judge_error(exc): + if attempt == _JUDGE_MAX_RETRIES - 1 or not _is_retryable_judge_error( + exc + ): raise - delay = min(_JUDGE_BASE_DELAY_S * (2 ** attempt), _JUDGE_MAX_DELAY_S) + delay = min(_JUDGE_BASE_DELAY_S * (2**attempt), _JUDGE_MAX_DELAY_S) delay += random.uniform(0.0, delay * 0.25) # jitter LOGGER.warning( "judge call failed (attempt %d/%d): %s — retrying in %.1fs", - attempt + 1, _JUDGE_MAX_RETRIES, exc, delay, + attempt + 1, + _JUDGE_MAX_RETRIES, + exc, + delay, ) time.sleep(delay) # Unreachable (loop either returns or raises), but keeps type-checkers happy. diff --git a/src/openjarvis/learning/intelligence/orchestrator/eval_backend.py b/src/openjarvis/learning/intelligence/orchestrator/eval_backend.py index 6aa34c19..3fe76e28 100644 --- a/src/openjarvis/learning/intelligence/orchestrator/eval_backend.py +++ b/src/openjarvis/learning/intelligence/orchestrator/eval_backend.py @@ -22,6 +22,14 @@ import re import time from typing import Any, Dict, Optional +from openjarvis.agents.hybrid.expert_registry import orchestrator_catalog +from openjarvis.agents.hybrid.toolorchestra import rollout as _rollout_mod +from openjarvis.agents.hybrid.toolorchestra.unified import ( + make_call_orchestrator, + make_dispatch, +) +from openjarvis.evals.core.backend import InferenceBackend + # The served fine-tuned model sometimes over-emits the answer marker, e.g. # "FINAL_ANSWER: FINAL_ANSWER: 42" — a doubled prefix that breaks answer # extraction and auto-scores the sample 0. Collapse any run of FINAL_ANSWER @@ -34,20 +42,12 @@ def _clean_final_answer(text: str) -> str: marks = list(_FA_MARK.finditer(t)) if not marks: return t - answer = t[marks[-1].end():].strip() + answer = t[marks[-1].end() :].strip() return f"FINAL_ANSWER: {answer}" if answer else t _OBS_CAP = 4000 # cap persisted observations so the eval JSONL stays readable -from openjarvis.agents.hybrid.expert_registry import orchestrator_catalog -from openjarvis.agents.hybrid.toolorchestra import rollout as _rollout_mod -from openjarvis.agents.hybrid.toolorchestra.unified import ( - make_call_orchestrator, - make_dispatch, -) -from openjarvis.evals.core.backend import InferenceBackend - DEFAULT_ENDPOINT = "http://localhost:8001/v1" DEFAULT_MODEL = "qwen3-8b" @@ -215,11 +215,17 @@ class OrchestratorBackend(InferenceBackend): { "reasoning": t.reasoning or "", "tool_name": t.tool_name, - "real_model": anon_map.get(t.tool_name) if t.tool_name else None, + "real_model": anon_map.get(t.tool_name) + if t.tool_name + else None, "arguments": t.arguments, "observation": ( (t.observation or "")[:_OBS_CAP] - + ("…[truncated]" if t.observation and len(t.observation) > _OBS_CAP else "") + + ( + "…[truncated]" + if t.observation and len(t.observation) > _OBS_CAP + else "" + ) ), } for t in (getattr(rollout, "turns", []) or []) diff --git a/src/openjarvis/learning/intelligence/orchestrator/sft_data/datasets.py b/src/openjarvis/learning/intelligence/orchestrator/sft_data/datasets.py index 6ed6a6b8..1c79f53d 100644 --- a/src/openjarvis/learning/intelligence/orchestrator/sft_data/datasets.py +++ b/src/openjarvis/learning/intelligence/orchestrator/sft_data/datasets.py @@ -38,10 +38,10 @@ class Task: task_id: str question: str answer: str - domain: str # coarse area: math / code / science / medical / chat / misc - difficulty: str = "" # OpenThoughts difficulty tier (GeneralThought has none) - dataset: str = "" # source dataset: GeneralThought / OpenThoughts3 - subsector: str = "" # fine source: NuminaMath / NHSQA / TACO / glaive / ... + domain: str # coarse area: math / code / science / medical / chat / misc + difficulty: str = "" # OpenThoughts difficulty tier (GeneralThought has none) + dataset: str = "" # source dataset: GeneralThought / OpenThoughts3 + subsector: str = "" # fine source: NuminaMath / NHSQA / TACO / glaive / ... @property def instruction(self) -> str: @@ -194,7 +194,9 @@ def _normalize_openthoughts(row: Dict[str, Any], *, index: int = 0) -> Optional[ if not question or not answer: return None domain = str(row.get("domain") or "unknown").strip().lower() or "unknown" - task_id = str(row.get("id") or row.get("source") or f"openthoughts-{index}") + f"-{index}" + task_id = ( + str(row.get("id") or row.get("source") or f"openthoughts-{index}") + f"-{index}" + ) return Task( task_id=task_id, question=question, @@ -324,9 +326,7 @@ def load_grpo_prompts(*, n: int = 30000, seed: int = 42) -> List[Task]: per = max(n // 4 + 1, 1) pool: List[Task] = [] pool.extend(load_generalthought(n=per, seed=seed)) - pool.extend( - load_openthoughts(n_code=per, n_math=per, n_science=per, seed=seed) - ) + pool.extend(load_openthoughts(n_code=per, n_math=per, n_science=per, seed=seed)) rng = random.Random(seed) rng.shuffle(pool) diff --git a/src/openjarvis/learning/intelligence/orchestrator/sft_data/hotpotqa.py b/src/openjarvis/learning/intelligence/orchestrator/sft_data/hotpotqa.py index 36b2f92b..b3a463f7 100644 --- a/src/openjarvis/learning/intelligence/orchestrator/sft_data/hotpotqa.py +++ b/src/openjarvis/learning/intelligence/orchestrator/sft_data/hotpotqa.py @@ -20,7 +20,7 @@ from __future__ import annotations import re import string from dataclasses import dataclass -from typing import Any, Dict, Iterable, Iterator, List, Optional +from typing import Any, Dict, Iterable, Iterator, Optional DATASET_ID = "hotpotqa/hotpot_qa" CONFIG = "fullwiki" @@ -32,8 +32,8 @@ class HotpotTask: task_id: str question: str answer: str - level: str = "" # easy | medium | hard - qtype: str = "" # comparison | bridge + level: str = "" # easy | medium | hard + qtype: str = "" # comparison | bridge # Parity with ToolScaleTask so the rejection-sampling loop is dataset-agnostic. @property diff --git a/src/openjarvis/learning/intelligence/orchestrator/sft_data/reject_sample.py b/src/openjarvis/learning/intelligence/orchestrator/sft_data/reject_sample.py index 1563a71d..8ac96e9b 100644 --- a/src/openjarvis/learning/intelligence/orchestrator/sft_data/reject_sample.py +++ b/src/openjarvis/learning/intelligence/orchestrator/sft_data/reject_sample.py @@ -22,9 +22,7 @@ import logging import re from collections import Counter from pathlib import Path -from typing import Callable, Iterable, List, Optional - -from typing import Any +from typing import Any, Callable, Iterable, List, Optional from openjarvis.agents.hybrid.expert_registry import ExpertTool from openjarvis.agents.hybrid.toolorchestra.rollout import UnifiedRollout @@ -107,7 +105,9 @@ def _target_is_clean(roll: UnifiedRollout) -> bool: return False if any(m in fa for m in _ERR_MARKERS): return False - if fa.count("") != fa.count(""): # unbalanced think tags (truncated OR stray ) + if fa.count("") != fa.count( + "" + ): # unbalanced think tags (truncated OR stray ) return False if _TRUNCATED_TAIL_RE.search(fa): # truncated mid-expression (e.g. "a=1, b=") return False @@ -150,28 +150,36 @@ def _target_is_clean(roll: UnifiedRollout) -> bool: # the correct hard-task answers. Keep the STRUCTURAL essay signals (many # numbered sections, markdown headers, tables) which still catch real essays. _marks = list(re.finditer(r"(?im)FINAL[_\s]?ANSWER\s*:?", fa)) - _ans = fa[_marks[-1].end():].strip() if _marks else fa - if len(_ans) > 2000: # was 700 — too tight for code/math + _ans = fa[_marks[-1].end() :].strip() if _marks else fa + if len(_ans) > 2000: # was 700 — too tight for code/math return False # Tool-status echo as the final answer: on code tasks the orchestrator # sometimes ends with a file_write/shell status line ("Successfully wrote to # /tmp/x.py") instead of the actual answer — a non-answer that slips the # length/format checks. Reject. (Audit: ~code/math trajectories where the # trace collapsed but the row was still scored correct.) - if re.search(r"(?i)\b(successfully (wrote|created|saved|executed|ran)|written to /|" - r"file (written|saved|created)|no further actions?)\b", _ans): + if re.search( + r"(?i)\b(successfully (wrote|created|saved|executed|ran)|written to /|" + r"file (written|saved|created)|no further actions?)\b", + _ans, + ): return False - if len(re.findall(r"(?m)^\s*\d+\.\s", _ans)) >= 6: # many numbered sections = essay (was 4) + if ( + len(re.findall(r"(?m)^\s*\d+\.\s", _ans)) >= 6 + ): # many numbered sections = essay (was 4) return False - if re.search(r"(?m)^\s*#{1,6}\s", _ans): # markdown headers = essay + if re.search(r"(?m)^\s*#{1,6}\s", _ans): # markdown headers = essay return False # Garbled / shouty final: a multi-word answer that's mostly UPPERCASE, or has # an absurdly long merged all-letter token, is decode garble — the audit found # e.g. "RABES PEPTETANUS BOOSTERS" (for "rabies PEP, tetanus boosters") passing. if len(_ans) > 20: _alpha = [c for c in _ans if c.isalpha()] - if (_alpha and sum(c.isupper() for c in _alpha) / len(_alpha) > 0.7 - and len(_ans.split()) >= 3): + if ( + _alpha + and sum(c.isupper() for c in _alpha) / len(_alpha) > 0.7 + and len(_ans.split()) >= 3 + ): return False if any(len(w) > 25 and w.isalpha() for w in _ans.split()): return False @@ -184,7 +192,9 @@ def _target_is_clean(roll: UnifiedRollout) -> bool: for t in roll.turns: if t.observation: obs_head = re.sub(r"\s+", " ", t.observation[:200]).strip() - if fa_head and fa_head == obs_head and len(fa) > 200: # long final = copied tool dump (short relays are legit) + if ( + fa_head and fa_head == obs_head and len(fa) > 200 + ): # long final = copied tool dump (short relays are legit) return False # "routed" = delegated to a model EXPERT (not just a utility like web_search). # When anonymized, expert calls are the anon labels in anon_map; otherwise @@ -310,8 +320,8 @@ def generate_sft_dataset( written = 0 correct_written = 0 incorrect_written = 0 - dropped = 0 # tasks that produced no kept record - tasks_solved = 0 # tasks with >=1 correct sample + dropped = 0 # tasks that produced no kept record + tasks_solved = 0 # tasks with >=1 correct sample domain_counts: Counter[str] = Counter() def _work(task: TaskLike) -> tuple[TaskLike, List[tuple[UnifiedRollout, bool]]]: @@ -329,8 +339,12 @@ def generate_sft_dataset( nonlocal written, correct_written, incorrect_written reward = reward_fn(roll) if reward_fn else 0.0 record = trajectory_to_record( - task.task_id, task.instruction, tools, roll, - reward=reward, domain=task.domain, + task.task_id, + task.instruction, + tools, + roll, + reward=reward, + domain=task.domain, ) record["correct"] = ok record["kept"] = kept @@ -380,8 +394,9 @@ def generate_sft_dataset( if not correct: dropped += 1 return - for roll in sorted((r for r, _ in correct), - key=lambda r: r.cost_usd)[:max_keep_per_task]: + for roll in sorted((r for r, _ in correct), key=lambda r: r.cost_usd)[ + :max_keep_per_task + ]: _emit(fh, task, roll, True, kept=(roll is cheapest)) with out.open("w") as fh: @@ -400,15 +415,22 @@ def generate_sft_dataset( try: _, results = fut.result() except Exception as exc: # a task's worker died; count + skip - logger.warning("task %s failed: %s", - getattr(task, "task_id", "?"), exc) + logger.warning( + "task %s failed: %s", getattr(task, "task_id", "?"), exc + ) results = [] _write(fh, task, results) done += 1 if done % 50 == 0 or done == seen: - logger.info("rejection-sampling: %d/%d tasks done " - "(%d written, %d solved, %d dropped)", - done, seen, written, tasks_solved, dropped) + logger.info( + "rejection-sampling: %d/%d tasks done " + "(%d written, %d solved, %d dropped)", + done, + seen, + written, + tasks_solved, + dropped, + ) stats = { "out_path": str(out), @@ -425,8 +447,13 @@ def generate_sft_dataset( "domain_distribution": dict(domain_counts), } out.with_suffix(out.suffix + ".stats.json").write_text(json.dumps(stats, indent=2)) - logger.info("Wrote %d SFT records to %s (%d correct, %d incorrect)", - written, out, correct_written, incorrect_written) + logger.info( + "Wrote %d SFT records to %s (%d correct, %d incorrect)", + written, + out, + correct_written, + incorrect_written, + ) return stats diff --git a/src/openjarvis/learning/intelligence/orchestrator/sft_data/toolscale.py b/src/openjarvis/learning/intelligence/orchestrator/sft_data/toolscale.py index 6e1aa824..4e946394 100644 --- a/src/openjarvis/learning/intelligence/orchestrator/sft_data/toolscale.py +++ b/src/openjarvis/learning/intelligence/orchestrator/sft_data/toolscale.py @@ -76,19 +76,25 @@ def normalize_row(row: Dict[str, Any], *, index: int = 0) -> ToolScaleTask: gold: List[GoldAction] = [] for a in _as_list(crit.get("actions")): if isinstance(a, dict) and a.get("name"): - gold.append(GoldAction( - name=str(a["name"]), - arguments=a.get("arguments") or a.get("args") or {}, - action_id=a.get("action_id"), - )) + gold.append( + GoldAction( + name=str(a["name"]), + arguments=a.get("arguments") or a.get("args") or {}, + action_id=a.get("action_id"), + ) + ) required = _str_list(crit.get("communicate_info")) nl = _str_list(crit.get("nl_assertions")) task_id = str(row.get("id") or row.get("task_id") or f"toolscale-{index}") return ToolScaleTask( - task_id=task_id, domain=str(domain), instruction=str(instruction), - gold_actions=gold, required_info=required, nl_assertions=nl, + task_id=task_id, + domain=str(domain), + instruction=str(instruction), + gold_actions=gold, + required_info=required, + nl_assertions=nl, ) diff --git a/src/openjarvis/learning/intelligence/orchestrator/sft_data/unified_serialize.py b/src/openjarvis/learning/intelligence/orchestrator/sft_data/unified_serialize.py index c6ee5f40..2bc24566 100644 --- a/src/openjarvis/learning/intelligence/orchestrator/sft_data/unified_serialize.py +++ b/src/openjarvis/learning/intelligence/orchestrator/sft_data/unified_serialize.py @@ -12,7 +12,6 @@ prompt), ``assistant`` (reasoning + a ```` tag, or the final answer), from __future__ import annotations -import json import re from typing import Any, Dict, List @@ -79,23 +78,26 @@ def _debox(text: str) -> str: # reasoning about its own harness. Never a good target: drop the whole line. _LEAK_RE = re.compile( r"(?im)^.*\b(?:" - r"the user (?:is|was|wants|has|had|needs|'s|would|will|asked|'ve asked|has asked)\b|" - r"asked me to\b|" # echoes of an injected nudge turn + r"the user (?:is|was|wants|has|had|needs|'s|would|will|asked" + r"|'ve asked|has asked)\b|" + r"asked me to\b|" # echoes of an injected nudge turn r"make progress by delegat|" r"delegat\w+[^\n]*sub-?question|" r"(?:to |ask |send .{0,20}to )?one of the models\b|" - r"instructions?|" # any self-reference to the rules + r"instructions?|" # any self-reference to the rules r"requirement|" - r"at least one (?:model|tool|expert)|" # "(call|invoke) at least one model" + r"at least one (?:model|tool|expert)|" # "(call|invoke) at least one model" r"(?:call|invoke|route to|delegate to|use)\s+(?:a|one|at least one|the)\s+model\b|" r"route\s+(?:this|it|that)\b[^\n]{0,40}\bthrough\s+(?:a|one|the)\s+model\b|" r"without mentioning (?:any )?(?:tools?|reasoning|meta|steps)|" r"in the required format\b|" - r"based on the (?:model|response|analysis|expert)|" # response-acknowledgment - r"the model (?:provided|confirmed|gave|returned|correctly|analy\w+|respon\w+|said|indicated|identified)|" - r"i (?:got|received|have) (?:a |the )?(?:comprehensive |clear |good |detailed )?answer\b|" + r"based on the (?:model|response|analysis|expert)|" # response-acknowledgment + r"the model (?:provided|confirmed|gave|returned|correctly|analy\w+" + r"|respon\w+|said|indicated|identified)|" + r"i (?:got|received|have) (?:a |the )?" + r"(?:comprehensive |clear |good |detailed )?answer\b|" r"now i can (?:confidently )?(?:give|provide|state|answer)|" - r"perfect[!,]|great[!,]|excellent[!,]|" # narration exclamations + r"perfect[!,]|great[!,]|excellent[!,]|" # narration exclamations r"as an orchestrator|" r"whose job is to (?:route|delegate|orchestrate)|" r"testing whether i (?:need|have|should|must)\b|" @@ -121,7 +123,7 @@ def _scrub_meta(text: str) -> str: m = re.match(r"(?is)^\s*()\s*", body) if m: think = "\n" - body = body[m.end():] + body = body[m.end() :] body = _LEAK_RE.sub("", body) body = _META_OPEN_RE.sub("", body).lstrip() return (think + body).strip() if think else body.strip() @@ -139,9 +141,9 @@ def _scrub_meta(text: str) -> str: # match, so legitimate ``<``/``>`` in math/code (``x < 3 and y > 2``) is left # untouched. _CONTROL_TOKEN_RE = re.compile( - r"<\|[^>]*?\|>" # closed pipe token - r"|<\|[^|>]*>" # unclosed pipe token - r"|" # named angle tokens + r"<\|[^>]*?\|>" # closed pipe token + r"|<\|[^|>]*>" # unclosed pipe token + r"|" # named angle tokens ) @@ -173,7 +175,7 @@ def _final_answer_block(text: str) -> str: marks = list(re.finditer(r"(?im)FINAL[_\s]?ANSWER\s*:?", text)) if marks: last = marks[-1] - answer = _strip_control_tokens(text[last.end():].strip()) + answer = _strip_control_tokens(text[last.end() :].strip()) # Final turn = the bare short answer only. Drop EVERYTHING before # FINAL_ANSWER — both the visible "Perfect! Based on the model's response…" # narration AND the final-turn , which is just confirmatory @@ -186,7 +188,11 @@ def _final_answer_block(text: str) -> str: # Drop the think block (same rationale as above) — keep only the answer. _, _, answer = text.rpartition("") answer = _strip_control_tokens(answer.strip()) - return f"FINAL_ANSWER: {answer}" if answer else f"FINAL_ANSWER: {_strip_control_tokens(_scrub_meta(text))}" + return ( + f"FINAL_ANSWER: {answer}" + if answer + else f"FINAL_ANSWER: {_strip_control_tokens(_scrub_meta(text))}" + ) return f"FINAL_ANSWER: {_strip_control_tokens(text)}" @@ -209,29 +215,39 @@ def trajectory_to_record( if turn.tool_name is None: # Final-answer turn. turn.reasoning is the model's actual output for # this turn (which already contains the answer); render it once. - conversations.append({ - "role": "assistant", - "content": _final_answer_block(turn.reasoning or rollout.final_answer), - }) + conversations.append( + { + "role": "assistant", + "content": _final_answer_block( + turn.reasoning or rollout.final_answer + ), + } + ) continue tag = tool_call_tag(turn.tool_name, turn.arguments) reasoning = _scrub_meta(_normalize_think((turn.reasoning or "").rstrip())) - conversations.append({ - "role": "assistant", - "content": (reasoning + "\n" + tag).strip(), - }) - conversations.append({ - "role": "tool", - "name": turn.tool_name, - "content": turn.observation or "", - }) + conversations.append( + { + "role": "assistant", + "content": (reasoning + "\n" + tag).strip(), + } + ) + conversations.append( + { + "role": "tool", + "name": turn.tool_name, + "content": turn.observation or "", + } + ) # If the rollout terminated on max_turns (no None turn), append the answer. if not rollout.turns or rollout.turns[-1].tool_name is not None: - conversations.append({ - "role": "assistant", - "content": _final_answer_block(rollout.final_answer), - }) + conversations.append( + { + "role": "assistant", + "content": _final_answer_block(rollout.final_answer), + } + ) return { "conversations": conversations, diff --git a/src/openjarvis/learning/intelligence/orchestrator/sft_data/verify.py b/src/openjarvis/learning/intelligence/orchestrator/sft_data/verify.py index c7e97add..fd8b5044 100644 --- a/src/openjarvis/learning/intelligence/orchestrator/sft_data/verify.py +++ b/src/openjarvis/learning/intelligence/orchestrator/sft_data/verify.py @@ -173,7 +173,9 @@ def _gemini_judge(task: Task, prediction: str) -> Optional[bool]: # caller falls back to string/f1. global _JUDGE_CLIENT if _JUDGE_CLIENT is None: - _JUDGE_CLIENT = anthropic.Anthropic(api_key=api_key, max_retries=0, timeout=15) + _JUDGE_CLIENT = anthropic.Anthropic( + api_key=api_key, max_retries=0, timeout=15 + ) client = _JUDGE_CLIENT prompt = ( "You are grading a candidate answer against a gold reference.\n" @@ -203,7 +205,7 @@ def _gemini_judge(task: Task, prediction: str) -> Optional[bool]: def verify_answer(task: Task, prediction: str) -> bool: - """Domain-dispatched correctness check for ``prediction`` against ``task.answer``.""" + """Domain-dispatched correctness check for ``prediction`` vs ``task.answer``.""" pred = (prediction or "").strip() if not pred or not (task.answer or "").strip(): return False diff --git a/src/openjarvis/learning/intelligence/orchestrator/sft_trainer.py b/src/openjarvis/learning/intelligence/orchestrator/sft_trainer.py index 9e467588..4a41a832 100644 --- a/src/openjarvis/learning/intelligence/orchestrator/sft_trainer.py +++ b/src/openjarvis/learning/intelligence/orchestrator/sft_trainer.py @@ -296,12 +296,14 @@ class OrchestratorSFTTrainer: make_call_orchestrator, make_dispatch, ) + from openjarvis.learning.intelligence.orchestrator.sft_data import ( + reject_sample as _reject_sample, + ) from openjarvis.learning.intelligence.orchestrator.sft_data.datasets import ( load_sft_tasks, ) - from openjarvis.learning.intelligence.orchestrator.sft_data.reject_sample import ( - generate_sft_dataset, - ) + + generate_sft_dataset = _reject_sample.generate_sft_dataset from openjarvis.learning.intelligence.orchestrator.sft_data.verify import ( make_verifier, ) @@ -321,8 +323,10 @@ class OrchestratorSFTTrainer: def rollout_fn(task: Any) -> Any: try: return run_unified_rollout( - task.instruction, tools, - call_orchestrator=call_orch, dispatch=dispatch, + task.instruction, + tools, + call_orchestrator=call_orch, + dispatch=dispatch, max_turns=self.config.max_rollout_turns, ) except Exception as exc: # network/key failures shouldn't kill the run diff --git a/src/openjarvis/tools/calculator.py b/src/openjarvis/tools/calculator.py index b4685b3b..3e391706 100644 --- a/src/openjarvis/tools/calculator.py +++ b/src/openjarvis/tools/calculator.py @@ -52,7 +52,6 @@ _MATH_FUNCS = { "radians": math.radians, "degrees": math.degrees, "exp": math.exp, - "abs": abs, "fabs": math.fabs, "hypot": math.hypot, "factorial": math.factorial, diff --git a/src/openjarvis/tools/web_search.py b/src/openjarvis/tools/web_search.py index 542e5705..950cc132 100644 --- a/src/openjarvis/tools/web_search.py +++ b/src/openjarvis/tools/web_search.py @@ -123,7 +123,7 @@ class WebSearchTool(BaseTool): request (DDGS timeout) AND wrap the call in a hard wall-clock deadline so a flaky search fast-fails instead of stalling data generation. """ - from concurrent.futures import ThreadPoolExecutor, TimeoutError as _FTO + from concurrent.futures import ThreadPoolExecutor from ddgs import DDGS diff --git a/tests/agents/test_expert_registry.py b/tests/agents/test_expert_registry.py index e04c7538..a14389b7 100644 --- a/tests/agents/test_expert_registry.py +++ b/tests/agents/test_expert_registry.py @@ -10,9 +10,9 @@ from openjarvis.agents.hybrid.expert_registry import ( CATEGORY_BASIC, CATEGORY_CLOUD_FRONTIER, CATEGORY_LOCAL_OSS, - ExpertTool, KIND_MODEL, KIND_TOOL, + ExpertTool, build_tool_specs, default_catalog, openjarvis_tool, @@ -28,8 +28,14 @@ def test_each_model_is_its_own_tool(): cat = default_catalog() names = {t.name for t in cat} # Distinct model tools, each with its own name. - for n in ("gpt_5", "gpt_5_mini", "qwen3_32b", "qwen_2_5_coder_32b_instruct", - "llama_3_3_70b_instruct", "claude_opus_4_7"): + for n in ( + "gpt_5", + "gpt_5_mini", + "qwen3_32b", + "qwen_2_5_coder_32b_instruct", + "llama_3_3_70b_instruct", + "claude_opus_4_7", + ): assert n in names, f"missing model tool {n}" # No meta-tool / slot vocabulary leaks in. assert "answer" not in names and "enhance_reasoning" not in names @@ -56,7 +62,9 @@ def test_invalid_tool_rejected(): with pytest.raises(ValueError): ExpertTool(name="x", kind="bogus", backend_type="openai", summary="", model="m") with pytest.raises(ValueError): - ExpertTool(name="x", kind=KIND_MODEL, backend_type="openai", summary="", model=None) + ExpertTool( + name="x", kind=KIND_MODEL, backend_type="openai", summary="", model=None + ) def test_specs_shape_and_pricing_in_description(): @@ -79,20 +87,31 @@ def test_sample_is_deterministic_and_well_formed(): b = sample_tool_config(cat, rng=random.Random(0), min_tools=4) assert [t.name for t in a] == [t.name for t in b] # deterministic assert len(a) >= 4 - assert any(t.kind == KIND_MODEL for t in a) # can reason - assert any(t.kind != KIND_MODEL for t in a) # can act + assert any(t.kind == KIND_MODEL for t in a) # can reason + assert any(t.kind != KIND_MODEL for t in a) # can act assert {t.name for t in a} <= {t.name for t in cat} # subset def test_price_jitter_changes_prices_reproducibly(): cat = default_catalog() - base = {t.name: t for t in sample_tool_config(cat, rng=random.Random(3), min_tools=8)} - jit = {t.name: t for t in sample_tool_config( - cat, rng=random.Random(3), min_tools=8, price_jitter=0.5)} + base = { + t.name: t for t in sample_tool_config(cat, rng=random.Random(3), min_tools=8) + } + jit = { + t.name: t + for t in sample_tool_config( + cat, rng=random.Random(3), min_tools=8, price_jitter=0.5 + ) + } # Same subset (same seed/sequence up to jitter draws), but model prices move. - moved = [n for n in base - if base[n].kind == KIND_MODEL and base[n].price_in - and n in jit and jit[n].price_in != base[n].price_in] + moved = [ + n + for n in base + if base[n].kind == KIND_MODEL + and base[n].price_in + and n in jit + and jit[n].price_in != base[n].price_in + ] assert moved, "expected jitter to change at least one model price" for n in moved: f = jit[n].price_in / base[n].price_in @@ -137,8 +156,12 @@ def test_orchestrator_catalog_two_model_classes_plus_basics(): assert by["gpt_5_5"].category == CATEGORY_CLOUD_FRONTIER assert by["claude_opus_4_8"].category == CATEGORY_CLOUD_FRONTIER # Default routing for every model tool is OpenRouter (no servers required). - for n in ("qwen3_5_9b", "qwen3_6_27b_fp8", "qwen3_5_122b_a10b_fp8", - "qwen3_5_397b_a17b_fp8"): + for n in ( + "qwen3_5_9b", + "qwen3_6_27b_fp8", + "qwen3_5_122b_a10b_fp8", + "qwen3_5_397b_a17b_fp8", + ): assert by[n].category == CATEGORY_LOCAL_OSS assert by[n].backend_type == "openrouter" assert by[n].base_url is None @@ -175,7 +198,8 @@ def test_orchestrator_local_models_get_base_url_when_provided(): local_endpoints={ "Qwen/Qwen3.5-9B": "http://x/v1", "Qwen/Qwen3.6-27B-FP8": "http://y/v1", - }) + } + ) by = tools_by_name(cat) assert by["qwen3_5_9b"].backend_type == "vllm" assert by["qwen3_5_9b"].base_url == "http://x/v1" @@ -213,7 +237,8 @@ def test_orchestrator_model_backends_override(): def test_orchestrator_openrouter_slug_override(): cat = orchestrator_catalog( - openrouter_slugs={"Qwen/Qwen3.5-9B": "qwen/qwen3.5-9b-custom"}) + openrouter_slugs={"Qwen/Qwen3.5-9B": "qwen/qwen3.5-9b-custom"} + ) by = tools_by_name(cat) assert by["qwen3_5_9b"].model == "qwen/qwen3.5-9b-custom" @@ -236,12 +261,19 @@ def test_openjarvis_tool_bridges_real_tool_with_custom_schema(): def test_build_tool_specs_includes_category_for_bridged_tools(): - specs = build_tool_specs([ - openjarvis_tool("calculator", summary="Math.", - params={"type": "object", - "properties": {"expression": {"type": "string"}}, - "required": ["expression"]}), - ]) + specs = build_tool_specs( + [ + openjarvis_tool( + "calculator", + summary="Math.", + params={ + "type": "object", + "properties": {"expression": {"type": "string"}}, + "required": ["expression"], + }, + ), + ] + ) assert specs[0]["function"]["category"] == CATEGORY_BASIC assert "expression" in specs[0]["function"]["parameters"]["properties"] @@ -250,6 +282,9 @@ def test_to_worker_dict_maps_backend(): cat = default_catalog(local_model="qwen3:8b", local_endpoint="http://x/v1") by = tools_by_name(cat) assert to_worker_dict(by["gpt_5"]) == { - "name": "gpt_5", "type": "openai", "model": "gpt-5"} + "name": "gpt_5", + "type": "openai", + "model": "gpt-5", + } local = to_worker_dict(by["qwen3_8b"]) assert local["type"] == "vllm" and local["base_url"] == "http://x/v1" diff --git a/tests/test_orchestrator_learning/sft_data/test_datasets.py b/tests/test_orchestrator_learning/sft_data/test_datasets.py index fe49a782..6d7f9592 100644 --- a/tests/test_orchestrator_learning/sft_data/test_datasets.py +++ b/tests/test_orchestrator_learning/sft_data/test_datasets.py @@ -60,7 +60,10 @@ OPENTHOUGHTS_ROWS = [ "difficulty": 5, "conversations": [ {"from": "human", "value": "What is 6 times 7?"}, - {"from": "gpt", "value": "6*7=42 The answer is \\boxed{42}."}, + { + "from": "gpt", + "value": "6*7=42 The answer is \\boxed{42}.", + }, ], }, { @@ -99,9 +102,7 @@ def test_load_generalthought_respects_n(): def test_load_openthoughts_fields_and_domains(): tasks = list( - load_openthoughts( - n_code=5, n_math=5, n_science=5, source=OPENTHOUGHTS_ROWS - ) + load_openthoughts(n_code=5, n_math=5, n_science=5, source=OPENTHOUGHTS_ROWS) ) assert all(isinstance(t, Task) for t in tasks) assert len(tasks) == 3 diff --git a/tests/test_orchestrator_learning/sft_data/test_rejection_pipeline.py b/tests/test_orchestrator_learning/sft_data/test_rejection_pipeline.py index da9c8de3..8cec1bf0 100644 --- a/tests/test_orchestrator_learning/sft_data/test_rejection_pipeline.py +++ b/tests/test_orchestrator_learning/sft_data/test_rejection_pipeline.py @@ -55,7 +55,7 @@ def test_run_unified_rollout_terminates_on_no_tool_call(): assert name in by scripted = [ - (f"reason 1\n", [(name, {"input": "do step 1"})], 5, 5), + ("reason 1\n", [(name, {"input": "do step 1"})], 5, 5), ("here is the answer", [], 3, 3), # no tool call -> terminate ] calls = iter(scripted) @@ -69,7 +69,11 @@ def test_run_unified_rollout_terminates_on_no_tool_call(): return (f"OBS for {tool.name}", 0.01, 10, False) roll = run_unified_rollout( - "What is X?", tools, call_orchestrator=call_orch, dispatch=dispatch, max_turns=5, + "What is X?", + tools, + call_orchestrator=call_orch, + dispatch=dispatch, + max_turns=5, ) assert roll.final_answer == "here is the answer" assert roll.num_tool_calls == 1 @@ -81,23 +85,34 @@ def test_serialize_record_shape_and_tool_call_tags(): tools = default_catalog() roll = UnifiedRollout( turns=[ - UnifiedTurn(reasoning="think", tool_name="qwen3_32b", - arguments={"input": "q"}, observation="obs"), + UnifiedTurn( + reasoning="think", + tool_name="qwen3_32b", + arguments={"input": "q"}, + observation="obs", + ), # The final turn's reasoning IS the model's real final output (it # already contains the answer). The serializer now renders this turn # via _final_answer_block(turn.reasoning), not rollout.final_answer. - UnifiedTurn(reasoning="The result is 42.\nFINAL_ANSWER: 42", - tool_name=None), + UnifiedTurn( + reasoning="The result is 42.\nFINAL_ANSWER: 42", tool_name=None + ), ], - final_answer="42", cost_usd=0.02, tokens=30, num_tool_calls=1, + final_answer="42", + cost_usd=0.02, + tokens=30, + num_tool_calls=1, ) rec = trajectory_to_record("t1", "Q?", tools, roll, reward=0.5, domain="math") roles = [m["role"] for m in rec["conversations"]] assert roles[0] == "system" and roles[1] == "user" assert "tool" in roles and roles[-1] == "assistant" # Tool call is emitted as a tag (what the parser reads back). - assert any("" in m["content"] and "qwen3_32b" in m["content"] - for m in rec["conversations"] if m["role"] == "assistant") + assert any( + "" in m["content"] and "qwen3_32b" in m["content"] + for m in rec["conversations"] + if m["role"] == "assistant" + ) assert "FINAL_ANSWER: 42" in rec["conversations"][-1]["content"] assert rec["reward"] == 0.5 and rec["domain"] == "math" @@ -112,10 +127,15 @@ def test_gold_coverage_verify(): final_answer="done", ) missing = UnifiedRollout( - turns=[UnifiedTurn("", "cancel", {}, "ok")], final_answer="done") + turns=[UnifiedTurn("", "cancel", {}, "ok")], final_answer="done" + ) empty_ans = UnifiedRollout( - turns=[UnifiedTurn("", "cancel", {}, "ok"), - UnifiedTurn("", "refund", {}, "ok")], final_answer="") + turns=[ + UnifiedTurn("", "cancel", {}, "ok"), + UnifiedTurn("", "refund", {}, "ok"), + ], + final_answer="", + ) assert gold_coverage_verify(t, good) is True assert gold_coverage_verify(t, missing) is False assert gold_coverage_verify(t, empty_ans) is False @@ -123,10 +143,16 @@ def test_gold_coverage_verify(): def test_generate_sft_dataset_end_to_end(tmp_path): tools = default_catalog() - tasks = [normalize_row(_RAW_ROW), normalize_row({ - **_RAW_ROW, "id": "unsolvable", - "evaluation_criteria": {"actions": [{"name": "never_called"}]}, - })] + tasks = [ + normalize_row(_RAW_ROW), + normalize_row( + { + **_RAW_ROW, + "id": "unsolvable", + "evaluation_criteria": {"actions": [{"name": "never_called"}]}, + } + ), + ] def rollout_fn(task): # Solve the first task; always miss the gold action of the second. @@ -139,18 +165,26 @@ def test_generate_sft_dataset_end_to_end(tmp_path): ], # num_tool_calls must reflect the two routed calls: the structural # _target_is_clean gate drops a trajectory with num_tool_calls < 1. - final_answer="refunded $20.90", cost_usd=0.03, num_tool_calls=2, + final_answer="refunded $20.90", + cost_usd=0.03, + num_tool_calls=2, ) - return UnifiedRollout(turns=[UnifiedTurn("x", "cancel", {}, "ok")], - final_answer="nope", cost_usd=0.05) + return UnifiedRollout( + turns=[UnifiedTurn("x", "cancel", {}, "ok")], + final_answer="nope", + cost_usd=0.05, + ) out = tmp_path / "sft.jsonl" stats = generate_sft_dataset( - str(out), tasks=tasks, tools=tools, rollout_fn=rollout_fn, + str(out), + tasks=tasks, + tools=tools, + rollout_fn=rollout_fn, samples_per_task=2, ) assert stats["tasks_seen"] == 2 - assert stats["records_written"] == 1 # only the solvable task + assert stats["records_written"] == 1 # only the solvable task assert stats["tasks_dropped"] == 1 lines = out.read_text().strip().splitlines() assert len(lines) == 1 diff --git a/tests/test_orchestrator_learning/test_build_orchestrator_sft.py b/tests/test_orchestrator_learning/test_build_orchestrator_sft.py index 696d0ad3..f088691b 100644 --- a/tests/test_orchestrator_learning/test_build_orchestrator_sft.py +++ b/tests/test_orchestrator_learning/test_build_orchestrator_sft.py @@ -30,11 +30,23 @@ def _canned_rollout(task: Task) -> UnifiedRollout: # One tool turn + a final-answer turn that echoes the gold answer. return UnifiedRollout( turns=[ - UnifiedTurn(reasoning="let me compute", tool_name="code_interpreter", - arguments={"code": "print(6*7)"}, observation="42"), - UnifiedTurn(reasoning=f"The tool returned {task.answer}. FINAL_ANSWER: {task.answer}", tool_name=None), + UnifiedTurn( + reasoning="let me compute", + tool_name="code_interpreter", + arguments={"code": "print(6*7)"}, + observation="42", + ), + UnifiedTurn( + reasoning=( + f"The tool returned {task.answer}. FINAL_ANSWER: {task.answer}" + ), + tool_name=None, + ), ], - final_answer=task.answer, cost_usd=0.01, tokens=20, num_tool_calls=1, + final_answer=task.answer, + cost_usd=0.01, + tokens=20, + num_tool_calls=1, ) @@ -44,7 +56,7 @@ def test_build_v1_writes_expected_conversations_shape(tmp_path, monkeypatch): # Accept-all verifier (the real make_verifier may hit Gemini / math checkers). monkeypatch.setattr( "openjarvis.learning.intelligence.orchestrator.sft_data.verify.make_verifier", - lambda: (lambda task, rollout: True), + lambda: lambda task, rollout: True, ) from openjarvis.learning.intelligence.orchestrator.sft_data.verify import ( make_verifier, @@ -105,23 +117,36 @@ def test_driver_runs_end_to_end_with_fakes(tmp_path, monkeypatch): drv = _load_driver() monkeypatch.setattr(drv, "load_sft_tasks", _fake_tasks) - monkeypatch.setattr(drv, "run_unified_rollout", - lambda question, tools, **kw: _canned_rollout(_fake_tasks()[0])) - monkeypatch.setattr(drv, "make_verifier", lambda: (lambda task, rollout: True)) + monkeypatch.setattr( + drv, + "run_unified_rollout", + lambda question, tools, **kw: _canned_rollout(_fake_tasks()[0]), + ) + monkeypatch.setattr(drv, "make_verifier", lambda: lambda task, rollout: True) # make_call_orchestrator builds an OpenAI client lazily, so it never connects # here (run_unified_rollout is stubbed out). # The driver treats --out as a TAG and always writes to - # data/orchestrator/raw/