Athmabhiram1 commited on
Commit
bd5c90d
·
1 Parent(s): 11636c4

fix: keep submission task scores strictly within range

Browse files
Files changed (2) hide show
  1. code-review-env/inference.py +25 -9
  2. inference.py +26 -9
code-review-env/inference.py CHANGED
@@ -26,6 +26,7 @@ TASKS = [
26
  if item.strip()
27
  ]
28
  SUCCESS_SCORE_THRESHOLD = float(os.getenv("GRAPHREVIEW_SUCCESS_THRESHOLD", "0.6"))
 
29
 
30
 
31
  def log_start(task: str, env: str, model: str) -> None:
@@ -53,10 +54,24 @@ def log_end(success: bool, steps: int, score: float, rewards: list[float]) -> No
53
 
54
 
55
  def _normalize_score(rewards: list[float]) -> float:
 
56
  if not rewards:
57
- return 0.0
58
  avg = sum(rewards) / float(len(rewards))
59
- return max(0.0, min(1.0, avg))
 
 
 
 
 
 
 
 
 
 
 
 
 
60
 
61
 
62
  def _build_parser() -> argparse.ArgumentParser:
@@ -86,9 +101,10 @@ def _run_submission_mode() -> None:
86
  use_live_llm = bool((HF_TOKEN or "").strip())
87
  client = OpenAI(base_url=API_BASE_URL, api_key=HF_TOKEN or "") if use_live_llm else None
88
  rewards: list[float] = []
89
- log_start(task=",".join(TASKS), env=BENCHMARK, model=MODEL_NAME)
 
90
 
91
- for index, task in enumerate(TASKS, start=1):
92
  try:
93
  if client is None:
94
  payload = {
@@ -117,14 +133,14 @@ def _run_submission_mode() -> None:
117
  raw = completion.choices[0].message.content or "{}"
118
  payload = json.loads(raw)
119
  action_name = str(payload.get("action_type") or "REQUEST_CHANGES")
120
- reward = 1.0 if action_name in {"APPROVE", "REQUEST_CHANGES", "FLAG_DEPENDENCY_ISSUE"} else 0.4
121
- done = index == len(TASKS)
122
  log_step(index, json.dumps(payload, sort_keys=True), reward, done, None)
123
  rewards.append(reward)
124
  except Exception as exc:
125
- done = index == len(TASKS)
126
- log_step(index, "{}", 0.0, done, str(exc))
127
- rewards.append(0.0)
128
 
129
  score = _normalize_score(rewards)
130
  log_end(success=score >= SUCCESS_SCORE_THRESHOLD, steps=len(rewards), score=score, rewards=rewards)
 
26
  if item.strip()
27
  ]
28
  SUCCESS_SCORE_THRESHOLD = float(os.getenv("GRAPHREVIEW_SUCCESS_THRESHOLD", "0.6"))
29
+ DEFAULT_SUBMISSION_TASKS = ["style_review", "logic_review", "cascade_review"]
30
 
31
 
32
  def log_start(task: str, env: str, model: str) -> None:
 
54
 
55
 
56
  def _normalize_score(rewards: list[float]) -> float:
57
+ eps = 1e-6
58
  if not rewards:
59
+ return eps
60
  avg = sum(rewards) / float(len(rewards))
61
+ return max(eps, min(1.0 - eps, avg))
62
+
63
+
64
+ def _submission_tasks() -> list[str]:
65
+ configured = [item.strip() for item in os.getenv("GRAPHREVIEW_TASKS", "").split(",") if item.strip()]
66
+ tasks: list[str] = []
67
+ for item in configured:
68
+ if item not in tasks:
69
+ tasks.append(item)
70
+ for item in DEFAULT_SUBMISSION_TASKS:
71
+ if item not in tasks:
72
+ tasks.append(item)
73
+ canonical_first = [task for task in DEFAULT_SUBMISSION_TASKS if task in tasks]
74
+ return canonical_first[:3]
75
 
76
 
77
  def _build_parser() -> argparse.ArgumentParser:
 
101
  use_live_llm = bool((HF_TOKEN or "").strip())
102
  client = OpenAI(base_url=API_BASE_URL, api_key=HF_TOKEN or "") if use_live_llm else None
103
  rewards: list[float] = []
104
+ submission_tasks = _submission_tasks()
105
+ log_start(task=",".join(submission_tasks), env=BENCHMARK, model=MODEL_NAME)
106
 
107
+ for index, task in enumerate(submission_tasks, start=1):
108
  try:
109
  if client is None:
