(label_name_list, tokenizer, end_symbol='<end>')
| 24 | |
| 25 | |
| 26 | def get_label_name_tree(label_name_list, tokenizer, end_symbol='<end>'): |
| 27 | sub_token_tree = dict() |
| 28 | |
| 29 | label_tree = dict() |
| 30 | for typename in label_name_list: |
| 31 | after_tokenized = tokenizer.encode(typename, add_special_tokens=False) |
| 32 | label_tree[typename] = after_tokenized |
| 33 | |
| 34 | for _, sub_label_seq in label_tree.items(): |
| 35 | parent = sub_token_tree |
| 36 | for value in sub_label_seq: |
| 37 | if value not in parent: |
| 38 | parent[value] = dict() |
| 39 | parent = parent[value] |
| 40 | |
| 41 | parent[end_symbol] = None |
| 42 | |
| 43 | return sub_token_tree |
| 44 | |
| 45 | |
| 46 | class PrefixTree: |
no outgoing calls
no test coverage detected