Ares_v1 / src /ares /tokenizer /tokenizer.py
Ares Deployer
Deploy Ares full from scratch: BPE 128K, RoPE 8192, GQA+KV, RMSNorm, SwiGLU, RAG SQLite, CoT/ToT/Planner, SFT/RLHF, code+search
701cf7d
Raw
History Blame Contribute Delete
8.92 kB
"""
Final Tokenizer wrapper - encode/decode with BPE merges, handles special tokens, sections.
Implements token sections for prompt formatting.
"""
import json
import re
from typing import List, Dict
import os
class AresTokenizer:
def __init__(self, vocab_file: str = None, vocab_size: int = 128256):
self.vocab_file = vocab_file
self.vocab = {}
self.inverse_vocab = {}
self.merges = {} # pair -> rank (lower = earlier)
self.merge_ranks = {}
self.special_tokens = ["<|pad|>", "<|bos|>", "<|eos|>", "<|unk|>", "<|im_start|>", "<|im_end|>"]
self.special_ids = {}
# byte encoder mapping from trainer
self.byte_encoder = {}
self.byte_decoder = {}
# GPT-2 style byte mapping fallback
if not self.byte_encoder:
bs = list(range(ord("!"), ord("~")+1)) + list(range(ord("¡"), ord("¬")+1)) + list(range(ord("®"), ord("ÿ")+1))
cs = bs[:]
n=0
for b in range(256):
if b not in bs:
bs.append(b)
cs.append(256+n)
n+=1
self.byte_encoder = dict(zip(bs, [chr(c) for c in cs]))
self.byte_decoder = {v:k for k,v in self.byte_encoder.items()}
self.pat = re.compile(r"""'s|'t|'re|'ve|'m|'ll|'d| ?[A-Za-z]+| ?[0-9]+| ?[^\sA-Za-z0-9]+|\s+(?!\S)|\s+""")
if vocab_file and os.path.exists(vocab_file):
self.load(vocab_file)
else:
# Dummy small vocab for tiny config test
self._build_dummy_vocab(vocab_size)
# Special token ids
for tok in self.special_tokens:
if tok in self.vocab:
self.special_ids[tok] = self.vocab[tok]
# Defaults
self.pad_token_id = self.special_ids.get("<|pad|>", 0)
self.bos_token_id = self.special_ids.get("<|bos|>", 1)
self.eos_token_id = self.special_ids.get("<|eos|>", 2)
self.unk_token_id = self.special_ids.get("<|unk|>", 3)
def _build_dummy_vocab(self, size):
# Build minimal byte-level vocab for tests without trained file, but pad to requested size
base = {tok:i for i,tok in enumerate(self.special_tokens)}
for b in range(256):
ch = self.byte_encoder[b]
if ch not in base:
base[ch] = len(base)
# Pad remaining with tok_#
idx = len(base)
while len(base) < size:
tok = f"tok_{idx}"
if tok not in base:
base[tok] = len(base)
idx+=1
if idx>size+10000:
break
self.vocab = base
self.inverse_vocab = {v:k for k,v in base.items()}
# No merges
self.merges = {}
self.merge_ranks = {}
def load(self, path):
with open(path, 'r', encoding='utf-8') as f:
data = json.load(f)
self.vocab = data["vocab"]
# ensure int ids
# merges stored as "a b": rank
raw_merges = data.get("merges", {})
self.merges = {}
for k,v in raw_merges.items():
parts = k.split(' ')
if len(parts)==2:
self.merges[(parts[0], parts[1])] = v
# sort by rank for BPE
self.merge_ranks = {pair: rank for pair, rank in sorted(self.merges.items(), key=lambda x: x[1])}
self.special_tokens = data.get("special_tokens", self.special_tokens)
self.byte_encoder = data.get("byte_encoder", self.byte_encoder)
self.byte_decoder = {v:k for k,v in self.byte_encoder.items()}
self.inverse_vocab = {v:k for k,v in self.vocab.items()}
def get_pairs(self, word):
pairs = set()
prev = word[0]
for ch in word[1:]:
pairs.add((prev, ch))
prev = ch
return pairs
def bpe(self, token):
# token is string of byte-encoded chars
if token in self.vocab:
return [token]
word = list(token)
pairs = self.get_pairs(word)
if not pairs:
return [token]
while True:
# Find best pair by lowest rank
bigram = min(pairs, key=lambda pair: self.merge_ranks.get(pair, float('inf')))
if bigram not in self.merge_ranks:
break
first, second = bigram
new_word = []
i=0
while i < len(word):
try:
j = word.index(first, i)
except:
new_word.extend(word[i:])
break
new_word.extend(word[i:j])
i=j
if i < len(word)-1 and word[i]==first and word[i+1]==second:
new_word.append(first+second)
i+=2
else:
new_word.append(word[i])
i+=1
word = new_word
if len(word)==1:
break
else:
pairs = self.get_pairs(word)
return word
def encode(self, text: str, add_bos=False, add_eos=False) -> List[int]:
tokens = []
if add_bos:
tokens.append(self.bos_token_id)
# handle special tokens as separate sections
# Split by special tokens
# Simple: iterate over pat
for tok in re.findall(self.pat, text):
# byte encode
try:
encoded = ''.join(self.byte_encoder[b] for b in tok.encode('utf-8'))
except:
encoded = tok
bpe_tokens = self.bpe(encoded)
for bt in bpe_tokens:
tokens.append(self.vocab.get(bt, self.unk_token_id))
if add_eos:
tokens.append(self.eos_token_id)
return tokens
def decode(self, ids: List[int], skip_special=True) -> str:
text = ''
for i in ids:
token = self.inverse_vocab.get(i, "<|unk|>")
if skip_special and token in self.special_tokens:
continue
text += token
# byte decode
# Convert unicode chars back to bytes
try:
# text is sequence of byte-encoded unicode chars
byte_vals = []
for c in text:
if c in self.byte_decoder:
byte_vals.append(self.byte_decoder[c])
else:
# might be merged token containing multiple bytes encoded chars
# Decode each char inside token
for ch in c:
if ch in self.byte_decoder:
byte_vals.append(self.byte_decoder[ch])
decoded = bytes(byte_vals).decode('utf-8', errors='replace')
return decoded
except Exception:
return text
def encode_with_sections(self, messages: List[Dict[str,str]]) -> List[int]:
"""
Token sections for chat: [{"role":"user","content":"..."}]
Format: <|im_start|>role\ncontent<|im_end|>
"""
ids = []
for m in messages:
role = m.get("role","user")
content = m.get("content","")
start = f"<|im_start|>{role}\n"
end = "<|im_end|>\n"
# Encode sections
for part in [start, content, end]:
if part in self.special_tokens or part.startswith("<|im_"):
# handle special split
# start contains special token substring
if "<|im_start|>" in part:
ids.append(self.vocab.get("<|im_start|>", self.special_ids.get("<|im_start|>", 4)))
# remainder after
remainder = part.replace("<|im_start|>","")
if remainder:
ids.extend(self.encode(remainder))
elif "<|im_end|>" in part:
# check if content includes im_end
if part.strip()=="<|im_end|>":
ids.append(self.vocab.get("<|im_end|>", self.special_ids.get("<|im_end|>",5)))
else:
# split
ids.extend(self.encode(part.replace("<|im_end|>","")))
ids.append(self.vocab.get("<|im_end|>",5))
else:
ids.append(self.vocab.get(part, self.unk_token_id))
else:
ids.extend(self.encode(part))
return ids
def save(self, path):
data = {
"vocab": self.vocab,
"merges": {f"{k[0]} {k[1]}": v for k,v in self.merges.items()},
"special_tokens": self.special_tokens,
"byte_encoder": self.byte_encoder
}
with open(path,'w',encoding='utf-8') as f:
json.dump(data,f,ensure_ascii=False)