MCPcopy Create free account
hub / github.com/PaddlePaddle/Research / pad_batch_data

Function pad_batch_data

NLP/EMNLP2019-MAL/src/reader.py:563–617  ·  view source on GitHub ↗

Pad the instances to the max sequence length in batch, and generate the corresponding position data and attention bias.

(insts,
                   pad_idx,
                   n_head,
                   is_target=False,
                   is_label=False,
                   return_attn_bias=True,
                   return_max_len=True,
                   return_num_token=False)

Source from the content-addressed store, hash-verified

561
562
563def pad_batch_data(insts,
564 pad_idx,
565 n_head,
566 is_target=False,
567 is_label=False,
568 return_attn_bias=True,
569 return_max_len=True,
570 return_num_token=False):
571 """
572 Pad the instances to the max sequence length in batch, and generate the
573 corresponding position data and attention bias.
574 """
575 return_list = []
576 max_len = max(len(inst) for inst in insts)
577 # Any token included in dict can be used to pad, since the paddings' loss
578 # will be masked out by weights and make no effect on parameter gradients.
579 inst_data = np.array(
580 [inst + [pad_idx] * (max_len - len(inst)) for inst in insts])
581 return_list += [inst_data.astype("int64").reshape([-1, 1])]
582 if is_label: # label weight
583 inst_weight = np.array(
584 [[1.] * len(inst) + [0.] * (max_len - len(inst)) for inst in insts])
585 return_list += [inst_weight.astype("float32").reshape([-1, 1])]
586 else: # position data
587 inst_pos = np.array([
588 list(range(0, len(inst))) + [0] * (max_len - len(inst))
589 for inst in insts
590 ])
591 return_list += [inst_pos.astype("int64").reshape([-1, 1])]
592 if return_attn_bias:
593 if is_target:
594 # This is used to avoid attention on paddings and subsequent
595 # words.
596 slf_attn_bias_data = np.ones((inst_data.shape[0], max_len, max_len))
597 slf_attn_bias_data = np.triu(slf_attn_bias_data,
598 1).reshape([-1, 1, max_len, max_len])
599 slf_attn_bias_data = np.tile(slf_attn_bias_data,
600 [1, n_head, 1, 1]) * [-1e9]
601 else:
602 # This is used to avoid attention on paddings.
603 slf_attn_bias_data = np.array([[0] * len(inst) + [-1e9] *
604 (max_len - len(inst))
605 for inst in insts])
606 slf_attn_bias_data = np.tile(
607 slf_attn_bias_data.reshape([-1, 1, 1, max_len]),
608 [1, n_head, max_len, 1])
609 return_list += [slf_attn_bias_data.astype("float32")]
610 if return_max_len:
611 return_list += [max_len]
612 if return_num_token:
613 num_token = 0
614 for inst in insts:
615 num_token += len(inst)
616 return_list += [num_token]
617 return return_list if len(return_list) > 1 else return_list[0]

Callers 2

prepare_batch_inputFunction · 0.70
prepare_batch_inputFunction · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected