ares-static-lab / ares_core /tokenizer_train.py
jacmor64's picture
Deploy Ares Static Lab Colab training pipeline
8fa3dd6 verified
Raw
History Blame Contribute Delete
2.86 kB
from __future__ import annotations
import argparse
from pathlib import Path
from typing import Iterable, List
def iter_files(paths: List[str]) -> Iterable[str]:
for item in paths:
p = Path(item)
if p.is_dir():
for child in sorted(p.rglob("*")):
if child.is_file() and child.suffix.lower() in {".txt", ".md", ".jsonl", ".json"}:
yield str(child)
elif p.is_file():
yield str(p)
else:
raise FileNotFoundError(item)
def train_bpe(input_paths: List[str], output: str, vocab_size: int, min_frequency: int = 2) -> None:
try:
from tokenizers import Tokenizer
from tokenizers.models import BPE
from tokenizers.pre_tokenizers import ByteLevel
from tokenizers.decoders import ByteLevel as ByteLevelDecoder
from tokenizers.trainers import BpeTrainer
from tokenizers.processors import TemplateProcessing
except ImportError as exc:
raise SystemExit("Install tokenizers first: pip install tokenizers") from exc
special_tokens = [
"<|pad|>",
"<|bos|>",
"<|eos|>",
"<|unk|>",
"<|system|>",
"<|user|>",
"<|assistant|>",
"<|tool|>",
"<|end|>",
]
files = list(iter_files(input_paths))
if not files:
raise ValueError("No input files found")
tokenizer = Tokenizer(BPE(unk_token="<|unk|>"))
tokenizer.pre_tokenizer = ByteLevel(add_prefix_space=False)
tokenizer.decoder = ByteLevelDecoder()
trainer = BpeTrainer(
vocab_size=vocab_size,
min_frequency=min_frequency,
show_progress=True,
special_tokens=special_tokens,
)
tokenizer.train(files, trainer)
bos_id = tokenizer.token_to_id("<|bos|>")
eos_id = tokenizer.token_to_id("<|eos|>")
tokenizer.post_processor = TemplateProcessing(
single="<|bos|> $A <|eos|>",
pair="<|bos|> $A <|end|> $B <|eos|>",
special_tokens=[("<|bos|>", bos_id), ("<|eos|>", eos_id), ("<|end|>", tokenizer.token_to_id("<|end|>"))],
)
out = Path(output)
out.parent.mkdir(parents=True, exist_ok=True)
tokenizer.save(str(out))
print(f"Saved tokenizer with vocab_size={tokenizer.get_vocab_size()} to {out}")
def main() -> None:
parser = argparse.ArgumentParser(description="Train Ares BPE tokenizer from scratch.")
parser.add_argument("--input", nargs="+", required=True, help="Input files or directories")
parser.add_argument("--output", required=True, help="Output tokenizer.json path")
parser.add_argument("--vocab-size", type=int, default=32000)
parser.add_argument("--min-frequency", type=int, default=2)
args = parser.parse_args()
train_bpe(args.input, args.output, args.vocab_size, args.min_frequency)
if __name__ == "__main__":
main()