-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy pathrunner.py
More file actions
255 lines (199 loc) · 9.1 KB
/
Copy pathrunner.py
File metadata and controls
255 lines (199 loc) · 9.1 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
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
import _io
import copy
import os
import sys
import traceback
import ray
from ray import ObjectRef
from ray.actor import ActorHandle
from datetime import datetime
import tyro
import json
import orjson
from tqdm import tqdm
from openai import OpenAI
from dataclasses import dataclass, field, fields, asdict
from typing import Any, Set, Dict, List, Tuple, Union, Optional, Literal, TypedDict, NamedTuple, Iterable
from transformers import PreTrainedTokenizer, AutoTokenizer
from src.inference.define import SamplingParams, DeleteAction, ExecuteStepRecord, ExecuteResult
from src.inference.executor import execute_solving_question
def load_jsonline(fp: str, limit: int = -1) -> List[Any]:
items = []
with open(fp, 'r', encoding='utf-8') as f:
for idx, i in enumerate(f):
if limit != -1 and idx > limit:
break
items.append(orjson.loads(i))
return items
def write_jsonline(fp: str, obj: List[Any]):
with open(fp, 'wb') as f:
for i in obj:
f.write(orjson.dumps(i) + b"\n")
def load_json(fp: str) -> Any:
with open(fp, 'rb') as f:
return orjson.loads(f.read())
def write_json(fp: str, obj: Any):
with open(fp, 'w', encoding='utf-8') as f:
f.write(json.dumps(obj, ensure_ascii=False, indent=4))
#
DEL_SYSTEM_PROMPT = """
Role: AI Reasoning Analyst
Your task is to act as an AI Reasoning Analyst. You will be given a Chain-of-Thought (CoT) reasoning and your goal is to identify and mark for deletion any paragraphs from the CoT reasoning that are redundant or irrelevant.
### Instructions:
1. **Analyze the CoT Reasoning:** Carefully read and understand the reasoning, context, and goals of the CoT reasoning.
2. **Identify Redundant or Irrelevant Paragraphs:** A paragraph is considered redundant or irrelevant if it does not directly support or inform the overall reasoning process. This could include tangential thoughts, corrected errors, or superseded lines of reasoning.
3. **Mark for Deletion:** For every irrelevant paragraph you detect, generate a JSON object that uniquely identifies it. Each object must include a `prefix` and a `suffix` extracted from the paragraph. The paragraph to be removed is defined as the text spanning from the `prefix` through to the `suffix`, inclusive. Both `prefix` and `suffix` must be non-empty and together **uniquely specify** the target paragraph.
4. **Format the Output:** Your final output must be a single JSON object containing a list of these `prefix`/`suffix` objects.
5. <DELETED> in the CoT reasoning indicates that the content has already been removed.
6. It is possible that you don't need to delete any paragraphs in the CoT reasoning.
### Constraints:
* If no paragraphs in the CoT reasoning are redundant or irrelevant, output an empty list: `[]`.
* If multiple paragraphs are redundant or irrelevant, include a separate JSON object for each one in the list. The list can be non-continuous.
* Do not include any explanatory text in your output, only the JSON.
* If two adjacent paragraphs are redundant or irrelevant, include both in the list as a single object with a combined prefix and suffix.
* Do not delete the first or the last paragraph in the CoT reasoning.
* Limit the number of deletions: aim to delete one or two paragraphs, with a maximum of three.
""".strip()
DEL_USER_PROMPT = """
### CoT Reasoning:
{{GENERATION}}
""".strip()
def get_timestamp() -> str:
now = datetime.now()
filename = now.strftime(f"%Y_%m_%d_%H_%M_%S")
return filename
def load_existed_samples(dir_path: str, ds_name: str) -> List[ExecuteResult]:
items = []
for fname in os.listdir(dir_path):
if fname.startswith(ds_name) and fname.endswith(".jsonl"):
fp = os.path.join(dir_path, fname)
items.extend(load_jsonline(fp=fp))
items = [ExecuteResult.from_dict(obj=i) for i in items]
items = [r for r in items if r.steps[-1].generate_success]
return items
@dataclass
class Args:
ds_name: Literal["aime_2425", "brumo25", "hmmt25", "imo"]
input: str
output_dir: str = field(default="")
gen_model_path: str = field(default="")
del_model_path: str = field(default="")
adapter_name: str = field(default="lora")
use_json_schema: bool = field(default=True)
max_token_per_generation: int = field(default=5000)
max_gen_model_len: int = field(default=32 * 1024)
max_del_model_len: int = field(default=32 * 1024)
num_delete_times: int = field(default=50)
temperature: Optional[float] = field(default=None)
top_p: Optional[float] = field(default=None)
top_k: Optional[int] = field(default=None)
num_samples: Optional[int] = field(default=1)
retry: int = field(default=3)
@property
def output(self) -> str:
return os.path.join(
self.output_dir,
self.ds_name,
f"{self.ds_name}_schema-{self.use_json_schema}_{self.max_token_per_generation}-{self.max_gen_model_len}-{self.max_del_model_len}_{get_timestamp()}.jsonl"
)
@ray.remote
class GenerateActor:
def __init__(
self,
url: str,
del_system_prompt: str,
del_user_prompt: str,
args: Args,
):
self.args = args
self.del_system_prompt = del_system_prompt
self.del_user_prompt = del_user_prompt
self.client = OpenAI(api_key="", base_url=url)
self.gen_tokenizer: PreTrainedTokenizer = AutoTokenizer.from_pretrained(
pretrained_model_name_or_path=self.args.gen_model_path
)
self.del_tokenizer: PreTrainedTokenizer = AutoTokenizer.from_pretrained(
pretrained_model_name_or_path=self.args.del_model_path
)
self.sampling_parms = SamplingParams(
temperature=args.temperature,
top_p=args.top_p,
top_k=args.top_k
)
def process(self, ins: Dict[str, Any]) -> ExecuteResult:
response, steps = execute_solving_question(
gen_client=self.client,
del_client=self.client,
model=self.args.gen_model_path,
adapter_name=self.args.adapter_name,
question=ins["prompt"],
sampling_params=self.sampling_parms,
max_token_per_generation=self.args.max_token_per_generation,
max_gen_model_len=self.args.max_gen_model_len,
del_system_prompt=self.del_system_prompt,
del_user_prompt=self.del_user_prompt,
max_del_model_len=self.args.max_del_model_len,
use_json_schema=self.args.use_json_schema,
num_delete_times=self.args.num_delete_times,
del_tokenizer=self.del_tokenizer,
gen_tokenizer=self.gen_tokenizer,
retry=self.args.retry
)
ins.update({"response": response, "steps": []})
result = ExecuteResult.from_dict(obj=ins)
result.steps = steps
return result
def process_ready_handlers(handlers: List[Any], writer: _io.TextIOWrapper):
for h in handlers:
try:
r: ExecuteResult = ray.get(h)
if not r.steps[-1].generate_success:
# 说明这个问题失败了
continue
safe_str = json.dumps(r.to_dict(), ensure_ascii=False)
safe_str = safe_str.encode("utf-8", "ignore").decode()
writer.write(safe_str + "\n")
except Exception:
print(f"[!] process_ready_handlers | error: {traceback.format_exc()}")
writer.flush()
if __name__ == '__main__':
args: Args = tyro.cli(Args)
os.system(f"mkdir -p {os.path.join(args.output_dir, args.ds_name)}")
service_ips = [
("<ip>", "<port>")
]
actors: List[ActorHandle] = []
for i in range(0, len(service_ips)):
ip, port = service_ips[i % len(service_ips)]
actors.append(GenerateActor.options(max_concurrency=16).remote(
url=f"http://{ip}:{port}/v1",
del_system_prompt=DEL_SYSTEM_PROMPT,
del_user_prompt=DEL_USER_PROMPT,
args=args,
))
instances = load_jsonline(fp=args.input)
print("Num Instances: ", len(instances))
samples = []
for i in range(0, len(instances)):
for _ in range(0, args.num_samples):
samples.append(copy.deepcopy(instances[i]))
instances = samples
del samples
handlers: List[ObjectRef] = []
pbr = tqdm(total=len(instances), desc=f"Process {args.ds_name} -> {args.output}")
writer = open(args.output, "w", encoding="utf-8")
a_cur = 0
for i in range(0, len(instances)):
h: ObjectRef = actors[i % len(actors)].process.remote(ins=instances[i])
handlers.append(h)
while len(handlers) > 480:
ready_handlers, handlers = ray.wait(handlers, num_returns=min(1, len(handlers)))
process_ready_handlers(handlers=ready_handlers, writer=writer)
pbr.update(n=len(ready_handlers))
print("Num Handlers: ", len(handlers))
while len(handlers) > 0:
ready_handlers, handlers = ray.wait(handlers, num_returns=min(1, len(handlers)))
process_ready_handlers(handlers=ready_handlers, writer=writer)
pbr.update(n=len(ready_handlers))
pbr.close()
writer.close()