MCPcopy Create free account
hub / github.com/THUDM/AgentTuning / get_data_split

Function get_data_split

AgentBench.old/src/tasks/mind2web/dataloader.py:219–278  ·  view source on GitHub ↗
(data_dir, split_file, candidate_results=None, cache_dir=None, is_train=False, is_debug=False)

Source from the content-addressed store, hash-verified

217
218
219def get_data_split(data_dir, split_file, candidate_results=None, cache_dir=None, is_train=False, is_debug=False):
220 def flatten_actions(samples):
221 outputs = {
222 "website": [],
223 "confirmed_task": [],
224 "annotation_id": [],
225 "previous_actions": [],
226 "action_uid": [],
227 "operation": [],
228 "pos_candidates": [],
229 "neg_candidates": [],
230 "cleaned_html": [],
231 }
232 num_actions = [len(actions) for actions in samples["actions"]]
233 for key in ["website", "confirmed_task", "annotation_id"]:
234 for idx, value in enumerate(samples[key]):
235 outputs[key] += [value] * num_actions[idx]
236 for actions, action_reprs in zip(samples["actions"], samples["action_reprs"]):
237 for a_idx, action in enumerate(actions):
238 outputs["previous_actions"].append(action_reprs[:a_idx])
239 for key in [
240 "action_uid",
241 "operation",
242 "pos_candidates",
243 "neg_candidates",
244 "cleaned_html",
245 ]:
246 outputs[key].append(action[key])
247 return outputs
248 dataset = load_dataset(data_dir, data_files=split_file, split="all", cache_dir=cache_dir)
249 if is_debug:
250 dataset = dataset.select(range(3))
251 flatten_dataset = dataset.map(
252 flatten_actions,
253 batched=True,
254 remove_columns=dataset.column_names,
255 batch_size=10,
256 num_proc=4,
257 )
258 if candidate_results is not None:
259 candidate_scores = candidate_results["scores"]
260 candidate_ranks = candidate_results["ranks"]
261
262 def get_score(sample):
263 sample_id = f"{sample['annotation_id']}_{sample['action_uid']}"
264 for candidates in [sample["pos_candidates"], sample["neg_candidates"]]:
265 for candidate in candidates:
266 candidate_id = candidate["backend_node_id"]
267 candidate["score"] = candidate_scores[sample_id][candidate_id]
268 candidate["rank"] = candidate_ranks[sample_id][candidate_id]
269 return {
270 "pos_candidates": sample["pos_candidates"],
271 "neg_candidates": sample["neg_candidates"],
272 }
273
274 flatten_dataset = flatten_dataset.map(get_score)
275 if is_train:
276 flatten_dataset = flatten_dataset.filter(lambda x: len(x["pos_candidates"]) > 0)

Callers 1

__init__Method · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected