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)
| 561 | |
| 562 | |
| 563 | def 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] |
no outgoing calls
no test coverage detected