Team Ai
Apppublic

bglearning/tapas-tokenizer-viz

sourceHugging Facebsd-3-clauseupdated 3y agoView on Hugging Face
0likes
tapas_visualizer.py191 linesDownload Raw Back to root
1"""Visualizer for TAPAS2 3Implementation heavily based on4`EncodingVisualizer` from `tokenizers.tools`.5"""6import os7from typing import Any, List, Dict8 9from collections import defaultdict10 11import pandas as pd12 13from transformers import TapasTokenizer14 15dirname = os.path.dirname(__file__)16css_filename = os.path.join(dirname, "tapas-styles.css")17with open(css_filename) as f:18    css = f.read()19 20 21def HTMLBody(table_html: str, css_styles: str = css) -> str:22    """23    Generates the full html with css from a list of html spans24 25    Args:26        table_html (str):27            The html string of the table28 29        css_styles (str):30            CSS styling to be embedded inline31 32    Returns:33        :obj:`str`: An HTML string with style markup34    """35    return f"""36    <html>37        <head>38            <style>39                {css_styles}40            </style>41        </head>42        <body>43            <div class="tokenized-text" dir=auto>44            {table_html}45            </div>46        </body>47    </html>48    """49 50 51class TapasVisualizer:52    def __init__(self, tokenizer: TapasTokenizer) -> None:53        self.tokenizer = tokenizer54 55    def normalize_token_str(self, token_str: str) -> str:56        # Normalize subword tokens to org subword str57        return token_str.replace("##", "")58 59    def style_span(self, span_text: str, css_classes: List[str]) -> str:60        css = f'''class="{' '.join(css_classes)}"'''61        return f"<span {css} >{span_text}</span>"62 63    def text_to_html(self, org_text: str, tokens: List[str]) -> str:64        """Create html based on the original text and its tokens.65 66        Note: The tokens need to be in same order as in the original text67 68        Args:69            org_text (str): Original string before tokenization70            tokens (List[str]): The tokens of org_text71 72        Returns:73            str: html with styling for the tokens74        """75        if len(tokens) == 0:76            print(f"Empty tokens for: {org_text}")77            return ""78 79        cur_token_id = 080        cur_token = self.normalize_token_str(tokens[cur_token_id])81 82        # Loop through each character83        next_start = 084        last_end = 085        spans = []86 87        while next_start < len(org_text):88            candidate = org_text[next_start : next_start + len(cur_token)]89 90            # The tokenizer performs lowercasing; so check against lowercase91            if candidate.lower() == cur_token:92                if last_end != next_start:93                    # There was token-less text (probably whitespace)94                    # in the middle95                    spans.append(96                        self.style_span(org_text[last_end:next_start], ["non-token"])97                    )98 99                odd_or_even = "even-token" if cur_token_id % 2 == 0 else "odd-token"100                spans.append(self.style_span(candidate, ["token", odd_or_even]))101                next_start += len(cur_token)102                last_end = next_start103                cur_token_id += 1104                if cur_token_id >= len(tokens):105                    break106                cur_token = self.normalize_token_str(tokens[cur_token_id])107            else:108                next_start += 1109 110        if last_end != len(org_text):111            spans.append(self.style_span(org_text[last_end:next_start], ["non-token"]))112 113        return spans114 115    def cells_to_html(116        self,117        cell_vals: List[List[str]],118        cell_tokens: Dict,119        row_id_start: int = 0,120        cell_element: str = "td",121        cumulative_cnt: int = 0,122        table_html: str = "",123    ) -> str:124        for row_id, row in enumerate(cell_vals, start=row_id_start):125            row_html = ""126            row_token_cnt = 0127            for col_id, cell in enumerate(row, start=1):128                cur_cell_tokens = cell_tokens[(row_id, col_id)]129                span_htmls = self.text_to_html(cell, cur_cell_tokens)130                cell_html = "".join(span_htmls)131                row_html += f"<{cell_element}>{cell_html}</{cell_element}>"132                row_token_cnt += len(cur_cell_tokens)133            cumulative_cnt += row_token_cnt134            cnt_html = (135                f'<td style="border: none;" align="right">'136                f'{self.style_span(str(cumulative_cnt), ["non-token", "count"])}'137                "</td>"138                f'<td style="border: none;" align="right">'139                f'{self.style_span(f"<+{row_token_cnt}", ["non-token", "count"])}'140                "</td>"141            )142            row_html = cnt_html + row_html143            table_html += f"<tr>{row_html}</tr>"144 145        return table_html, cumulative_cnt146 147    def __call__(self, table: pd.DataFrame) -> Any:148        tokenized = self.tokenizer(table)149 150        cell_tokens = defaultdict(list)151 152        for id_ind, input_id in enumerate(tokenized["input_ids"]):153            input_id = int(input_id)154            # 'prev_label', 'column_rank', 'inv_column_rank', 'numeric_relation'155            # not required156            segment_id, col_id, row_id, *_ = tokenized["token_type_ids"][id_ind]157            token_text = self.tokenizer._convert_id_to_token(input_id)158            if int(segment_id) == 1:159                cell_tokens[(row_id, col_id)].append(token_text)160 161        table_html, cumulative_cnt = self.cells_to_html(162            cell_vals=[table.columns],163            cell_tokens=cell_tokens,164            row_id_start=0,165            cell_element="th",166            cumulative_cnt=0,167            table_html="",168        )169 170        table_html, cumulative_cnt = self.cells_to_html(171            cell_vals=table.values,172            cell_tokens=cell_tokens,173            row_id_start=1,174            cell_element="td",175            cumulative_cnt=cumulative_cnt,176            table_html=table_html,177        )178        top_label = self.style_span("#Tokens", ["count"])179        top_label_cnt = self.style_span(f"(Total: {cumulative_cnt})", ["count"])180 181        table_html = (182            '<tr style="line-height: 2rem">'183            f'<td style="border: none;" colspan="2" align="left">{top_label}</td>'184            f'<td style="border: none;" colspan="1" align="left">{top_label_cnt}</td>'185            "</tr>"186            f"{table_html}"187        )188 189        table_html = f"<table>{table_html}</table>"190        return HTMLBody(table_html)191