Contents

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) noexcept

Builds a view that records through recorder, which must outlive the view.

explicit SpatialOps(Engine &engine) noexcept

Builds a view that records on engine's own graph.

SpatialOps(const SpatialOps &)=default
SpatialOps & operator=(const SpatialOps &)=delete
SpatialOps(SpatialOps &&)=default
SpatialOps & operator=(SpatialOps &&)=delete
~SpatialOps()=default
TaskHandle 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_h

The output's row count.

out_w

The output's column count.

Exceptions

std::runtime_error

at 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_error

as 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_h

The window grid's row count, im2col's H_o.

out_w

The window grid's column count, im2col's W_o.

Exceptions

std::runtime_error

at 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

indices

Must be INT64 and shaped like out.

Exceptions

std::runtime_error

at 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

indices

The INT64 argmax buffer pool_max wrote.

g

Must be the geometry the forward used.

Exceptions

std::runtime_error

at 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_error

at 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

g

Must be the geometry the forward used.

Exceptions

std::runtime_error

at 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

out

Its extents must be exactly scale_h * H by scale_w * W.

scale_h

Must be positive, as must scale_w.

Exceptions

std::runtime_error

at 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_h

The factor the forward used, as is scale_w; both must be positive.

Exceptions

std::runtime_error

at 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_error

at 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

dst

Its channel count decides how many channels are copied.

Exceptions

std::runtime_error

at 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

gamma

The rank-1 [C] scale.

beta

The rank-1 [C] shift.

Exceptions

std::runtime_error

at 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_error

at 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

out

Its logical shape must match in's.

Exceptions

std::runtime_error

at record time when the ranks, dtypes or logical shapes disagree.

See also

include/SushiBLAS/engine/README.md