mistral-7b-tokenizer-ja / create_tokenizer.py
jondurbin's picture
Update create_tokenizer.py
75a21c6
Raw
History Blame Contribute Delete
1.33 kB
import gc
import os
import glob
import random
from huggingface_hub import snapshot_download
from datasets import load_dataset, concatenate_datasets, Dataset
from transformers import AutoTokenizer
os.environ["HUGGING_FACE_HUB_TOKEN"] = "[redacted]"
# Download a random sampling of the JA data (159 total files, select those with 3 as last digit).
snapshot_download(
repo_id="uonlp/CulturaX",
local_dir="/workspace/CulturaX",
cache_dir="/workspace/.cache",
allow_patterns=["ja/*3.parquet"],
repo_type="dataset"
)
# Combine the files into one dataset (using a random subset of 2 datasets). More would be better, but the RAM usage is insane...
paths = list(map(str, glob.glob("CulturaX/ja/*.parquet")))
random.shuffle(paths)
dataset = concatenate_datasets([Dataset.from_parquet(path) for path in paths[0:2]])
gc.collect()
# Yield dataset in batches.
batch_size = 1000
def batch_iterator():
for i in range(0, len(dataset), batch_size):
yield dataset[i : i + batch_size]["text"]
gc.collect()
# Train a new mistral tokenizer from the dataset.
tokenizer = AutoTokenizer.from_pretrained("mistralai/Mistral-7B-v0.1")
new_tokenizer = tokenizer.train_new_from_iterator(batch_iterator(), vocab_size=65536, max_token_length=8)
new_tokenizer.save_pretrained("/workspace/mistral-7b-tokenizer-ja")