diff --git a/datasets/agenttuning_alfworld/extract_raw.py b/datasets/agenttuning_alfworld/extract_raw.py index 29cdd3c3..b5b501d3 100644 --- a/datasets/agenttuning_alfworld/extract_raw.py +++ b/datasets/agenttuning_alfworld/extract_raw.py @@ -16,12 +16,6 @@ "assistant": "assistant", "system": "system", } -SAMPLE_IDS = { - "alfworld": ["alfworld_155", "alfworld_219", "alfworld_56", "alfworld_58", "alfworld_149"], - "db": ["db_100", "db_356", "db_493", "db_516", "db_531"], - "kg": ["kg_132", "kg_260", "kg_164", "kg_221", "kg_12"], - "webshop": ["webshop_330", "webshop_190", "webshop_243", "webshop_89", "webshop_26"], -} def json_safe(value: Any) -> Any: @@ -50,18 +44,11 @@ def has_valid_conversation(sample: dict[str, Any]) -> bool: def main(config_name: str) -> None: try: - dataset = load_dataset("THUDM/AgentInstruct")[config_name] - sample_ids = SAMPLE_IDS.get(config_name) - if sample_ids is not None: - selected = {} - sample_id_set = set(sample_ids) - for sample in dataset: - if sample.get("id") in sample_id_set and has_valid_conversation(sample): - selected[sample["id"]] = json_safe(sample) - for sample_id in sample_ids: - if sample_id in selected: - print(json.dumps(selected[sample_id], ensure_ascii=False)) - return + dataset = load_dataset( + "parquet", + data_files=f"hf://datasets/THUDM/AgentInstruct/data/{config_name}-*.parquet", + split="train", + ) for sample in dataset: if not has_valid_conversation(sample): continue diff --git a/datasets/agenttuning_db/extract_raw.py b/datasets/agenttuning_db/extract_raw.py index d8c9f8f9..fd73731e 100644 --- a/datasets/agenttuning_db/extract_raw.py +++ b/datasets/agenttuning_db/extract_raw.py @@ -16,12 +16,6 @@ "assistant": "assistant", "system": "system", } -SAMPLE_IDS = { - "alfworld": ["alfworld_155", "alfworld_219", "alfworld_56", "alfworld_58", "alfworld_149"], - "db": ["db_100", "db_356", "db_493", "db_516", "db_531"], - "kg": ["kg_132", "kg_260", "kg_164", "kg_221", "kg_12"], - "webshop": ["webshop_330", "webshop_190", "webshop_243", "webshop_89", "webshop_26"], -} def json_safe(value: Any) -> Any: @@ -50,18 +44,11 @@ def has_valid_conversation(sample: dict[str, Any]) -> bool: def main(config_name: str) -> None: try: - dataset = load_dataset("THUDM/AgentInstruct")[config_name] - sample_ids = SAMPLE_IDS.get(config_name) - if sample_ids is not None: - selected = {} - sample_id_set = set(sample_ids) - for sample in dataset: - if sample.get("id") in sample_id_set and has_valid_conversation(sample): - selected[sample["id"]] = json_safe(sample) - for sample_id in sample_ids: - if sample_id in selected: - print(json.dumps(selected[sample_id], ensure_ascii=False)) - return + dataset = load_dataset( + "parquet", + data_files=f"hf://datasets/THUDM/AgentInstruct/data/{config_name}-*.parquet", + split="train", + ) for sample in dataset: if not has_valid_conversation(sample): continue diff --git a/datasets/agenttuning_db/metadata.json b/datasets/agenttuning_db/metadata.json index dfa26e78..41212e45 100644 --- a/datasets/agenttuning_db/metadata.json +++ b/datasets/agenttuning_db/metadata.json @@ -1,5 +1,7 @@ { "custom_tools": [], - "code_enabled": [], + "code_enabled": [ + "bash" + ], "browser_enabled": false } diff --git a/datasets/agenttuning_kg/extract_raw.py b/datasets/agenttuning_kg/extract_raw.py index 1eea87d5..ef31d786 100644 --- a/datasets/agenttuning_kg/extract_raw.py +++ b/datasets/agenttuning_kg/extract_raw.py @@ -16,12 +16,6 @@ "assistant": "assistant", "system": "system", } -SAMPLE_IDS = { - "alfworld": ["alfworld_155", "alfworld_219", "alfworld_56", "alfworld_58", "alfworld_149"], - "db": ["db_100", "db_356", "db_493", "db_516", "db_531"], - "kg": ["kg_132", "kg_260", "kg_164", "kg_221", "kg_12"], - "webshop": ["webshop_330", "webshop_190", "webshop_243", "webshop_89", "webshop_26"], -} def json_safe(value: Any) -> Any: @@ -50,18 +44,11 @@ def has_valid_conversation(sample: dict[str, Any]) -> bool: def main(config_name: str) -> None: try: - dataset = load_dataset("THUDM/AgentInstruct")[config_name] - sample_ids = SAMPLE_IDS.get(config_name) - if sample_ids is not None: - selected = {} - sample_id_set = set(sample_ids) - for sample in dataset: - if sample.get("id") in sample_id_set and has_valid_conversation(sample): - selected[sample["id"]] = json_safe(sample) - for sample_id in sample_ids: - if sample_id in selected: - print(json.dumps(selected[sample_id], ensure_ascii=False)) - return + dataset = load_dataset( + "parquet", + data_files=f"hf://datasets/THUDM/AgentInstruct/data/{config_name}-*.parquet", + split="train", + ) for sample in dataset: if not has_valid_conversation(sample): continue diff --git a/datasets/agenttuning_mind2web/extract_raw.py b/datasets/agenttuning_mind2web/extract_raw.py index ac053480..b481de34 100644 --- a/datasets/agenttuning_mind2web/extract_raw.py +++ b/datasets/agenttuning_mind2web/extract_raw.py @@ -16,12 +16,6 @@ "assistant": "assistant", "system": "system", } -SAMPLE_IDS = { - "alfworld": ["alfworld_155", "alfworld_219", "alfworld_56", "alfworld_58", "alfworld_149"], - "db": ["db_100", "db_356", "db_493", "db_516", "db_531"], - "kg": ["kg_132", "kg_260", "kg_164", "kg_221", "kg_12"], - "webshop": ["webshop_330", "webshop_190", "webshop_243", "webshop_89", "webshop_26"], -} def json_safe(value: Any) -> Any: @@ -50,18 +44,11 @@ def has_valid_conversation(sample: dict[str, Any]) -> bool: def main(config_name: str) -> None: try: - dataset = load_dataset("THUDM/AgentInstruct")[config_name] - sample_ids = SAMPLE_IDS.get(config_name) - if sample_ids is not None: - selected = {} - sample_id_set = set(sample_ids) - for sample in dataset: - if sample.get("id") in sample_id_set and has_valid_conversation(sample): - selected[sample["id"]] = json_safe(sample) - for sample_id in sample_ids: - if sample_id in selected: - print(json.dumps(selected[sample_id], ensure_ascii=False)) - return + dataset = load_dataset( + "parquet", + data_files=f"hf://datasets/THUDM/AgentInstruct/data/{config_name}-*.parquet", + split="train", + ) for sample in dataset: if not has_valid_conversation(sample): continue diff --git a/datasets/agenttuning_os/extract_raw.py b/datasets/agenttuning_os/extract_raw.py index b0e634bc..1fd1c288 100644 --- a/datasets/agenttuning_os/extract_raw.py +++ b/datasets/agenttuning_os/extract_raw.py @@ -16,12 +16,6 @@ "assistant": "assistant", "system": "system", } -SAMPLE_IDS = { - "alfworld": ["alfworld_155", "alfworld_219", "alfworld_56", "alfworld_58", "alfworld_149"], - "db": ["db_100", "db_356", "db_493", "db_516", "db_531"], - "kg": ["kg_132", "kg_260", "kg_164", "kg_221", "kg_12"], - "webshop": ["webshop_330", "webshop_190", "webshop_243", "webshop_89", "webshop_26"], -} def json_safe(value: Any) -> Any: @@ -50,18 +44,11 @@ def has_valid_conversation(sample: dict[str, Any]) -> bool: def main(config_name: str) -> None: try: - dataset = load_dataset("THUDM/AgentInstruct")[config_name] - sample_ids = SAMPLE_IDS.get(config_name) - if sample_ids is not None: - selected = {} - sample_id_set = set(sample_ids) - for sample in dataset: - if sample.get("id") in sample_id_set and has_valid_conversation(sample): - selected[sample["id"]] = json_safe(sample) - for sample_id in sample_ids: - if sample_id in selected: - print(json.dumps(selected[sample_id], ensure_ascii=False)) - return + dataset = load_dataset( + "parquet", + data_files=f"hf://datasets/THUDM/AgentInstruct/data/{config_name}-*.parquet", + split="train", + ) for sample in dataset: if not has_valid_conversation(sample): continue diff --git a/datasets/agenttuning_webshop/extract_raw.py b/datasets/agenttuning_webshop/extract_raw.py index 899f1cc1..db84047a 100644 --- a/datasets/agenttuning_webshop/extract_raw.py +++ b/datasets/agenttuning_webshop/extract_raw.py @@ -16,12 +16,6 @@ "assistant": "assistant", "system": "system", } -SAMPLE_IDS = { - "alfworld": ["alfworld_155", "alfworld_219", "alfworld_56", "alfworld_58", "alfworld_149"], - "db": ["db_100", "db_356", "db_493", "db_516", "db_531"], - "kg": ["kg_132", "kg_260", "kg_164", "kg_221", "kg_12"], - "webshop": ["webshop_330", "webshop_190", "webshop_243", "webshop_89", "webshop_26"], -} def json_safe(value: Any) -> Any: @@ -50,18 +44,11 @@ def has_valid_conversation(sample: dict[str, Any]) -> bool: def main(config_name: str) -> None: try: - dataset = load_dataset("THUDM/AgentInstruct")[config_name] - sample_ids = SAMPLE_IDS.get(config_name) - if sample_ids is not None: - selected = {} - sample_id_set = set(sample_ids) - for sample in dataset: - if sample.get("id") in sample_id_set and has_valid_conversation(sample): - selected[sample["id"]] = json_safe(sample) - for sample_id in sample_ids: - if sample_id in selected: - print(json.dumps(selected[sample_id], ensure_ascii=False)) - return + dataset = load_dataset( + "parquet", + data_files=f"hf://datasets/THUDM/AgentInstruct/data/{config_name}-*.parquet", + split="train", + ) for sample in dataset: if not has_valid_conversation(sample): continue diff --git a/datasets/agenttuning_webshop/metadata.json b/datasets/agenttuning_webshop/metadata.json index 76c2ba3a..1ee4f344 100644 --- a/datasets/agenttuning_webshop/metadata.json +++ b/datasets/agenttuning_webshop/metadata.json @@ -42,7 +42,7 @@ } ], "code_enabled": [], - "browser_enabled": true, + "browser_enabled": false, "sample_expectations": { "min_std_steps": 25, "min_std_tool_calls": 12, diff --git a/datasets/android_in_the_wild/extract_raw.py b/datasets/android_in_the_wild/extract_raw.py index 34b5186f..a7cb267c 100644 --- a/datasets/android_in_the_wild/extract_raw.py +++ b/datasets/android_in_the_wild/extract_raw.py @@ -59,10 +59,7 @@ def parse_image_data(image_bytes, height, width, nb_channels) -> Image.Image: dataset = tf.data.TFRecordDataset(file_names, compression_type="GZIP") json_list: List[Dict[str, Any]] = [] - for i, rcd in enumerate(dataset): - if i >= 5: - break - + for rcd in dataset: example = tf.train.Example() example.ParseFromString(rcd.numpy()) diff --git a/datasets/androidcontrol/extract_raw.py b/datasets/androidcontrol/extract_raw.py index 4ecacf90..5cb00b03 100644 --- a/datasets/androidcontrol/extract_raw.py +++ b/datasets/androidcontrol/extract_raw.py @@ -103,10 +103,22 @@ def process_tfrecord_file(tfrecord_file): # Ensure the output directory exists if not os.path.exists(output_dir): os.makedirs(output_dir) -# Get the list of TFRecord files -tfrecord_files = [ - os.path.join(data_dir, f) for f in os.listdir(data_dir) if not f.endswith(".json") +# Get the list of TFRecord files. The bundled tf_record_sample is only for +# sample generation; extract_raw.py should process the full corpus by default. +local_tfrecord_files = [ + os.path.join(data_dir, f) + for f in os.listdir(data_dir) + if not f.endswith(".json") and f != "tf_record_sample" ] +tfrecord_files = local_tfrecord_files or tf.io.gfile.glob( + "gs://gresearch/android_control/android_control*" +) +if not tfrecord_files: + raise FileNotFoundError( + "No full AndroidControl TFRecord files found. Download " + "gs://gresearch/android_control/android_control* into " + f"{data_dir}, or make sure TensorFlow can read that public GCS path." + ) # Create a list to store parsed data data = [] # Use ThreadPoolExecutor to process files in parallel diff --git a/scripts/atif_to_std_common.py b/scripts/atif_to_std_common.py index 19d4cafb..55c9d035 100644 --- a/scripts/atif_to_std_common.py +++ b/scripts/atif_to_std_common.py @@ -108,6 +108,19 @@ def normalize_prompt_boilerplate(trajectory: ATIFTrajectory) -> None: step.source = "system" continue step.message = RESPONSE_FORMAT_PROMPT.sub("", step.message).strip() + if ( + len(trajectory.steps) >= 3 + and trajectory.steps[0].source == "system" + and trajectory.steps[1].source == "agent" + and trajectory.steps[2].source == "user" + and isinstance(trajectory.steps[1].message, str) + and trajectory.steps[1].message.lower().strip().startswith(("ok.", "okay", "sure")) + and not trajectory.steps[1].tool_calls + and trajectory.steps[1].observation is None + ): + del trajectory.steps[1] + for index, step in enumerate(trajectory.steps, start=1): + step.step_id = index def renumber_steps(steps: list[Step]) -> list[Step]: