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