Truncates a sequence pair in place to the maximum length.
(tokens_a, tokens_b, max_length)
| 47 | return vector |
| 48 | |
| 49 | def truncate_seq_pair(tokens_a, tokens_b, max_length): |
| 50 | """Truncates a sequence pair in place to the maximum length.""" |
| 51 | |
| 52 | # This is a simple heuristic which will always truncate the longer sequence |
| 53 | # one token at a time. This makes more sense than truncating an equal percent |
| 54 | # of tokens from each, since if one sequence is very short then each token |
| 55 | # that's truncated likely contains more information than a longer sequence. |
| 56 | while True: |
| 57 | total_length = len(tokens_a) + len(tokens_b) |
| 58 | if total_length <= max_length: |
| 59 | break |
| 60 | if len(tokens_a) > len(tokens_b): |
| 61 | tokens_a.pop() |
| 62 | else: |
| 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).""" |