Files
AL1-model-B/mlb/build_sft.py

47 lines
1.5 KiB
Python
Raw Permalink Normal View History

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()