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