Batched Strided Convolution
Problem statement
Implement a batched, multi-channel two-dimensional convolution using direct loops.
The logical input tensor has shape [batchSize][channels][height][width]. In the callable representation, input flattens the first two dimensions into shape [batchSize * channels][height][width]; plane b * channels + c is channel c of batch item b.
Each row kernels[f] stores output filter f in channel-major order. Its entry for input channel c, kernel row kr, and kernel column kc is at index c * kernelHeight * kernelWidth + kr * kernelWidth + kc.
Use the supplied positive stride, no padding, and neural-network cross-correlation semantics: do not reverse the kernel, add a bias, or apply an activation.
The output height is floor((height - kernelHeight) / stride) + 1, and the output width is defined analogously. Return the output with batch and output-channel dimensions flattened in batch-major order: plane b * outputChannels + f stores filter f for batch item b.
Function
convolveBatched(input: int[][][], batchSize: int, channels: int, kernels: int[][], kernelHeight: int, kernelWidth: int, stride: int) → int[][][]Examples
Example 1
input = [[[1,2,3],[4,5,6],[7,8,9]]]batchSize = 1channels = 1kernels = [[1,0,0,-1]]kernelHeight = 2kernelWidth = 2stride = 1return = [[[-4,-4],[-4,-4]]]There is one batch item and one filter. At the top-left position, the sum is 1 * 1 + 2 * 0 + 4 * 0 + 5 * (-1) = -4.
Example 2
input = [[[1,2,3],[4,5,6],[7,8,9]],[[9,8,7],[6,5,4],[3,2,1]]]batchSize = 2channels = 1kernels = [[1,1,1,1]]kernelHeight = 2kernelWidth = 2stride = 2return = [[[12]],[[28]]]The stride leaves one valid window per batch item. Their sums are 12 and 28.
Example 3
input = [[[1,2,3],[4,5,6]],[[10,20,30],[40,50,60]]]batchSize = 1channels = 2kernels = [[1,1,0,0],[0,0,1,-1]]kernelHeight = 1kernelWidth = 2stride = 1return = [[[3,5],[9,11]],[[-10,-10],[-10,-10]]]The first filter adds adjacent values from channel 0. The second subtracts adjacent values in channel 1.
Constraints
1 <= batchSize <= 4.1 <= channels <= 8.input.length == batchSize * channels.1 <= height, width <= 20.1 <= kernels.length <= 8.1 <= kernelHeight <= heightand1 <= kernelWidth <= width.- Every row of
kernelshas lengthchannels * kernelHeight * kernelWidth. 1 <= stride <= max(height, width).- All input and kernel values are between
-100and100, inclusive. - Every output value fits in a signed 32-bit integer.