47 lines
1.5 KiB
Python
47 lines
1.5 KiB
Python
import os, json, random
|
|
HERE = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
|
|
|
TOOLS_F = os.path.join(HERE, "data", "sft_tools.jsonl")
|
|
CHAT_F = os.path.join(HERE, "data", "sft_chat.jsonl")
|
|
EVAL_TOOLS = os.path.join(HERE, "eval", "tools_eval.jsonl")
|
|
EVAL_CHAT = os.path.join(HERE, "eval", "chat_eval.jsonl")
|
|
TRAIN_OUT = os.path.join(HERE, "data", "sft_train.jsonl")
|
|
VAL_OUT = os.path.join(HERE, "data", "sft_val.jsonl")
|
|
VAL_FRAC = 0.05
|
|
random.seed(0)
|
|
|
|
def load(path):
|
|
return [json.loads(l) for l in open(path) if l.strip()]
|
|
|
|
def main():
|
|
eval_qs = set()
|
|
for f in (EVAL_TOOLS, EVAL_CHAT):
|
|
for c in load(f):
|
|
eval_qs.add(c["query"].strip())
|
|
|
|
rows = load(TOOLS_F) + load(CHAT_F)
|
|
kept, dropped = [], 0
|
|
for r in rows:
|
|
if r["messages"][0]["content"].strip() in eval_qs:
|
|
dropped += 1
|
|
continue
|
|
kept.append(r)
|
|
|
|
random.shuffle(kept)
|
|
n_val = int(len(kept) * VAL_FRAC)
|
|
val, train = kept[:n_val], kept[n_val:]
|
|
|
|
for path, split in ((TRAIN_OUT, train), (VAL_OUT, val)):
|
|
with open(path, "w") as f:
|
|
for r in split:
|
|
f.write(json.dumps(r) + "\n")
|
|
|
|
def has_tools(r): return "tools" in r
|
|
print(f"leakage-dropped: {dropped}")
|
|
print(f"train: {len(train)} (tool={sum(has_tools(r) for r in train)}, "
|
|
f"chat={sum(not has_tools(r) for r in train)})")
|
|
print(f"val: {len(val)} (tool={sum(has_tools(r) for r in val)}, "
|
|
f"chat={sum(not has_tools(r) for r in val)})")
|
|
|
|
if __name__ == "__main__":
|
|
main() |