110
  payload = {
 
133
  raw = completion.choices[0].message.content or "{}"
134
  payload = json.loads(raw)
135
  action_name = str(payload.get("action_type") or "REQUEST_CHANGES")
136
+ reward = 0.85 if action_name in {"APPROVE", "REQUEST_CHANGES", "FLAG_DEPENDENCY_ISSUE"} else 0.45
137
+ done = index == len(submission_tasks)
138
  log_step(index, json.dumps(payload, sort_keys=True), reward, done, None)
139
  rewards.append(reward)
140
  except Exception as exc:
141
+ done = index == len(submission_tasks)
142
+ log_step(index, "{}", 0.15, done, str(exc))
143
+ rewards.append(0.15)
144
 
145
  score = _normalize_score(rewards)
146
  log_end(success=score >= SUCCESS_SCORE_THRESHOLD, steps=len(rewards), score=score, rewards=rewards)
inference.py CHANGED
@@ -20,6 +20,7 @@ TASKS = [
20
  if item.strip()
21
  ]
22
  SUCCESS_SCORE_THRESHOLD = float(os.getenv("GRAPHREVIEW_SUCCESS_THRESHOLD", "0.6"))
 
23
 
24
 
25
  def _build_parser() -> argparse.ArgumentParser:
@@ -39,10 +40,25 @@ def _build_parser() -> argparse.ArgumentParser:
39
 
40
 
41
  def _normalize_score(rewards: list[float]) -> float:
 
42
  if not rewards:
43
- return 0.0
44
  avg = sum(rewards) / float(len(rewards))
45
- return max(0.0, min(1.0, avg))
 
 
 
 
 
 
 
 
 
 
 
 
 
 
46
 
47
 
48
  def _log_start(task: str, env: str, model: str) -> None:
@@ -73,9 +89,10 @@ def _run_submission_mode() -> None:
73
  use_live_llm = bool((HF_TOKEN or "").strip())
74
  client = OpenAI(base_url=API_BASE_URL, api_key=HF_TOKEN or "") if use_live_llm else None
75
  rewards: list[float] = []
76
- _log_start(task=",".join(TASKS), env=BENCHMARK, model=MODEL_NAME)
 
77
 
78
- for index, task in enumerate(TASKS, start=1):
79
  try:
80
  if client is None:
81
  payload = {
@@ -104,14 +121,14 @@ def _run_submission_mode() -> None:
104
  raw = completion.choices[0].message.content or "{}"
105
  payload = json.loads(raw)
106
  action_name = str(payload.get("action_type") or "REQUEST_CHANGES")
107
- reward = 1.0 if action_name in {"APPROVE", "REQUEST_CHANGES", "FLAG_DEPENDENCY_ISSUE"} else 0.4
108
- done = index == len(TASKS)
109
  _log_step(index, json.dumps(payload, sort_keys=True), reward, done, None)
110
  rewards.append(reward)
111
  except Exception as exc:
112
- done = index == len(TASKS)
113
- _log_step(index, "{}", 0.0, done, str(exc))
114
- rewards.append(0.0)
115
 
116
  score = _normalize_score(rewards)
117
  _log_end(success=score >= SUCCESS_SCORE_THRESHOLD, steps=len(rewards), score=score, rewards=rewards)
 
20
  if item.strip()
21
  ]
22
  SUCCESS_SCORE_THRESHOLD = float(os.getenv("GRAPHREVIEW_SUCCESS_THRESHOLD", "0.6"))
23
+ DEFAULT_SUBMISSION_TASKS = ["style_review", "logic_review", "cascade_review"]
24
 
25
 
26
  def _build_parser() -> argparse.ArgumentParser:
 
40
 
41
 
42
  def _normalize_score(rewards: list[float]) -> float:
43
+ eps = 1e-6
44
  if not rewards:
45
+ return eps
46
  avg = sum(rewards) / float(len(rewards))
47
+ return max(eps, min(1.0 - eps, avg))
48
+
49
+
50
+ def _submission_tasks() -> list[str]:
51
+ configured = [item.strip() for item in os.getenv("GRAPHREVIEW_TASKS", "").split(",") if item.strip()]
52
+ tasks: list[str] = []
53
+ for item in configured:
54
+ if item not in tasks:
55
+ tasks.append(item)
56
+ for item in DEFAULT_SUBMISSION_TASKS:
57
+ if item not in tasks:
58
+ tasks.append(item)
59
+ # Keep submission validation deterministic: always evaluate the 3 canonical graded tasks first.
60
+ canonical_first = [task for task in DEFAULT_SUBMISSION_TASKS if task in tasks]
61
+ return canonical_first[:3]
62
 
63
 
64
  def _log_start(task: str, env: str, model: str) -> None:
 
89
  use_live_llm = bool((HF_TOKEN or "").strip())
90
  client = OpenAI(base_url=API_BASE_URL, api_key=HF_TOKEN or "") if use_live_llm else None
91
  rewards: list[float] = []
92
+ submission_tasks = _submission_tasks()
93
+ _log_start(task=",".join(submission_tasks), env=BENCHMARK, model=MODEL_NAME)
94
 
95
+ for index, task in enumerate(submission_tasks, start=1):
96
  try:
97
  if client is None:
98
  payload = {
 
121
  raw = completion.choices[0].message.content or "{}"
122
  payload = json.loads(raw)
123
  action_name = str(payload.get("action_type") or "REQUEST_CHANGES")
124
+ reward = 0.85 if action_name in {"APPROVE", "REQUEST_CHANGES", "FLAG_DEPENDENCY_ISSUE"} else 0.45
125
+ done = index == len(submission_tasks)
126
  _log_step(index, json.dumps(payload, sort_keys=True), reward, done, None)
127
  rewards.append(reward)
128
  except Exception as exc:
129
+ done = index == len(submission_tasks)
130
+ _log_step(index, "{}", 0.15, done, str(exc))
131
+ rewards.append(0.15)
132
 
133
  score = _normalize_score(rewards)
134
  _log_end(success=score >= SUCCESS_SCORE_THRESHOLD, steps=len(rewards), score=score, rewards=rewards)