MCPcopy Create free account
hub / github.com/arrayfire/arrayfire / conv2DataGradient

Function conv2DataGradient

src/backend/cpu/convolve.cpp:180–211  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

178
179template<typename T>
180Array<T> conv2DataGradient(const Array<T> &incoming_gradient,
181 const Array<T> &original_signal,
182 const Array<T> &original_filter,
183 const Array<T> & /*convolved_output*/,
184 af::dim4 stride, af::dim4 padding,
185 af::dim4 dilation) {
186 const dim4 &cDims = incoming_gradient.dims();
187 const dim4 &sDims = original_signal.dims();
188 const dim4 &fDims = original_filter.dims();
189
190 Array<T> collapsed_filter = flip(original_filter, {1, 1, 0, 0});
191 collapsed_filter = modDims(collapsed_filter,
192 dim4(fDims[0] * fDims[1] * fDims[2], fDims[3]));
193
194 Array<T> collapsed_gradient = incoming_gradient;
195 collapsed_gradient = reorder(collapsed_gradient, dim4(0, 1, 3, 2));
196 collapsed_gradient = modDims(
197 collapsed_gradient, dim4(cDims[0] * cDims[1] * cDims[3], cDims[2]));
198
199 Array<T> res =
200 matmul(collapsed_gradient, collapsed_filter, AF_MAT_NONE, AF_MAT_TRANS);
201 res = modDims(res, dim4(res.dims()[0] / sDims[3], sDims[3],
202 fDims[0] * fDims[1], sDims[2]));
203 res = reorder(res, dim4(0, 2, 3, 1));
204
205 const bool retCols = false;
206 res = wrap_dilated(res, sDims[0], sDims[1], fDims[0], fDims[1], stride[0],
207 stride[1], padding[0], padding[1], dilation[0],
208 dilation[1], retCols);
209
210 return res;
211}
212
213template<typename T>
214Array<T> conv2FilterGradient(const Array<T> &incoming_gradient,

Callers

nothing calls this directly

Calls 7

dim4Class · 0.70
reorderFunction · 0.70
matmulFunction · 0.70
wrap_dilatedFunction · 0.70
flipFunction · 0.50
modDimsFunction · 0.50
dimsMethod · 0.45

Tested by

no test coverage detected