Skip to content
heliaCORE
API reference
HELIA HUB

arm_nn_broadcast_walk.h

Machine-readable model

function

Check that one NHWC dimension of two operands broadcasts to the output dimension.

Include/Internal/arm_nn_broadcast_walk.h:34

static int32_t arm_nn_broadcast_dim_valid(const int32_t dim_1, const int32_t dim_2, const int32_t dim_out)

Check that one NHWC dimension of two operands broadcasts to the output dimension.

TensorFlow Lite broadcast rules: both inputs must be at least 1 (an empty tensor is rejected rather than treated as a no-op), they must be equal or one of them must be 1, and the output dimension must be the larger of the two.

Parameters of arm_nn_broadcast_dim_valid
NameTypeDescription
dim_1const int32_t
dim_2const int32_t
dim_outconst int32_t
function

Check that two NHWC operands broadcast to the given output shape.

Include/Internal/arm_nn_broadcast_walk.h:54

static int32_t arm_nn_broadcast_dims_valid(
const cmsis_nn_dims *dims_1,
const cmsis_nn_dims *dims_2,
const cmsis_nn_dims *dims_out
)

Check that two NHWC operands broadcast to the given output shape.

Every kernel that uses ARM_NN_BROADCAST_WALK_NHWC must reject arguments that fail this check, since the walk indexes each input by its own dims and writes the output by the output dims.

Parameters of arm_nn_broadcast_dims_valid
NameTypeDescription
dims_1const cmsis_nn_dims *
dims_2const cmsis_nn_dims *
dims_outconst cmsis_nn_dims *
macro

Walk an NHWC broadcast of two operands, calling a contiguous kernel on each run.

Include/Internal/arm_nn_broadcast_walk.h:89

#define ARM_NN_BROADCAST_WALK_NHWC(IN_TYPE, OUT_TYPE, in_1, dims_1, in_2, dims_2, out, dims_out, FULL, SCALAR_1, SCALAR_2)

Walk an NHWC broadcast of two operands, calling a contiguous kernel on each run.

Each input is indexed by its own dims: a dimension of 1 has stride 0 and is broadcast, any other dimension equals the output dimension and strides normally. The longest contiguous run whose shapes agree is handed to the caller’s kernels, so the common cases (identical shapes, a single scalar, per-batch, per-row, per-channel) each cost one call per run.

Preconditions: arm_nn_broadcast_dims_valid(dims_1, dims_2, dims_out) is non-zero.

Parameters of ARM_NN_BROADCAST_WALK_NHWC
NameDescription
IN_TYPEelement type of the inputs
OUT_TYPEelement type of the output
in_1const IN_TYPE * first input
dims_1const `cmsis_nn_dims` * dims of in_1
in_2const IN_TYPE * second input
dims_2const `cmsis_nn_dims` * dims of in_2
outOUT_TYPE * output, sized by dims_out
dims_outconst `cmsis_nn_dims` * broadcast output dims (see arm_nn_broadcast_dims_valid)
FULLFULL(const IN_TYPE *a, const IN_TYPE *b, OUT_TYPE *o, int32_t n) elementwise kernel over n elements of a and b
SCALAR_1SCALAR_1(const IN_TYPE *scalar, const IN_TYPE *vec, OUT_TYPE *o, int32_t n) kernel where *scalar is one element of in_1 broadcast against n elements of in_2
SCALAR_2SCALAR_2(const IN_TYPE *scalar, const IN_TYPE *vec, OUT_TYPE *o, int32_t n) kernel where *scalar is one element of in_2 broadcast against n elements of in_1; note the operands arrive in reversed order, so an asymmetric kernel must swap them