37 lines
1.2 KiB
Python
37 lines
1.2 KiB
Python
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() |