初始化项目,由ModelHub XC社区提供模型
Model: karthik-2905/AL1-model-B Source: Original Platform
This commit is contained in:
90
mlb/gen_sft_tools.py
Normal file
90
mlb/gen_sft_tools.py
Normal file
@@ -0,0 +1,90 @@
|
||||
import os, json, random
|
||||
HERE = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
import sys; sys.path.insert(0, os.path.join(HERE, "mlb"))
|
||||
from tools_common import TOOLS
|
||||
|
||||
OUT = os.path.join(HERE, "data", "sft_tools.jsonl")
|
||||
N = 1000
|
||||
random.seed(0)
|
||||
|
||||
FILES = ["config.py", "app.py", "main.py", "README.md", "utils.py", "src/utils.py",
|
||||
"tests/test_api.py", "server.js", "index.html", "data/train.csv", "Makefile",
|
||||
"requirements.txt", "model.py", "train.py", ".env.example", "docker-compose.yml"]
|
||||
DIRS = ["src", "tests", "data", ".", "configs", "scripts", "app", "lib", "models"]
|
||||
CMDS = ["pytest", "git status", "ls -la", "npm test", "python train.py",
|
||||
"pip install -r requirements.txt", "git log --oneline", "make build",
|
||||
"docker build .", "ruff check .", "python -m pytest tests/", "git diff"]
|
||||
EDITS = [("5000", "8080"), ("foo", "bar"), ("debug=False", "debug=True"),
|
||||
("localhost", "0.0.0.0"), ("v1", "v2"), ("GET", "POST"),
|
||||
("255", "128"), ("lr=1e-3", "lr=3e-4"), ("True", "False")]
|
||||
WRITES = [("hello.py", "print('hi')"), ("notes.txt", "hello world"),
|
||||
(".gitignore", "__pycache__/"), ("VERSION", "1.0.0"),
|
||||
("run.sh", "#!/bin/bash\npython main.py"), ("todo.md", "- ship it")]
|
||||
|
||||
READ_T = ["Show me the contents of {p}", "Read {p}", "Open {p}", "What's in {p}?",
|
||||
"Display {p}", "Print out {p} for me", "Can you show me {p}?"]
|
||||
LIST_T = ["List the files in {p}", "What's in the {p} directory?", "Show the entries in {p}",
|
||||
"List everything under {p}", "ls {p}", "Show me what's inside {p}"]
|
||||
WRITE_T = ["Create a file {p} containing {c}", "Make a new file {p} with the text {c}",
|
||||
"Write {c} to {p}", "Save {c} into a file called {p}"]
|
||||
EDIT_T = ["In {p} replace {a} with {b}", "Change {a} to {b} in {p}",
|
||||
"Update {p}: swap {a} for {b}", "In the file {p}, replace {a} with {b}",
|
||||
"Edit {p} and change {a} into {b}", "Modify {p} so {a} becomes {b}"]
|
||||
BASH_T = ["Run {c}", "Execute {c}", "Run the command {c}", "Please run {c}",
|
||||
"Can you run {c}?", "Kick off {c}", "Go ahead and run {c}"]
|
||||
|
||||
WEIGHTS = [("edit_file", 0.28), ("run_bash", 0.28), ("read_file", 0.16),
|
||||
("list_dir", 0.14), ("write_file", 0.14)]
|
||||
|
||||
def pick_tool():
|
||||
r, acc = random.random(), 0.0
|
||||
for name, w in WEIGHTS:
|
||||
acc += w
|
||||
if r <= acc:
|
||||
return name
|
||||
return WEIGHTS[-1][0]
|
||||
|
||||
def make_example():
|
||||
t = pick_tool()
|
||||
if t == "read_file":
|
||||
p = random.choice(FILES)
|
||||
q = random.choice(READ_T).format(p=p); args = {"path": p}
|
||||
elif t == "list_dir":
|
||||
p = random.choice(DIRS)
|
||||
q = random.choice(LIST_T).format(p=p); args = {"path": p}
|
||||
elif t == "write_file":
|
||||
p, c = random.choice(WRITES)
|
||||
q = random.choice(WRITE_T).format(p=p, c=c); args = {"path": p, "content": c}
|
||||
elif t == "edit_file":
|
||||
p = random.choice(FILES); a, b = random.choice(EDITS)
|
||||
q = random.choice(EDIT_T).format(p=p, a=a, b=b)
|
||||
args = {"path": p, "old_string": a, "new_string": b}
|
||||
else: # run_bash
|
||||
c = random.choice(CMDS)
|
||||
q = random.choice(BASH_T).format(c=c); args = {"command": c}
|
||||
call = json.dumps({"name": t, "arguments": args})
|
||||
assistant = f"<tool_call>\n{call}\n</tool_call>"
|
||||
return {"messages": [{"role": "user", "content": q},
|
||||
{"role": "assistant", "content": assistant}],
|
||||
"tools": TOOLS}
|
||||
|
||||
def main():
|
||||
seen, rows = set(), []
|
||||
while len(rows) < N:
|
||||
ex = make_example()
|
||||
key = ex["messages"][0]["content"] + ex["messages"][1]["content"]
|
||||
if key in seen:
|
||||
continue
|
||||
seen.add(key); rows.append(ex)
|
||||
with open(OUT, "w") as f:
|
||||
for r in rows:
|
||||
f.write(json.dumps(r) + "\n")
|
||||
dist = {}
|
||||
for r in rows:
|
||||
name = json.loads(r["messages"][1]["content"].split("<tool_call>\n")[1].split("\n</tool_call>")[0])["name"]
|
||||
dist[name] = dist.get(name, 0) + 1
|
||||
print(f"wrote {len(rows)} -> {OUT}")
|
||||
print("distribution:", dist)
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user