初始化项目,由ModelHub XC社区提供模型
Model: karthik-2905/AL1-model-B Source: Original Platform
This commit is contained in:
37
mlb/gen_sft_chat.py
Normal file
37
mlb/gen_sft_chat.py
Normal file
@@ -0,0 +1,37 @@
|
||||
import os, json, random
|
||||
HERE = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
os.environ["HF_HOME"] = os.path.join(HERE, "hf_cache")
|
||||
|
||||
from datasets import load_dataset
|
||||
|
||||
OUT = os.path.join(HERE, "data", "sft_chat.jsonl")
|
||||
N = 500
|
||||
random.seed(0)
|
||||
|
||||
def main():
|
||||
ds = load_dataset("databricks/databricks-dolly-15k", split="train")
|
||||
idx = list(range(len(ds)))
|
||||
random.shuffle(idx)
|
||||
|
||||
rows = []
|
||||
for i in idx:
|
||||
r = ds[i]
|
||||
instr, ctx, resp = r["instruction"].strip(), r["context"].strip(), r["response"].strip()
|
||||
if not instr or not resp:
|
||||
continue
|
||||
if len(resp) > 1200: # keep seq length sane
|
||||
continue
|
||||
user = f"{instr}\n\n{ctx}" if ctx else instr
|
||||
rows.append({"messages": [{"role": "user", "content": user},
|
||||
{"role": "assistant", "content": resp}]})
|
||||
if len(rows) >= N:
|
||||
break
|
||||
|
||||
with open(OUT, "w") as f:
|
||||
for r in rows:
|
||||
f.write(json.dumps(r) + "\n")
|
||||
print(f"wrote {len(rows)} chat examples -> {OUT}")
|
||||
print("sample user:", rows[0]["messages"][0]["content"][:80])
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user