Convert a list of 2d tensors into a padded 3d tensor.
(values, pad_idx=0, left_pad=False, shift_right=False, max_len=None)
| 60 | |
| 61 | |
| 62 | def collate_2d(values, pad_idx=0, left_pad=False, shift_right=False, max_len=None): |
| 63 | """Convert a list of 2d tensors into a padded 3d tensor.""" |
| 64 | size = max(v.size(0) for v in values) if max_len is None else max_len |
| 65 | res = values[0].new(len(values), size, values[0].shape[1]).fill_(pad_idx) |
| 66 | |
| 67 | def copy_tensor(src, dst): |
| 68 | assert dst.numel() == src.numel() |
| 69 | if shift_right: |
| 70 | dst[1:] = src[:-1] |
| 71 | else: |
| 72 | dst.copy_(src) |
| 73 | |
| 74 | for i, v in enumerate(values): |
| 75 | copy_tensor(v, res[i][size - len(v):] if left_pad else res[i][:len(v)]) |
| 76 | return res |
| 77 | |
| 78 | |
| 79 | def _is_batch_full(batch, num_tokens, max_tokens, max_sentences): |
nothing calls this directly
no test coverage detected