Pad sequence by some side Args: sequences: The input sequences in tensor. padding_value: The padding value padding_side: The padding side Returns: A tensor after padding
(sequences: List[torch.Tensor],
padding_value: float = 0.,
padding_side: Literal['right', 'left'] = 'right')
| 865 | |
| 866 | @staticmethod |
| 867 | def pad_sequence(sequences: List[torch.Tensor], |
| 868 | padding_value: float = 0., |
| 869 | padding_side: Literal['right', 'left'] = 'right') -> torch.Tensor: |
| 870 | """Pad sequence by some side |
| 871 | |
| 872 | Args: |
| 873 | sequences: The input sequences in tensor. |
| 874 | padding_value: The padding value |
| 875 | padding_side: The padding side |
| 876 | |
| 877 | Returns: |
| 878 | A tensor after padding |
| 879 | """ |
| 880 | padding_right = padding_side == 'right' |
| 881 | if padding_right: |
| 882 | return pad_sequence(sequences, batch_first=True, padding_value=padding_value) |
| 883 | |
| 884 | max_len = max([s.size(0) for s in sequences]) |
| 885 | |
| 886 | padded_sequences = [] |
| 887 | for seq in sequences: |
| 888 | pad_length = max_len - seq.size(0) |
| 889 | pad_tuple = [0] * ((seq.dim() - 1) * 2) + [pad_length, 0] |
| 890 | padded_seq = F.pad(seq, tuple(pad_tuple), 'constant', padding_value) |
| 891 | padded_sequences.append(padded_seq) |
| 892 | |
| 893 | return torch.stack(padded_sequences) |
| 894 | |
| 895 | def data_collator(self, batch: List[Dict[str, Any]], padding_to: Optional[int] = None) -> Dict[str, Any]: |
| 896 | """ |
no test coverage detected