cmpatino/tokenization_diff
0
1import html2 3import gradio as gr4from datasets import load_dataset5from transformers import AutoTokenizer6 7 8def build_alignment_groups_from_ids(student_tokenizer, teacher_tokenizer, student_token_ids, teacher_token_ids):9 """10 Build alignment groups using a greedy substring-equality algorithm on decoded token pieces.11 Adapted from TRL's GoldTrainer._build_alignment_groups_from_ids.12 """13 14 def to_canonical_pieces(tok, ids):15 pieces = []16 prev = ""17 for k in range(len(ids)):18 cur = tok.decode(ids[: k + 1], skip_special_tokens=False, clean_up_tokenization_spaces=False)19 pieces.append(cur[len(prev):])20 prev = cur21 return pieces22 23 s_pieces = to_canonical_pieces(student_tokenizer, student_token_ids)24 t_pieces = to_canonical_pieces(teacher_tokenizer, teacher_token_ids)25 26 i = j = 027 s_buf = t_buf = ""28 s_group = []29 t_group = []30 s_groups = []31 t_groups = []32 33 def flush():34 if s_group and t_group:35 s_groups.append(s_group.copy())36 t_groups.append(t_group.copy())37 38 while i < len(s_pieces) or j < len(t_pieces):39 if s_buf == t_buf and s_buf != "":40 flush()41 s_buf = t_buf = ""42 s_group = []43 t_group = []44 continue45 46 if s_buf == "" and i < len(s_pieces):47 s_buf += s_pieces[i]48 s_group.append(i)49 i += 150 continue51 if t_buf == "" and j < len(t_pieces):52 t_buf += t_pieces[j]53 t_group.append(j)54 j += 155 continue56 57 if len(s_buf) <= len(t_buf):58 if i < len(s_pieces):59 s_buf += s_pieces[i]60 s_group.append(i)61 i += 162 elif j < len(t_pieces):63 t_buf += t_pieces[j]64 t_group.append(j)65 j += 166 else:67 if j < len(t_pieces):68 t_buf += t_pieces[j]69 t_group.append(j)70 j += 171 elif i < len(s_pieces):72 s_buf += s_pieces[i]73 s_group.append(i)74 i += 175 76 if s_buf == t_buf and s_group and t_group:77 flush()78 elif s_group or t_group:79 if not s_group:80 s_group = []81 if not t_group:82 t_group = []83 if s_group or t_group:84 s_groups.append(s_group.copy() if s_group else [])85 t_groups.append(t_group.copy() if t_group else [])86 87 return s_groups, t_groups88 89 90def _decode_pieces(tokenizer, token_ids, indices):91 """Decode individual token pieces for a group of token indices."""92 return [93 tokenizer.decode([token_ids[idx]], skip_special_tokens=False, clean_up_tokenization_spaces=False)94 for idx in indices95 ]96 97 98def _format_pieces(pieces):99 """Format token pieces as a list, e.g. '["hel", "lo"]'."""100 inner = ", ".join(f'"{p}"' for p in pieces)101 return f"[{inner}]"102 103 104def highlight_groups(student_tokenizer, teacher_tokenizer, student_token_ids, teacher_token_ids, s_groups, t_groups):105 """Build an HTML string with highlighted misalignment regions."""106 parts = []107 first_purple = True108 for k in range(len(s_groups)):109 s_ids = [student_token_ids[idx] for idx in s_groups[k]]110 text = student_tokenizer.decode(s_ids, skip_special_tokens=False, clean_up_tokenization_spaces=False)111 escaped = html.escape(text)112 113 s_multi = len(s_groups[k]) > 1114 t_multi = len(t_groups[k]) > 1115 116 if s_multi and t_multi:117 if first_purple:118 s_pieces = _decode_pieces(student_tokenizer, student_token_ids, s_groups[k])119 t_pieces = _decode_pieces(teacher_tokenizer, teacher_token_ids, t_groups[k])120 tooltip = html.escape(f'Student: {_format_pieces(s_pieces)} / Teacher: {_format_pieces(t_pieces)}')121 parts.append(f'<span style="background-color: #b388ff;" title="{tooltip}">{escaped}</span>')122 first_purple = False123 else:124 parts.append(f'<span style="background-color: #b388ff;">{escaped}</span>')125 elif s_multi:126 s_pieces = _decode_pieces(student_tokenizer, student_token_ids, s_groups[k])127 tooltip = html.escape(f'Student: {_format_pieces(s_pieces)}')128 parts.append(f'<span style="background-color: #ffcc80;" title="{tooltip}">{escaped}</span>')129 elif t_multi:130 t_pieces = _decode_pieces(teacher_tokenizer, teacher_token_ids, t_groups[k])131 tooltip = html.escape(f'Teacher: {_format_pieces(t_pieces)}')132 parts.append(f'<span style="background-color: #90caf9;" title="{tooltip}">{escaped}</span>')133 else:134 parts.append(escaped)135 136 return "".join(parts)137 138 139def make_html_block(student_tokenizer, teacher_tokenizer, text, idx):140 """Process a single text and return its highlighted HTML block."""141 s_ids = student_tokenizer.encode(text, add_special_tokens=False)142 t_ids = teacher_tokenizer.encode(text, add_special_tokens=False)143 144 s_groups, t_groups = build_alignment_groups_from_ids(145 student_tokenizer, teacher_tokenizer, s_ids, t_ids146 )147 148 highlighted = highlight_groups(student_tokenizer, teacher_tokenizer, s_ids, t_ids, s_groups, t_groups)149 150 # Build tokenized views with alternating colors151 s_tokens = [student_tokenizer.decode([tid], skip_special_tokens=False, clean_up_tokenization_spaces=False) for tid in s_ids]152 t_tokens = [teacher_tokenizer.decode([tid], skip_special_tokens=False, clean_up_tokenization_spaces=False) for tid in t_ids]153 154 color1 = "#fff9c4"155 color2 = "#b2ebf2"156 157 s_tokens_html = "".join(158 f'<span style="background-color:{color1 if i % 2 == 0 else color2};">{html.escape(t)}</span>'159 for i, t in enumerate(s_tokens)160 )161 t_tokens_html = "".join(162 f'<span style="background-color:{color1 if i % 2 == 0 else color2};">{html.escape(t)}</span>'163 for i, t in enumerate(t_tokens)164 )165 166 tokenized_section = f'''167 <div style="margin-bottom:15px;">168 <details style="margin-bottom:10px;">169 <summary style="cursor:pointer; font-weight:bold; user-select:none;">Show tokenization details</summary>170 <div style="display:grid; grid-template-columns:1fr 1fr; gap:15px; margin-top:10px;">171 <div style="border:1px solid #ddd; padding:10px; border-radius:5px;">172 <strong style="color:#f57c00;">Student Tokens ({len(s_ids)})</strong>173 <div style="margin-top:8px; font-size:12px; word-break:break-word;">{s_tokens_html}</div>174 </div>175 <div style="border:1px solid #ddd; padding:10px; border-radius:5px;">176 <strong style="color:#1976d2;">Teacher Tokens ({len(t_ids)})</strong>177 <div style="margin-top:8px; font-size:12px; word-break:break-word;">{t_tokens_html}</div>178 </div>179 </div>180 </details>181 </div>182 '''183 184 return (185 f'<div style="border:1px solid #ccc; padding:10px; margin:10px 0; '186 f'border-radius:5px; white-space:pre-wrap; font-family:monospace; font-size:13px;">'187 f"<strong>Text {idx + 1}</strong> "188 f"(student tokens: {len(s_ids)}, teacher tokens: {len(t_ids)})<br><br>"189 f"{tokenized_section}"190 f"{highlighted}"191 f"</div>"192 )193 194 195def process_texts(student_model_id, teacher_model_id, dataset_id, dataset_config, progress=gr.Progress()):196 """Load tokenizers and dataset, compute first row only."""197 progress(0, desc="Loading tokenizers...")198 student_tokenizer = AutoTokenizer.from_pretrained(student_model_id)199 teacher_tokenizer = AutoTokenizer.from_pretrained(teacher_model_id)200 201 progress(0.5, desc="Loading dataset...")202 config = dataset_config.strip() if dataset_config and dataset_config.strip() else None203 ds = load_dataset(dataset_id, name=config, split="train")204 rows = ds.select(range(min(10, len(ds))))205 texts = ["".join(msg["content"] for msg in row["messages"]) for row in rows]206 207 progress(0.8, desc="Processing first text...")208 first_block = make_html_block(student_tokenizer, teacher_tokenizer, texts[0], 0)209 cache = {0: first_block}210 211 progress(1, desc="Done!")212 return student_tokenizer, teacher_tokenizer, texts, cache, 0, render_page(cache, 0, len(texts))213 214 215LEGEND = (216 '<div style="margin-bottom:15px; font-family:sans-serif;">'217 "<strong>Legend:</strong> "218 '<span style="background-color:#ffcc80; padding:2px 8px; margin-right:8px;">Student token split (orange)</span>'219 '<span style="background-color:#90caf9; padding:2px 8px; margin-right:8px;">Teacher token split (blue)</span>'220 '<span style="background-color:#b388ff; padding:2px 8px;">Both (purple)</span>'221 "</div>"222)223 224 225def render_page(cache, idx, total):226 if not cache:227 return ""228 counter = f'<div style="font-family:sans-serif; margin-bottom:10px;">Text {idx + 1} of {total}</div>'229 return LEGEND + counter + cache[idx]230 231 232def go_prev(cache, idx, texts):233 idx = max(0, idx - 1)234 return cache, idx, render_page(cache, idx, len(texts))235 236 237def go_next(student_tokenizer, teacher_tokenizer, texts, cache, idx):238 idx = min(len(texts) - 1, idx + 1)239 if idx not in cache:240 cache[idx] = make_html_block(student_tokenizer, teacher_tokenizer, texts[idx], idx)241 return cache, idx, render_page(cache, idx, len(texts))242 243 244with gr.Blocks(title="Tokenization Diff") as demo:245 gr.Markdown("# Tokenization Diff\nVisualize where two tokenizers differ in how they tokenize text.")246 247 with gr.Row():248 student_model = gr.Textbox(label="Student Model", value="Qwen/Qwen3-8B")249 teacher_model = gr.Textbox(label="Teacher Model", value="deepseek-ai/DeepSeek-Math-V2")250 dataset_id = gr.Textbox(label="Dataset ID", value="lm-provers/FineProofs-SFT")251 dataset_config = gr.Textbox(label="Dataset Config", value="default")252 253 submit_btn = gr.Button("Submit", variant="primary")254 255 student_tok_state = gr.State(None)256 teacher_tok_state = gr.State(None)257 texts_state = gr.State([])258 cache_state = gr.State({})259 idx_state = gr.State(0)260 261 output = gr.HTML(label="Tokenization Diff Output")262 263 with gr.Row():264 prev_btn = gr.Button("Previous")265 next_btn = gr.Button("Next")266 267 submit_btn.click(268 fn=process_texts,269 inputs=[student_model, teacher_model, dataset_id, dataset_config],270 outputs=[student_tok_state, teacher_tok_state, texts_state, cache_state, idx_state, output],271 )272 prev_btn.click(273 fn=go_prev,274 inputs=[cache_state, idx_state, texts_state],275 outputs=[cache_state, idx_state, output],276 )277 next_btn.click(278 fn=go_next,279 inputs=[student_tok_state, teacher_tok_state, texts_state, cache_state, idx_state],280 outputs=[cache_state, idx_state, output],281 )282 283if __name__ == "__main__":284 demo.launch()285 