Make features.
(features, local_radius, relative_pos_max_distance,
use_hard_g2l_mask, padding_id, eos_id, null_id, cls_id,
sep_id, sequence_length, global_sequence_length)
| 441 | |
| 442 | |
| 443 | def features_map_fn(features, local_radius, relative_pos_max_distance, |
| 444 | use_hard_g2l_mask, padding_id, eos_id, null_id, cls_id, |
| 445 | sep_id, sequence_length, global_sequence_length): |
| 446 | """Make features.""" |
| 447 | batch_size = tf.get_static_value(features['token_ids'].shape[0]) |
| 448 | # sequence_lengths = features['token_ids'].row_lengths() |
| 449 | question_lengths = tf.argmax( |
| 450 | tf.equal(features['token_ids'].to_tensor( |
| 451 | shape=(batch_size, global_sequence_length)), sep_id), -1) + 1 |
| 452 | mapped_features = dict( |
| 453 | token_ids=tf.cast( |
| 454 | features['token_ids'].to_tensor(shape=(batch_size, sequence_length)), |
| 455 | tf.int32), |
| 456 | global_token_ids=tf.cast( |
| 457 | features['global_token_ids'].to_tensor( |
| 458 | shape=(batch_size, global_sequence_length)), tf.int32), |
| 459 | segment_ids=tf.cast( |
| 460 | features['segment_ids'].to_tensor( |
| 461 | shape=(batch_size, sequence_length)), tf.int32), |
| 462 | ) |
| 463 | relative_pos_generator = RelativePositionGenerator( |
| 464 | max_distance=relative_pos_max_distance) |
| 465 | # Only do long-to-long attention for non-null tokens. |
| 466 | # Let the null token attend to itself. |
| 467 | l2l_att_mask = tf.ones((batch_size, sequence_length, 2 * local_radius + 1), |
| 468 | tf.int32) |
| 469 | l2l_att_mask *= 1 - tf.cast( |
| 470 | tf.logical_or( |
| 471 | tf.equal(mapped_features['token_ids'], padding_id), |
| 472 | tf.equal(mapped_features['token_ids'], null_id)), |
| 473 | tf.int32)[:, :, tf.newaxis] |
| 474 | l2l_relative_att_ids = relative_pos_generator.make_local_relative_att_ids( |
| 475 | seq_len=sequence_length, local_radius=local_radius, batch_size=batch_size) |
| 476 | # |
| 477 | l2g_att_mask = tf.ones((batch_size, sequence_length, global_sequence_length), |
| 478 | tf.int32) |
| 479 | l2g_att_mask *= tf.cast( |
| 480 | tf.not_equal(mapped_features['token_ids'], padding_id), |
| 481 | tf.int32)[:, :, tf.newaxis] |
| 482 | l2g_att_mask *= tf.cast( |
| 483 | tf.not_equal(mapped_features['global_token_ids'], padding_id), |
| 484 | tf.int32)[:, tf.newaxis, :] |
| 485 | l2g_relative_att_ids = tf.fill( |
| 486 | (batch_size, sequence_length, global_sequence_length), |
| 487 | relative_pos_generator.relative_vocab_size + 1) |
| 488 | # |
| 489 | g2g_att_mask = tf.ones( |
| 490 | (batch_size, global_sequence_length, global_sequence_length), tf.int32) |
| 491 | g2g_att_mask *= tf.cast( |
| 492 | tf.not_equal(mapped_features['global_token_ids'], padding_id), |
| 493 | tf.int32)[:, :, tf.newaxis] |
| 494 | g2g_relative_att_ids = relative_pos_generator.make_relative_att_ids( |
| 495 | seq_len=global_sequence_length, batch_size=batch_size) |
| 496 | global_sentence_mask = tf.equal(mapped_features['global_token_ids'], eos_id) |
| 497 | global_question_mask = tf.logical_not( |
| 498 | tf.logical_or( |
| 499 | tf.logical_or( |
| 500 | tf.equal(mapped_features['global_token_ids'], cls_id), |
no test coverage detected