MCPcopy Create free account
hub / github.com/apple/axlearn / conv_transpose_output_shape

Function conv_transpose_output_shape

axlearn/common/convolution.py:1077–1138  ·  view source on GitHub ↗

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],
)

Source from the content-addressed store, hash-verified

1075
1076
1077def 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

Callers 3

output_shapeMethod · 0.85
output_shapeMethod · 0.85
output_shapeMethod · 0.85

Calls 1

conv_dilate_windowFunction · 0.85

Tested by

no test coverage detected