Team Ai
Apppublic

cmpatino/tokenization_diff

sourceHugging Faceupdated 8mo agoView on Hugging Face
0likes
app.py285 linesDownload Raw Back to root
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