bglearning/tapas-tokenizer-viz
0
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 