-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmain.py
More file actions
151 lines (125 loc) · 4.88 KB
/
Copy pathmain.py
File metadata and controls
151 lines (125 loc) · 4.88 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
"""
main.py
CLI入口 — SAGE-v2
用法:
python main.py demo # 跑1题,verbose,验证启动
python main.py batch --n 10 # 跑前10题,串行
python main.py batch --n 100 --concurrency 2 # 并发2
python main.py baseline --n 100 # 跑基线并保存结果JSON
"""
from __future__ import annotations
import sys
import os
import asyncio
import json
import argparse
from datetime import datetime
from dotenv import load_dotenv
load_dotenv()
def cmd_demo(args):
"""跑1题,全verbose,用于验证启动"""
from data.loader import load_livecodebench
from pipeline import run_batch, print_summary
print("=" * 55)
print(" SAGE-v2 demo(1题)")
print("=" * 55)
problems = load_livecodebench(max_problems=1)
results = asyncio.run(run_batch(problems, verbose=True))
print_summary(results)
def cmd_batch(args):
"""批量跑n题"""
from data.loader import load_livecodebench
from pipeline import run_batch, print_summary
sort = not args.no_sort
print(f"[main] batch模式 n={args.n} concurrency={args.concurrency} sort_by_difficulty={sort}")
problems = load_livecodebench(max_problems=args.n, sort_by_difficulty=sort)
results = asyncio.run(run_batch(
problems,
verbose=args.verbose,
concurrency=args.concurrency,
))
print_summary(results)
def cmd_baseline(args):
"""跑基线并保存结果到data/"""
from data.loader import load_livecodebench
from pipeline import run_batch, print_summary
import os
sort = not args.no_sort
print(f"[main] baseline模式 n={args.n} sort_by_difficulty={sort}")
problems = load_livecodebench(max_problems=args.n, sort_by_difficulty=sort)
results = asyncio.run(run_batch(
problems,
verbose=args.verbose,
concurrency=args.concurrency,
))
print_summary(results)
# 保存结果
passed_count = sum(1 for r in results if r.passed)
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
out_path = os.path.join("data", f"results_livecodebench_{timestamp}.json")
total_cost = sum(r.cost_usd for r in results)
payload = {
"baseline": "B1",
"sort_by_difficulty": sort,
"dataset": "livecodebench",
"model_generator": os.getenv("GENERATOR_MODEL", "deepseek-chat"),
"model_reviewer": os.getenv("REVIEWER_MODEL", "deepseek-reasoner"),
"n": len(results),
"pass_at_1": round(passed_count / len(results), 4) if results else 0,
"passed": passed_count,
"total": len(results),
"total_cost_usd": round(total_cost, 4),
"cost_per_problem_usd": round(total_cost / len(results), 5) if results else 0,
"timestamp": datetime.now().isoformat(),
"results": [
{
"task_id": r.task_id,
"passed": r.passed,
"difficulty": r.difficulty,
"elapsed_s": round(r.elapsed_s, 2),
"changed_by_reviewer": r.changed_by_reviewer,
"gen_tokens_in": r.gen_tokens_in,
"gen_tokens_out": r.gen_tokens_out,
"rev_tokens_in": r.rev_tokens_in,
"rev_tokens_out": r.rev_tokens_out,
"cost_usd": round(r.cost_usd, 5),
"error": r.error,
}
for r in results
],
}
with open(out_path, "w", encoding="utf-8") as f:
json.dump(payload, f, ensure_ascii=False, indent=2)
print(f"\n[main] 结果已保存至 {out_path}")
def build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(
prog="sage",
description="SAGE-v2 Memory-Augmented Code Review Pipeline",
)
sub = parser.add_subparsers(dest="cmd", required=True)
# demo
sub.add_parser("demo", help="跑1题验证启动")
# batch
p_batch = sub.add_parser("batch", help="批量跑n题")
p_batch.add_argument("--n", type=int, default=10)
p_batch.add_argument("--concurrency", type=int, default=1)
p_batch.add_argument("--verbose", action="store_true", default=True)
p_batch.add_argument("--no-sort", action="store_true", default=False,
help="不按难度排序(默认按easy→medium→hard)")
# baseline
p_base = sub.add_parser("baseline", help="跑基线并保存JSON")
p_base.add_argument("--n", type=int, default=100)
p_base.add_argument("--concurrency", type=int, default=1)
p_base.add_argument("--verbose", action="store_true", default=False)
p_base.add_argument("--no-sort", action="store_true", default=False,
help="不按难度排序(默认按easy→medium→hard)")
return parser
if __name__ == "__main__":
parser = build_parser()
args = parser.parse_args()
if args.cmd == "demo":
cmd_demo(args)
elif args.cmd == "batch":
cmd_batch(args)
elif args.cmd == "baseline":
cmd_baseline(args)