Go from token-level labels to list of entities (start, end, class).
(sentence_tags, is_head=None, segment_id=None, inv_label_mapping=None, ignore_label=list([]))
| 63 | tokens_b.pop() |
| 64 | |
| 65 | def get_span_labels(sentence_tags, is_head=None, segment_id=None, inv_label_mapping=None, ignore_label=list([])): |
| 66 | """Go from token-level labels to list of entities (start, end, class).""" |
| 67 | if inv_label_mapping: |
| 68 | sentence_tags = [inv_label_mapping[i] for i in sentence_tags] |
| 69 | filtered_sentence_tag = [] |
| 70 | if is_head: |
| 71 | # assert(len(sentence_tags) == len(is_head)) |
| 72 | |
| 73 | for idx, (head, segment) in enumerate(zip(is_head, segment_id)): |
| 74 | if (head == 1 or head == True) and (segment == 0 or segment == True): |
| 75 | if sentence_tags[idx] != 'X': |
| 76 | filtered_sentence_tag.append(sentence_tags[idx]) |
| 77 | else: |
| 78 | filtered_sentence_tag.append("O") |
| 79 | if filtered_sentence_tag: |
| 80 | sentence_tags = filtered_sentence_tag |
| 81 | span_labels = [] |
| 82 | last = 'O' |
| 83 | start = -1 |
| 84 | for i, tag in enumerate(sentence_tags): |
| 85 | items = (None, 'O') if tag == 'O' else tag.split('-', 1) |
| 86 | pos, _ = items if len(items) == 2 else (items[0], None) |
| 87 | if (pos == 'S' or pos == 'B' or tag == 'O') and last != 'O': |
| 88 | span_labels.append((start, i - 1, None if len(last.split('-', 1)) != 2 else last.split('-', 1)[-1])) |
| 89 | if pos == 'B' or pos == 'S' or last == 'O': |
| 90 | start = i |
| 91 | last = tag |
| 92 | if sentence_tags[-1] != 'O': |
| 93 | span_labels.append((start, len(sentence_tags) - 1, |
| 94 | None if len(last.split('-', 1)) != 2 else last.split('-', 1)[-1])) |
| 95 | |
| 96 | # This code has problem! |
| 97 | # for item in span_labels: |
| 98 | # if item[2] in ignore_label: |
| 99 | # span_labels.remove(item) |
| 100 | |
| 101 | filtered_labels = [] |
| 102 | for item in span_labels: |
| 103 | if item[2] not in ignore_label: |
| 104 | filtered_labels.append(item) |
| 105 | |
| 106 | return set(filtered_labels), sentence_tags |
| 107 | |
| 108 | def filter_head_prediction(sentence_tags, is_head): |
| 109 | filtered_sentence_tag = [] |