初始化项目,由ModelHub XC社区提供模型
Model: karthik-2905/AL1-model-B Source: Original Platform
This commit is contained in:
47
mlb/build_sft.py
Normal file
47
mlb/build_sft.py
Normal file
@@ -0,0 +1,47 @@
|
||||
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()
|
||||
Reference in New Issue
Block a user