Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
23 changes: 5 additions & 18 deletions datasets/agenttuning_alfworld/extract_raw.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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
Expand Down
23 changes: 5 additions & 18 deletions datasets/agenttuning_db/extract_raw.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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
Expand Down
4 changes: 3 additions & 1 deletion datasets/agenttuning_db/metadata.json
Original file line number Diff line number Diff line change
@@ -1,5 +1,7 @@
{
"custom_tools": [],
"code_enabled": [],
"code_enabled": [
"bash"
],
"browser_enabled": false
}
23 changes: 5 additions & 18 deletions datasets/agenttuning_kg/extract_raw.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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
Expand Down
23 changes: 5 additions & 18 deletions datasets/agenttuning_mind2web/extract_raw.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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
Expand Down
23 changes: 5 additions & 18 deletions datasets/agenttuning_os/extract_raw.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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
Expand Down
23 changes: 5 additions & 18 deletions datasets/agenttuning_webshop/extract_raw.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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
Expand Down
2 changes: 1 addition & 1 deletion datasets/agenttuning_webshop/metadata.json
Original file line number Diff line number Diff line change
Expand Up @@ -42,7 +42,7 @@
}
],
"code_enabled": [],
"browser_enabled": true,
"browser_enabled": false,
"sample_expectations": {
"min_std_steps": 25,
"min_std_tool_calls": 12,
Expand Down
5 changes: 1 addition & 4 deletions datasets/android_in_the_wild/extract_raw.py
Original file line number Diff line number Diff line change
Expand Up @@ -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())

Expand Down
18 changes: 15 additions & 3 deletions datasets/androidcontrol/extract_raw.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
13 changes: 13 additions & 0 deletions scripts/atif_to_std_common.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]:
Expand Down
Loading