File size: 4,667 Bytes
116524e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
#!/usr/bin/env python3
"""Merge partial TAU-bench re-run results into an existing detailed JSON.



Replaces failed/dead task entries in the original with fresh results from

a partial re-run, then recomputes pass^k metrics over all tasks.



Usage:

    uv run python scripts/merge_tau_results.py \

        --original tau_benchmark_results/tau_airline_..._223116_detailed.json \

        --patch    tau_benchmark_results/tau_airline_..._HHMMSS_detailed.json \

        --output   tau_benchmark_results/tau_airline_..._merged

"""

from __future__ import annotations

import argparse
import json
from math import comb
from pathlib import Path


def pass_hat_k(n: int, s: int, k: int) -> float:
    """Compute pass^k = C(s, k) / C(n, k).



    Args:

        n: Total number of trials.

        s: Number of successes.

        k: k value for pass^k.



    Returns:

        Combinatorial probability that all k chosen trials succeed.

    """
    if k > n or k > s:
        return 0.0
    return comb(s, k) / comb(n, k)


def merge(original: dict, patch: dict) -> dict:
    """Replace task entries in original with matching entries from patch.



    Matching is by task_id. Only tasks present in the patch are replaced.

    """
    patch_by_id = {r["task_id"]: r for r in patch["results"]}

    merged_results = []
    replaced = []
    for task_result in original["results"]:
        tid = task_result["task_id"]
        if tid in patch_by_id:
            merged_results.append(patch_by_id[tid])
            replaced.append(tid)
        else:
            merged_results.append(task_result)

    print(f"Replaced {len(replaced)} tasks: {replaced}")

    # Recompute pass^k metrics
    k = original["k"]
    n_tasks = len(merged_results)
    pass_sums = {str(j): 0.0 for j in range(1, k + 1)}

    for task_result in merged_results:
        trials = task_result["trials"]
        n_trials = len(trials)
        n_successes = sum(1 for t in trials if t.get("success", False))

        # Recompute per-task pass_k_values
        task_pass_k = {}
        for j in range(1, k + 1):
            task_pass_k[str(j)] = pass_hat_k(n_trials, n_successes, j)
        task_result["pass_k_values"] = task_pass_k
        task_result["passed_all"] = all(t.get("success", False) for t in trials)

        for j in range(1, k + 1):
            pass_sums[str(j)] += task_pass_k[str(j)]

    metrics = {}
    for j in range(1, k + 1):
        metrics[f"pass_{j}"] = pass_sums[str(j)] / n_tasks if n_tasks > 0 else 0.0

    merged = {
        "tasks_evaluated": n_tasks,
        "k": k,
        "pass_sums": pass_sums,
        "metrics": metrics,
        "results": merged_results,
    }

    return merged


def main() -> None:
    parser = argparse.ArgumentParser(
        description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter
    )
    parser.add_argument(
        "--original", required=True, help="Path to the original detailed JSON"
    )
    parser.add_argument(
        "--patch", required=True, help="Path to the partial re-run detailed JSON"
    )
    parser.add_argument(
        "--output",
        required=True,
        help="Output path prefix (will create _detailed.json and _summary.json)",
    )
    args = parser.parse_args()

    original = json.loads(Path(args.original).read_text())
    patch = json.loads(Path(args.patch).read_text())

    merged = merge(original, patch)

    # Validate: no tasks with 0 steps remaining
    zero_step_tasks = [
        r["task_id"]
        for r in merged["results"]
        if all(t.get("steps", 0) == 0 for t in r["trials"])
    ]
    if zero_step_tasks:
        print(
            f"WARNING: {len(zero_step_tasks)} tasks still have all-zero steps: {zero_step_tasks}"
        )

    # Save detailed
    detailed_path = Path(f"{args.output}_detailed.json")
    detailed_path.write_text(json.dumps(merged, indent=2, default=str))
    print(f"Saved detailed: {detailed_path}")

    # Save summary
    summary = {
        "tasks_evaluated": merged["tasks_evaluated"],
        "k": merged["k"],
        "pass_sums": merged["pass_sums"],
        "metrics": merged["metrics"],
    }
    summary_path = Path(f"{args.output}_summary.json")
    summary_path.write_text(json.dumps(summary, indent=2))
    print(f"Saved summary: {summary_path}")

    # Print metrics
    print(f"\nMerged pass^k metrics ({merged['tasks_evaluated']} tasks):")
    for j in range(1, merged["k"] + 1):
        print(f"  pass^{j}: {merged['metrics'][f'pass_{j}']:.2%}")


if __name__ == "__main__":
    main()