class
SushiBLAS::SpatialOps
Records windowed gathers, windowed reductions and channel-wise broadcasts.
- Declared in
include/SushiBLAS/engine/spatial.hpp
Public member functions
explicit SpatialOps(TaskRecorder &recorder) noexceptBuilds a view that records through recorder, which must outlive the view.
explicit SpatialOps(Engine &engine) noexceptBuilds a view that records on engine's own graph.
SpatialOps(const SpatialOps &)=defaultSpatialOps & operator=(const SpatialOps &)=deleteSpatialOps(SpatialOps &&)=defaultSpatialOps & operator=(SpatialOps &&)=delete~SpatialOps()=defaultTaskHandle im2col(const Tensor &in, Tensor &cols, const Spatial::ConvGeometry &g, int64_t out_h, int64_t out_w)Gathers every window of an NHWC image into one row of a matrix.
Parameters
out_hThe output's row count.
out_wThe output's column count.
Exceptions
std::runtime_errorat record time when the ranks, dtypes or extents disagree.
See also
include/SushiBLAS/engine/README.md
TaskHandle im2col(const Tensor &in, Tensor &cols, const Spatial::ConvGeometry &g)Gathers every window of an NHWC image into matrix rows, deriving the output extents from cols.
Precondition
The padding is symmetric.
Exceptions
std::runtime_erroras the other overload, or when the width does not divide the row count.
See also
include/SushiBLAS/engine/README.md
TaskHandle col2im(const Tensor &cols, Tensor &image, const Spatial::ConvGeometry &g, int64_t out_h, int64_t out_w)Scatters an im2col matrix back onto an image by summing overlaps.
Parameters
out_hThe window grid's row count, im2col's
H_o.out_wThe window grid's column count, im2col's
W_o.
Exceptions
std::runtime_errorat record time when the ranks, dtypes or extents disagree.
See also
include/SushiBLAS/engine/README.md
TaskHandle pool_max(const Tensor &in, Tensor &out, Tensor &indices, const Spatial::ConvGeometry &g)Takes each window's maximum and records where it came from.
Parameters
indicesMust be INT64 and shaped like out.
Exceptions
std::runtime_errorat record time when the ranks, dtypes or extents disagree.
See also
include/SushiBLAS/engine/README.md
TaskHandle pool_max_backward(const Tensor &dy, const Tensor &indices, Tensor &dx, const Spatial::ConvGeometry &g)Routes each output gradient back to the element it was taken from.
Parameters
indicesThe INT64 argmax buffer pool_max wrote.
gMust be the geometry the forward used.
Exceptions
std::runtime_errorat record time when the ranks, dtypes or extents disagree.
See also
include/SushiBLAS/engine/README.md
TaskHandle pool_avg(const Tensor &in, Tensor &out, const Spatial::ConvGeometry &g)Averages each window, dividing by the window's area.
Exceptions
std::runtime_errorat record time when the ranks, dtypes or extents disagree.
See also
include/SushiBLAS/engine/README.md
TaskHandle pool_avg_backward(const Tensor &dy, Tensor &dx, const Spatial::ConvGeometry &g)Spreads each output gradient evenly over the window it came from.
Parameters
gMust be the geometry the forward used.
Exceptions
std::runtime_errorat record time when the ranks, dtypes or extents disagree.
See also
include/SushiBLAS/engine/README.md
TaskHandle upsample_nearest(const Tensor &in, Tensor &out, int64_t scale_h, int64_t scale_w)Repeats each input pixel over a scale_h by scale_w block.
Parameters
outIts extents must be exactly
scale_h * Hbyscale_w * W.scale_hMust be positive, as must scale_w.
Exceptions
std::runtime_errorat record time when the ranks, dtypes or extents disagree.
See also
include/SushiBLAS/engine/README.md
TaskHandle upsample_nearest_backward(const Tensor &dy, Tensor &dx, int64_t scale_h, int64_t scale_w)Sums each output block back onto the pixel that produced it.
Parameters
scale_hThe factor the forward used, as is scale_w; both must be positive.
Exceptions
std::runtime_errorat record time when the ranks, dtypes or extents disagree.
See also
include/SushiBLAS/engine/README.md
TaskHandle concat_channels(const Tensor &src, Tensor &dst, int64_t dst_channel_offset)Copies all of src's channels into dst starting at a channel offset.
Precondition
src and dst agree on N, H and W.
Exceptions
std::runtime_errorat record time, also when the range runs past dst's channel count.
See also
include/SushiBLAS/engine/README.md
TaskHandle slice_channels(const Tensor &src, int64_t src_channel_offset, Tensor &dst)Copies a channel range out of src, filling all of dst.
Parameters
dstIts channel count decides how many channels are copied.
Exceptions
std::runtime_errorat record time, also when the range runs past src's channel count.
See also
include/SushiBLAS/engine/README.md
TaskHandle channel_affine(const Tensor &in, const Tensor &gamma, const Tensor &beta, Tensor &out)Applies a per-channel scale and shift: out[n,h,w,c] = in[n,h,w,c] * gamma[c] + beta[c].
Parameters
gammaThe rank-1
[C]scale.betaThe rank-1
[C]shift.
Exceptions
std::runtime_errorat record time when the ranks, dtypes or extents disagree.
See also
include/SushiBLAS/engine/README.md
TaskHandle broadcast_spatial(const Tensor &in, Tensor &out, double scale)Spreads an [N, C] value over [N, H, W, C]: out[n,h,w,c] = in[n,c] * scale.
Exceptions
std::runtime_errorat record time when the ranks, dtypes or the N/C extents disagree.
See also
include/SushiBLAS/engine/README.md
TaskHandle layout_convert(const Tensor &in, Tensor &out)Copies a rank-4 tensor element by logical (n, c, h, w) position.
Parameters
outIts logical shape must match in's.
Exceptions
std::runtime_errorat record time when the ranks, dtypes or logical shapes disagree.
See also
include/SushiBLAS/engine/README.md

