| 492 | |
| 493 | |
| 494 | def split_tensor_with_padding(input_tensor, pos_idx_range, expand_length=0): |
| 495 | pos_idx_range = [ |
| 496 | [idx[0] - expand_length, idx[1] + expand_length] for idx in pos_idx_range |
| 497 | ] |
| 498 | sub_sequences = [] |
| 499 | seq_len = input_tensor.size(1) |
| 500 | max_valid_idx = seq_len - 1 |
| 501 | k_lens_list = [] |
| 502 | for start, end in pos_idx_range: |
| 503 | pad_front = max(-start, 0) |
| 504 | pad_back = max(end - max_valid_idx, 0) |
| 505 | |
| 506 | valid_start = max(start, 0) |
| 507 | valid_end = min(end, max_valid_idx) |
| 508 | |
| 509 | if valid_start <= valid_end: |
| 510 | valid_part = input_tensor[:, valid_start: valid_end + 1, :] |
| 511 | else: |
| 512 | valid_part = input_tensor.new_zeros((1, 0, input_tensor.size(2))) |
| 513 | |
| 514 | padded_subseq = F.pad( |
| 515 | valid_part, |
| 516 | (0, 0, 0, pad_back + pad_front, 0, 0), |
| 517 | mode="constant", |
| 518 | value=0, |
| 519 | ) |
| 520 | k_lens_list.append(padded_subseq.size(-2) - pad_back - pad_front) |
| 521 | |
| 522 | sub_sequences.append(padded_subseq) |
| 523 | return torch.stack(sub_sequences, dim=1), torch.tensor( |
| 524 | k_lens_list, dtype=torch.long |
| 525 | ) |