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

37 lines
1.2 KiB
Python
Raw Normal View History

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