Returns output size for conv transpose. Each mode follows the formulas below, * SAME: padding=(min(window-1, ceil((w+s-2)/2)), max(stride-1, floor((w+s-2)/2))) pad_total = window+stride-2 output_size = input_size*stride * VALID: padding=(window-1, max(stride-1, window-1)
(
in_shape: Sequence[Optional[int]],
*,
window: Sequence[int],
strides: Sequence[int],
padding: ConvPaddingType,
dilation: Sequence[int],
)
| 1075 | |
| 1076 | |
| 1077 | def conv_transpose_output_shape( |
| 1078 | in_shape: Sequence[Optional[int]], |
| 1079 | *, |
| 1080 | window: Sequence[int], |
| 1081 | strides: Sequence[int], |
| 1082 | padding: ConvPaddingType, |
| 1083 | dilation: Sequence[int], |
| 1084 | ) -> Sequence[int]: |
| 1085 | """Returns output size for conv transpose. |
| 1086 | |
| 1087 | Each mode follows the formulas below, |
| 1088 | * SAME: padding=(min(window-1, ceil((w+s-2)/2)), max(stride-1, floor((w+s-2)/2))) |
| 1089 | pad_total = window+stride-2 |
| 1090 | output_size = input_size*stride |
| 1091 | * VALID: padding=(window-1, max(stride-1, window-1)) |
| 1092 | pad_total = window+stride-2 + max(window-stride, 0) |
| 1093 | output_size = input_size*stride + max(window-stride, 0) |
| 1094 | * CAUSAL: padding=(window-1, stride-1) |
| 1095 | pad_total = window+stride-2 |
| 1096 | output_size = input_size*stride |
| 1097 | |
| 1098 | Note: In the above equation, `window` will be replaced with `dilate_window` when dilation > 1. |
| 1099 | dilate_window = (window - 1) * dilation + 1. Check conv_dilate_window() |
| 1100 | |
| 1101 | Refer to |
| 1102 | https://towardsdatascience.com/understand-transposed-convolutions-and-build-your-own-transposed-convolution-layer-from-scratch-4f5d97b2967 |
| 1103 | |
| 1104 | Args: |
| 1105 | in_shape: convolution lhs shape. |
| 1106 | window: convolution window. |
| 1107 | strides: convolution strides. |
| 1108 | padding: convolution padding. |
| 1109 | dilation: convolution dilation. |
| 1110 | |
| 1111 | Returns: |
| 1112 | The output shape. |
| 1113 | |
| 1114 | Raises: |
| 1115 | ValueError: If the length of in_shape, window, strides, and padding are not equal. |
| 1116 | """ |
| 1117 | if len(in_shape) != len(window) or len(in_shape) != len(strides): |
| 1118 | raise ValueError( |
| 1119 | f"len(in_shape) = {len(in_shape)} must be equal to " |
| 1120 | f"len(window) = {len(window)} and len(strides) = {len(strides)}" |
| 1121 | ) |
| 1122 | |
| 1123 | window = conv_dilate_window(window=window, dilation=dilation) |
| 1124 | |
| 1125 | def output_shape(in_shape: Optional[int], window: int, stride: int): |
| 1126 | if in_shape is None: |
| 1127 | return None |
| 1128 | |
| 1129 | if padding == "SAME": |
| 1130 | return in_shape * stride |
| 1131 | elif padding == "VALID": |
| 1132 | return in_shape * stride + max(window - stride, 0) |
| 1133 | elif padding == "CAUSAL": |
| 1134 | return in_shape * stride |
no test coverage detected