(tensor=Union[List[torch.Tensor], torch.Tensor],
ndim: int = 3,
batch_size: int = None,
padding_mode: Literal['ones', 'zeros', 'repeat',
'none'] = 'none')
| 71 | |
| 72 | |
| 73 | def align_input_to_padded(tensor=Union[List[torch.Tensor], torch.Tensor], |
| 74 | ndim: int = 3, |
| 75 | batch_size: int = None, |
| 76 | padding_mode: Literal['ones', 'zeros', 'repeat', |
| 77 | 'none'] = 'none'): |
| 78 | if isinstance(tensor, list): |
| 79 | for i in range(len(tensor)): |
| 80 | if tensor[i].dim == ndim: |
| 81 | tensor[i] = tensor[i][0] |
| 82 | tensor = list_to_padded(tensor, equisized=True) |
| 83 | assert tensor.ndim in (ndim, ndim - 1) |
| 84 | if tensor.ndim == ndim - 1: |
| 85 | tensor = tensor.unsqueeze(0) |
| 86 | |
| 87 | if batch_size is not None: |
| 88 | current_batch_size = tensor.shape[0] |
| 89 | if current_batch_size == 1: |
| 90 | tensor = tensor.repeat_interleave(batch_size, 0) |
| 91 | elif current_batch_size < batch_size: |
| 92 | if padding_mode == 'ones': |
| 93 | tensor = torch.cat([ |
| 94 | tensor, |
| 95 | torch.ones_like(tensor)[:1].repeat_interleave( |
| 96 | batch_size - current_batch_size, 0) |
| 97 | ]) |
| 98 | elif padding_mode == 'ones': |
| 99 | tensor = torch.cat([ |
| 100 | tensor, |
| 101 | torch.zeros_like(tensor)[:1].repeat_interleave( |
| 102 | batch_size - current_batch_size, 0) |
| 103 | ]) |
| 104 | elif padding_mode == 'repeat': |
| 105 | tensor = tensor.repeat_interleave( |
| 106 | batch_size // current_batch_size + 1, 0)[:batch_size] |
| 107 | else: |
| 108 | raise ValueError('Wrong batch_size to allocate,' |
| 109 | ' please specify padding mode.') |
| 110 | elif current_batch_size > batch_size: |
| 111 | tensor = tensor[:batch_size] |
| 112 | |
| 113 | return tensor |
no outgoing calls
no test coverage detected