MCPcopy Create free account
hub / github.com/google-research/language / bert_preprocess

Function bert_preprocess

language/conpono/reconstruct/preprocess.py:477–575  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

475
476
477def 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)

Callers

nothing calls this directly

Calls 12

wordpiece_tokenizeFunction · 0.85
collapse_dimsFunction · 0.85
truncate_segment_tokensFunction · 0.85
add_special_tokensFunction · 0.85
create_segment_idsFunction · 0.85
sample_mask_indicesFunction · 0.85
apply_maskingFunction · 0.85
pad_to_lengthFunction · 0.85
BertInputsClass · 0.85
constantMethod · 0.80
sizeMethod · 0.45

Tested by

no test coverage detected

Used in the wild real call sites across dependent graphs

searching dependent graphs…