| 217 | |
| 218 | |
| 219 | def 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) |