[bugfix] fix read data containing "tools" field (#255)

This commit is contained in:
Maybewuss
2025-09-03 22:34:43 +08:00
committed by GitHub
parent 7c0a0341d5
commit 5c4fce60fc
+11 -7
View File
@@ -1,6 +1,6 @@
import random
from datasets import Dataset as hf_ds
from pandas as pd
from slime.utils.types import Sample
@@ -9,15 +9,14 @@ __all__ = ["Dataset"]
# TODO: don't read the whole file into memory.
def read_file(path):
if path.endswith(".jsonl") or path.endswith(".json"):
ds = hf_ds.from_json(path)
if path.endswith(".jsonl"):
df = pd.read_json(path, lines=True)
elif path.endswith(".parquet"):
ds = hf_ds.from_parquet(path)
df = pd.read_parquet(path, dtype_backend="pyarrow")
else:
raise ValueError(f"Unsupported file format: {path}. Supported formats are .jsonl and .parquet.")
for data in ds:
yield data
for _, row in df.iterrows():
yield row.to_dict()
class Dataset:
@@ -40,6 +39,11 @@ class Dataset:
if apply_chat_template:
if tool_key is not None:
tools = data[tool_key]
if isinstance(tools, str):
tools = json.loads(tools)
elif isinstance(tools, np.ndarray):
tools = tools.tolist()
assert isinstance(tools, list), f"tools must be a list, got {type(tools)} instead"
else:
tools = None
prompt = tokenizer.apply_chat_template(prompt, tools, tokenize=False, add_generation_prompt=True)