Spaces:
Running
Running
Download tokenizer.py from ArushBuilds/Pragya: direct link, hf CLI and curl.
- Browser
- Download file 2.86 kB
-
https://huggingface.co/spaces/ArushBuilds/Pragya/resolve/90c3f2a4d2d42dd3c2eb0c82a077b1bd656d3c35/tokenizer.py
- Command line
-
hf download hf://spaces/ArushBuilds/Pragya@90c3f2a4d2d42dd3c2eb0c82a077b1bd656d3c35/tokenizer.py
-
curl -L -o tokenizer.py https://huggingface.co/spaces/ArushBuilds/Pragya/resolve/90c3f2a4d2d42dd3c2eb0c82a077b1bd656d3c35/tokenizer.py
2.86 kB
| from __future__ import annotations | |
| import json | |
| from pathlib import Path | |
| from typing import Sequence, List | |
| import sentencepiece as spm | |
| def train_sentencepiece( | |
| data_files: Sequence[str], | |
| model_prefix: str = 'tokenizer', | |
| vocab_size: int = 32000, | |
| model_type: str = 'bpe', | |
| character_coverage: float = 1.0, | |
| byte_fallback: bool = True, | |
| pad_id: int = 1, | |
| unk_id: int = 0, | |
| bos_id: int = 2, | |
| eos_id: int = 3, | |
| ) -> str: | |
| data_files = [str(Path(p)) for p in data_files] | |
| if not data_files: | |
| raise ValueError('data_files is empty') | |
| args = [ | |
| f"--input={','.join(data_files)}", | |
| f'--model_prefix={model_prefix}', | |
| f'--vocab_size={int(vocab_size)}', | |
| f'--model_type={model_type}', | |
| f'--character_coverage={character_coverage}', | |
| f'--pad_id={pad_id}', | |
| f'--unk_id={unk_id}', | |
| f'--bos_id={bos_id}', | |
| f'--eos_id={eos_id}', | |
| '--hard_vocab_limit=false', | |
| '--normalization_rule_name=nmt_nfkc', | |
| ] | |
| if byte_fallback: | |
| args.append('--byte_fallback=true') | |
| spm.SentencePieceTrainer.train(' '.join(args)) | |
| return f'{model_prefix}.model' | |
| class TokenizerWrapper: | |
| def __init__(self, model_path: str): | |
| model_path = str(Path(model_path)) | |
| if not Path(model_path).exists(): | |
| raise FileNotFoundError(model_path) | |
| self.sp = spm.SentencePieceProcessor(model_file=model_path) | |
| self.vocab_size = int(self.sp.vocab_size()) | |
| self.pad_id = self.sp.pad_id() | |
| self.unk_id = self.sp.unk_id() | |
| self.bos_id = self.sp.bos_id() | |
| self.eos_id = self.sp.eos_id() | |
| for name, val in [('pad', self.pad_id), ('unk', self.unk_id), ('bos', self.bos_id), ('eos', self.eos_id)]: | |
| if val < 0: | |
| raise ValueError(f'SentencePiece model missing <{name}>') | |
| def encode(self, text: str, add_bos: bool = True, add_eos: bool = False) -> List[int]: | |
| ids = list(self.sp.encode(text, out_type=int)) | |
| if add_bos: | |
| ids = [self.bos_id] + ids | |
| if add_eos: | |
| ids = ids + [self.eos_id] | |
| return ids | |
| def encode_batch(self, texts: Sequence[str], add_bos: bool = True, add_eos: bool = False) -> List[List[int]]: | |
| return [self.encode(t, add_bos=add_bos, add_eos=add_eos) for t in texts] | |
| def decode(self, ids: Sequence[int]) -> str: | |
| return self.sp.decode([int(i) for i in ids if int(i) != self.pad_id]) | |
| def save_config(self, path: str) -> None: | |
| Path(path).write_text(json.dumps({ | |
| 'vocab_size': self.vocab_size, | |
| 'pad_id': self.pad_id, | |
| 'unk_id': self.unk_id, | |
| 'bos_id': self.bos_id, | |
| 'eos_id': self.eos_id, | |
| }, indent=2), encoding='utf-8') | |