Contents

struct

SushiAI::NN::Conv2d

Convolves an NHWC activation with a bank of learnable filters.

Declared in
include/SushiAI/nn/conv2d.hpp

Public attributes

int64_t in_channels = 0

Holds the channel count of the input.

int64_t out_channels = 0

Holds the number of filters, which is the channel count of the result.

Graph::ConvGeometry geometry {}

Holds the window; its top and left pads apply to both sides of an axis.

bool use_bias = true

Tells whether the layer adds a learnable bias per filter.

Parameter weight {}

Holds the [out_channels, kernel_h, kernel_w, in_channels] filters.

Parameter bias {}

Holds the [out_channels] bias; undeclared when use_bias is false.

Static public member functions

static constexpr auto fields()

Lists this layer's parameters for the generic walk.

Public member functions

Graph::ValueId forward(const Graph::GraphBuilder &builder, Graph::ValueId x)

Traces the convolution, declaring its parameters on the first call.

Parameters

x

An NHWC [N, H, W, in_channels] value that the window fits.

Returns

The [N, out_h, out_w, out_channels] result's id.

Shape output_shape(const Shape &input) const

Computes the per-sample shape forward() produces.

Parameters

input

A per-sample [H, W, in_channels] shape that the window fits.

Returns

The per-sample [out_h, out_w, out_channels] shape.