Pre-processes a text tuple into BERT format. Args: text_segments (Tensor): 1-D Tensor of un-tokenized text segments (string). These will be concatenated into a single sequence. vocab_table (StaticHashTable): a map from string to integer. max_seq_length (int): max # of wordpiece
(text_segments,
vocab_table,
max_seq_length,
max_predictions_per_seq,
tokenize_fn=None,
mask_rate=0.15,
do_lower_case=True,
cls_token=_CLS_TOKEN,
sep_token=_SEP_TOKEN,
mask_token=_MASK_TOKEN)
| 475 | |
| 476 | |
| 477 | def bert_preprocess(text_segments, |
| 478 | vocab_table, |
| 479 | max_seq_length, |
| 480 | max_predictions_per_seq, |
| 481 | tokenize_fn=None, |
| 482 | mask_rate=0.15, |
| 483 | do_lower_case=True, |
| 484 | cls_token=_CLS_TOKEN, |
| 485 | sep_token=_SEP_TOKEN, |
| 486 | mask_token=_MASK_TOKEN): |
| 487 | """Pre-processes a text tuple into BERT format. |
| 488 | |
| 489 | Args: |
| 490 | text_segments (Tensor): 1-D Tensor of un-tokenized text segments (string). |
| 491 | These will be concatenated into a single sequence. |
| 492 | vocab_table (StaticHashTable): a map from string to integer. |
| 493 | max_seq_length (int): max # of wordpiece tokens, including special tokens. |
| 494 | max_predictions_per_seq (int): max # of masked positions for entire seq. |
| 495 | tokenize_fn: a function that performs a tokenization on a 1D string Tensor |
| 496 | of untokenized strings. By default this uses wordpiece tokenization. |
| 497 | mask_rate (float): percentage of tokens to mask out. If 0, no masking is |
| 498 | performed. |
| 499 | do_lower_case (bool): whether to lowercase text or not. Default is True. |
| 500 | cls_token (unicode): token representing CLS |
| 501 | sep_token (unicode): token representing SEP (separator) |
| 502 | mask_token (unicode): token representing MASK |
| 503 | |
| 504 | Returns: |
| 505 | BertInputs |
| 506 | """ |
| 507 | |
| 508 | if not tokenize_fn: |
| 509 | # pylint: disable=g-long-lambda |
| 510 | tokenize_fn = lambda text_input: wordpiece_tokenize( |
| 511 | text_input=text_segments, |
| 512 | vocab_table=vocab_table, |
| 513 | token_out_type=tf.string, |
| 514 | use_unknown_token=True, |
| 515 | lower_case=do_lower_case) |
| 516 | |
| 517 | segment_tokens_grouped = tokenize_fn(text_segments) |
| 518 | # TODO(kguu): check that BERT handles UNK in the same way. |
| 519 | |
| 520 | # This is a RaggedTensor of shape [batch_size, (num_wordpieces)] |
| 521 | # One row for each segment. Each row is a list of wordpiece tokens. |
| 522 | segment_tokens = collapse_dims(segment_tokens_grouped) |
| 523 | |
| 524 | # Truncate. |
| 525 | num_special_tokens = segment_tokens.nrows() + 1 |
| 526 | segment_tokens_truncated = truncate_segment_tokens( |
| 527 | segment_tokens, max_seq_length - num_special_tokens) |
| 528 | |
| 529 | # Add special tokens. |
| 530 | segment_tokens_with_special_tokens = add_special_tokens( |
| 531 | segment_tokens_truncated, cls_token, sep_token) |
| 532 | |
| 533 | # Compute segment IDs. |
| 534 | segment_ids_2d = create_segment_ids(segment_tokens_with_special_tokens) |
nothing calls this directly
no test coverage detected
searching dependent graphs…