Contents

struct

SushiAI::NN::Flatten

Folds [N, d1, ..., dk] into [N, d1 * ... * dk] without moving data.

Declared in
include/SushiAI/nn/flatten.hpp

Static public member functions

static constexpr auto fields()

Lists this layer's parameters, of which there are none.

Public member functions

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

Traces the layer as one RESHAPE, or as nothing for a rank-2 input.

Parameters

x

An NHWC input of rank 2 or more, batch axis first.

Returns

The [N, d1 * ... * dk] result's id; x itself when it is rank 2.

Shape output_shape(const Shape &input) const

Computes the per-sample shape forward() produces.

Parameters

input

A per-sample shape of rank 1 or more.

Returns

The rank-1 shape holding the element count of input.