38 lines
1.3 KiB
Python
38 lines
1.3 KiB
Python
|
|
import os, json
|
||
|
|
HERE = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||
|
|
os.environ["HF_HOME"] = os.path.join(HERE, "hf_cache")
|
||
|
|
import sys; sys.path.insert(0, os.path.join(HERE, "mlb"))
|
||
|
|
|
||
|
|
from transformers import AutoTokenizer
|
||
|
|
|
||
|
|
MODEL = "Qwen/Qwen3-0.6B"
|
||
|
|
DATA = os.path.join(HERE, "data", "sft_tools.jsonl")
|
||
|
|
|
||
|
|
def main():
|
||
|
|
tok = AutoTokenizer.from_pretrained(MODEL)
|
||
|
|
rows = [json.loads(l) for l in open(DATA) if l.strip()]
|
||
|
|
print(f"loaded {len(rows)} rows")
|
||
|
|
|
||
|
|
lengths, bad = [], 0
|
||
|
|
for r in rows:
|
||
|
|
text = tok.apply_chat_template(
|
||
|
|
r["messages"], tools=r["tools"], tokenize=False,
|
||
|
|
add_generation_prompt=False, enable_thinking=False)
|
||
|
|
if "<tool_call>" not in text:
|
||
|
|
bad += 1
|
||
|
|
lengths.append(len(tok(text).input_ids))
|
||
|
|
|
||
|
|
lengths.sort()
|
||
|
|
print("tool_call missing in rendered text:", bad)
|
||
|
|
print("token length min/mean/max:",
|
||
|
|
lengths[0], round(sum(lengths) / len(lengths), 1), lengths[-1])
|
||
|
|
print("p95 length:", lengths[int(0.95 * len(lengths))])
|
||
|
|
|
||
|
|
print("\n----- SAMPLE RENDERED TRAINING STRING -----\n")
|
||
|
|
s = tok.apply_chat_template(rows[0]["messages"], tools=rows[0]["tools"],
|
||
|
|
tokenize=False, add_generation_prompt=False,
|
||
|
|
enable_thinking=False)
|
||
|
|
print(s)
|
||
|
|
|
||
|
|
if __name__ == "__main__":
|
||
|
|
main()
|