blob: 7d2b2305c8aacf8c73d6f5580c076a2f18fe9b88 [file] [edit]
// Copyright 2020-2025 Google LLC
//
// This source code is licensed under the BSD-style license found in the
// LICENSE file in the root directory of this source tree.
#include "src/xnnpack/subgraph.h"
#include <assert.h>
#include <inttypes.h> // IWYU pragma: keep
#include <math.h>
#include <stddef.h>
#include <stdint.h>
#include <stdlib.h>
#include <string.h>
#include "include/experimental.h"
#include "include/xnnpack.h"
#include "src/subgraph/rewrites/cvt_to_fp32.h"
#include "src/xnnpack/allocation-type.h"
#include "src/xnnpack/allocator.h"
#include "src/xnnpack/common.h"
#include "src/xnnpack/config-types.h"
#include "src/xnnpack/config.h"
#include "src/xnnpack/datatype.h"
#include "src/xnnpack/fp16.h"
#include "src/xnnpack/hardware-config.h"
#include "src/xnnpack/internal.h"
#include "src/xnnpack/log.h"
#include "src/xnnpack/math.h"
#include "src/xnnpack/node-type.h"
#include "src/xnnpack/operator.h"
#include "src/xnnpack/params.h"
#ifndef XNN_ENABLE_SPARSE
#error "XNN_ENABLE_SPARSE not defined"
#endif
enum xnn_status xnn_insert_clamp_node(xnn_subgraph_t subgraph, float output_min,
float output_max, uint32_t node_id) {
XNN_RETURN_IF_ERROR(xnn_subgraph_reserve_nodes(subgraph, 1));
XNN_RETURN_IF_ERROR(xnn_subgraph_reserve_values(subgraph, 1));
struct xnn_node* node = &subgraph->nodes[node_id];
uint32_t output_id = node->outputs[0];
struct xnn_value* output_value = &subgraph->values[output_id];
uint32_t new_id = XNN_INVALID_VALUE_ID;
enum xnn_status status;
size_t num_dims = output_value->shape.num_dims;
size_t dims[XNN_MAX_TENSOR_DIMS];
memcpy(dims, output_value->shape.dim, num_dims * sizeof(size_t));
switch (output_value->datatype) {
case xnn_datatype_quint8:
status = xnn_define_quantized_tensor_value(
subgraph, xnn_datatype_quint8, output_value->quantization.zero_point,
output_value->quantization.scale, num_dims, dims, NULL,
/*external_id=*/XNN_INVALID_VALUE_ID, /*flags=*/0, &new_id);
break;
case xnn_datatype_qint8:
status = xnn_define_quantized_tensor_value(
subgraph, xnn_datatype_qint8, output_value->quantization.zero_point,
output_value->quantization.scale, num_dims, dims, NULL,
/*external_id=*/XNN_INVALID_VALUE_ID, /*flags=*/0, &new_id);
break;
default:
status = xnn_define_tensor_value(
subgraph, output_value->datatype, num_dims, dims, NULL,
/*external_id=*/XNN_INVALID_VALUE_ID, /*flags=*/0, &new_id);
break;
}
if (status != xnn_status_success) {
return status;
}
struct xnn_value* new_value = &subgraph->values[new_id];
new_value->size = 0;
node->outputs[0] = new_id;
node->activation.output_min = -INFINITY;
node->activation.output_max = INFINITY;
union xnn_unary_params params;
params.clamp.min = output_min;
params.clamp.max = output_max;
return xnn_define_unary(subgraph, xnn_unary_clamp, &params, new_id, output_id,
/*flags=*/0);
}
enum xnn_status xnn_insert_pack_lh_node(xnn_subgraph_t subgraph,
uint32_t input_id, uint32_t* new_id) {
XNN_RETURN_IF_ERROR(xnn_subgraph_reserve_nodes(subgraph, 1));
XNN_RETURN_IF_ERROR(xnn_subgraph_reserve_values(subgraph, 1));
const struct xnn_value* input = &subgraph->values[input_id];
enum xnn_status status = xnn_status_uninitialized;
switch (input->datatype) {
case xnn_datatype_qint8: {
// Create a copy of the input shape since it might be reallocated by the
// subgraph when the new tensor is added.
struct xnn_shape input_shape = input->shape;
status = xnn_define_quantized_tensor_value(
subgraph, input->datatype, input->quantization.zero_point,
input->quantization.scale, input_shape.num_dims, input_shape.dim,
/*data=*/input->data,
/*external_id=*/XNN_INVALID_VALUE_ID, /*flags=*/0, new_id);
break;
}
case xnn_datatype_fp16:
case xnn_datatype_fp32:
status = xnn_define_tensor_value(subgraph, input->datatype, 0, NULL, NULL,
/*external_id=*/XNN_INVALID_VALUE_ID,
/*flags=*/0, new_id);
break;
default:
XNN_UNREACHABLE;
}
if (status != xnn_status_success) {
return status;
}
return xnn_define_pack_lh(subgraph, input_id, *new_id, /*flags=*/0);
}
enum xnn_status xnn_create_subgraph(uint32_t external_value_ids, uint32_t flags,
xnn_subgraph_t* subgraph_out) {
struct xnn_subgraph* subgraph = NULL;
enum xnn_status status = xnn_status_uninitialized;
if ((xnn_params.init_flags & XNN_INIT_FLAG_XNNPACK) == 0) {
xnn_log_error("failed to create subgraph: XNNPACK is not initialized");
goto error;
}
status = xnn_status_out_of_memory;
subgraph = xnn_allocate_zero_memory(sizeof(struct xnn_subgraph));
if (subgraph == NULL) {
xnn_log_error("failed to allocate %zu bytes for subgraph descriptor",
sizeof(struct xnn_subgraph));
goto error;
}
subgraph->external_value_ids = external_value_ids;
subgraph->values =
xnn_allocate_zero_memory(external_value_ids * sizeof(struct xnn_value));
if (subgraph->values == NULL) {
xnn_log_error("failed to allocate %zu bytes for subgraph values",
(size_t)external_value_ids * sizeof(struct xnn_value));
goto error;
}
for (size_t i = 0; i < external_value_ids; i++) {
subgraph->values[i].id = i;
}
subgraph->num_values = external_value_ids;
subgraph->num_reserved_values = external_value_ids;
*subgraph_out = subgraph;
return xnn_status_success;
error:
xnn_delete_subgraph(subgraph);
return status;
}
enum xnn_status xnn_subgraph_reserve_values(xnn_subgraph_t subgraph,
size_t num_values) {
struct xnn_value* values = subgraph->values;
const size_t size = subgraph->num_values;
const size_t capacity = subgraph->num_reserved_values;
if (capacity < size + num_values) {
const size_t new_capacity =
max(min(capacity * 2, capacity + 512), capacity + max(num_values, 64));
assert(new_capacity >= size + num_values);
values =
xnn_reallocate_memory(values, new_capacity * sizeof(struct xnn_value));
if (values == NULL) {
xnn_log_error("failed to allocate %zu bytes for subgraph values",
new_capacity * sizeof(struct xnn_value));
return xnn_status_out_of_memory;
}
subgraph->num_reserved_values = new_capacity;
subgraph->values = values;
}
return xnn_status_success;
}
struct xnn_value* xnn_subgraph_new_internal_value(xnn_subgraph_t subgraph) {
if (xnn_subgraph_add_internal_values(subgraph, 1) != xnn_status_success) {
return NULL;
}
return subgraph->values + subgraph->num_values - 1;
}
enum xnn_status xnn_subgraph_add_internal_values(xnn_subgraph_t subgraph,
size_t num_values) {
XNN_RETURN_IF_ERROR(xnn_subgraph_reserve_values(subgraph, num_values));
struct xnn_value* new_values = subgraph->values + subgraph->num_values;
for (size_t i = 0; i < num_values; i++) {
assert(new_values + i != NULL);
memset(new_values + i, 0, sizeof(struct xnn_value));
new_values[i].id = subgraph->num_values + i;
}
subgraph->num_values += num_values;
return xnn_status_success;
}
void xnn_node_clear(struct xnn_node* node) {
assert(node != NULL);
memset(node, 0, sizeof(struct xnn_node));
}
void xnn_value_clear(struct xnn_value* value) {
assert(value != NULL);
if ((value->flags & XNN_VALUE_FLAG_NEEDS_CLEANUP) && value->data != NULL) {
XNN_PRAGMA_CLANG("clang diagnostic push")
XNN_PRAGMA_CLANG("clang diagnostic ignored \"-Wcast-qual\"")
xnn_release_memory((void*)value->data);
XNN_PRAGMA_CLANG("clang diagnostic pop")
}
memset(value, 0, sizeof(struct xnn_value));
}
// Copies all fields from `src_node` to `dst_node` but leaves the node ID
// unchanged.
void xnn_node_copy(struct xnn_node* dst_node, const struct xnn_node* src_node) {
const uint32_t node_id = dst_node->id;
*dst_node = *src_node;
dst_node->id = node_id;
}
// Copies all fields from `src_value` to `dst_value` but leaves the value ID
// unchanged.
void xnn_value_copy(struct xnn_value* dst_value,
const struct xnn_value* src_value) {
const uint32_t value_id = dst_value->id;
*dst_value = *src_value;
dst_value->id = value_id;
}
void xnn_runtime_value_copy(struct xnn_runtime_value* dst_value,
const struct xnn_value* src_value) {
// Note: Value ID stays unchanged
dst_value->type = src_value->type;
dst_value->datatype = src_value->datatype;
dst_value->quantization = src_value->quantization;
dst_value->shape = src_value->shape;
dst_value->size = src_value->size;
dst_value->allocation_type = src_value->allocation_type;
dst_value->flags = src_value->flags;
if (src_value->num_consumers == 1) {
dst_value->flags |= XNN_VALUE_FLAG_ONE_CONSUMER;
}
if (src_value->fp16_rewrite.fp16_compatible) {
dst_value->flags |= XNN_VALUE_FLAG_FP16_COMPATIBLE;
}
if (src_value->layout == xnn_layout_type_nchw) {
dst_value->flags |= XNN_VALUE_FLAG_LAYOUT_NCHW;
}
dst_value->data = src_value->data;
dst_value->first_consumer = src_value->first_consumer;
dst_value->fp32_data = src_value->fp32_data;
dst_value->gemm_config = src_value->gemm_config;
}
struct xnn_node* xnn_subgraph_new_node(xnn_subgraph_t subgraph) {
if (xnn_subgraph_add_nodes(subgraph, 1) != xnn_status_success) {
return NULL;
}
return subgraph->nodes + subgraph->num_nodes - 1;
}
enum xnn_status xnn_subgraph_reserve_nodes(xnn_subgraph_t subgraph,
size_t num_nodes) {
struct xnn_node* nodes = subgraph->nodes;
const size_t size = subgraph->num_nodes;
const size_t capacity = subgraph->num_reserved_nodes;
if (capacity < size + num_nodes) {
const size_t new_capacity =
max(min(capacity * 2, capacity + 512), capacity + max(num_nodes, 64));
assert(new_capacity >= size + num_nodes);
nodes =
xnn_reallocate_memory(nodes, new_capacity * sizeof(struct xnn_node));
if (nodes == NULL) {
xnn_log_error("failed to allocate %zu bytes for subgraph nodes",
new_capacity * sizeof(struct xnn_node));
return xnn_status_out_of_memory;
}
subgraph->num_reserved_nodes = new_capacity;
subgraph->nodes = nodes;
}
return xnn_status_success;
}
enum xnn_status xnn_subgraph_add_nodes(xnn_subgraph_t subgraph,
size_t num_nodes) {
XNN_RETURN_IF_ERROR(xnn_subgraph_reserve_nodes(subgraph, num_nodes));
struct xnn_node* new_nodes = subgraph->nodes + subgraph->num_nodes;
for (size_t i = 0; i < num_nodes; i++) {
xnn_node_clear(&new_nodes[i]);
new_nodes[i].id = subgraph->num_nodes + i;
}
subgraph->num_nodes += num_nodes;
return xnn_status_success;
}
static bool is_repeated_input(const struct xnn_node* node, uint32_t idx) {
const uint32_t input_id = node->inputs[idx];
for (int k = 0; k < idx; ++k) {
if (node->inputs[k] == input_id) {
return true;
}
}
return false;
}
void xnn_subgraph_analyze_consumers_and_producers(xnn_subgraph_t subgraph) {
// Initialize producer/consumer fields to safe defaults.
for (uint32_t i = 0; i < subgraph->num_values; i++) {
struct xnn_value* value = &subgraph->values[i];
value->producer = XNN_INVALID_NODE_ID;
value->first_consumer = XNN_INVALID_NODE_ID;
value->num_consumers = 0;
}
// Analyse Nodes' inputs and output and update Values' producer/consumer
// fields
for (uint32_t n = 0; n < subgraph->num_nodes; n++) {
struct xnn_node* node = &subgraph->nodes[n];
if (node->type == xnn_node_type_invalid) {
continue;
}
for (uint32_t i = 0; i < node->num_inputs; i++) {
if (is_repeated_input(node, i)) {
continue;
}
const uint32_t input_id = node->inputs[i];
assert(input_id != XNN_INVALID_VALUE_ID);
assert(input_id < subgraph->num_values);
if (subgraph->values[input_id].num_consumers++ == 0) {
assert(subgraph->values[input_id].first_consumer ==
XNN_INVALID_NODE_ID);
subgraph->values[input_id].first_consumer = n;
subgraph->values[input_id].all_consumers_types_same = true;
} else {
enum xnn_node_type first_consumer_type =
subgraph->nodes[subgraph->values[input_id].first_consumer].type;
subgraph->values[input_id].all_consumers_types_same &=
(first_consumer_type == node->type);
}
}
for (uint32_t o = 0; o < node->num_outputs; o++) {
const uint32_t output_id = node->outputs[o];
assert(output_id < subgraph->num_values);
assert(subgraph->values[output_id].producer == XNN_INVALID_NODE_ID);
subgraph->values[output_id].producer = n;
}
}
// Count extra consumer for Values which are external outputs.
// Remove unreferenced values.
for (uint32_t i = 0; i < subgraph->num_values; i++) {
struct xnn_value* value = &subgraph->values[i];
if (xnn_value_is_external_output(value->flags)) {
value->num_consumers += 1;
}
}
}
#define XNN_LAYOUT_FLAG_COMPATIBLE_NCHW 1
#define XNN_LAYOUT_FLAG_COMPATIBLE_NHWC2NCHW 2
#define XNN_LAYOUT_FLAG_COMPATIBLE_NCHW2NHWC 4
#define XNN_LAYOUT_FLAG_INCOMPATIBLE_CLUSTER 8
static bool all_values_fp(xnn_subgraph_t subgraph,
const struct xnn_node* node) {
for (uint32_t i = 0; i < node->num_inputs; i++) {
const uint32_t input_id = node->inputs[i];
assert(input_id != XNN_INVALID_VALUE_ID);
if (subgraph->values[input_id].datatype != xnn_datatype_fp16 &&
subgraph->values[input_id].datatype != xnn_datatype_fp32) {
return false;
}
}
for (uint32_t i = 0; i < node->num_outputs; i++) {
const uint32_t output_id = node->outputs[i];
assert(output_id != XNN_INVALID_VALUE_ID);
if (subgraph->values[output_id].datatype != xnn_datatype_fp16 &&
subgraph->values[output_id].datatype != xnn_datatype_fp32) {
return false;
}
}
return true;
}
uint32_t xnn_check_nchw_compatibility(xnn_subgraph_t subgraph,
struct xnn_node* node) {
if (!all_values_fp(subgraph, node)) {
if (node->type != xnn_node_type_invalid) {
xnn_log_info("Node %s compute type is incompatible with sparse inference",
xnn_node_type_to_string(node->type));
}
return 0;
}
switch (node->type) {
case xnn_node_type_fully_connected:
return XNN_LAYOUT_FLAG_COMPATIBLE_NCHW;
case xnn_node_type_convolution_2d:
// Supported cases:
// - 1x1 convolution (no stride, no dilation, no padding, no groups)
// - 3x3 stride-2 convolution (no dilation, padding 1 on each side, no
// groups, 3 input channels)
if (node->params.convolution_2d.groups != 1) {
xnn_log_info("Node %s groups (%" PRIu32
") is incompatible with sparse inference",
xnn_node_type_to_string(node->type),
node->params.convolution_2d.groups);
return 0;
}
if ((node->params.convolution_2d.dilation_height |
node->params.convolution_2d.dilation_width) != 1) {
xnn_log_info("Node %s dilation (height=%" PRIu32 ", width=%" PRIu32
") is incompatible with sparse inference",
xnn_node_type_to_string(node->type),
node->params.convolution_2d.dilation_height,
node->params.convolution_2d.dilation_width);
return 0;
}
if ((node->params.convolution_2d.kernel_height |
node->params.convolution_2d.kernel_width) == 1) {
if ((node->params.convolution_2d.input_padding_top |
node->params.convolution_2d.input_padding_right |
node->params.convolution_2d.input_padding_bottom |
node->params.convolution_2d.input_padding_left) != 0) {
xnn_log_info("Node %s (1x1 kernel) padding (top=%" PRIu32
", right=%" PRIu32 ", bottom=%" PRIu32 ", left=%" PRIu32
") is incompatible with sparse inference",
xnn_node_type_to_string(node->type),
node->params.convolution_2d.input_padding_top,
node->params.convolution_2d.input_padding_right,
node->params.convolution_2d.input_padding_bottom,
node->params.convolution_2d.input_padding_left);
return 0;
}
if ((node->params.convolution_2d.subsampling_height |
node->params.convolution_2d.subsampling_width) != 1) {
xnn_log_info("Node %s (1x1 kernel) subsampling (height=%" PRIu32
", width=%" PRIu32
") is incompatible with sparse inference",
xnn_node_type_to_string(node->type),
node->params.convolution_2d.subsampling_height,
node->params.convolution_2d.subsampling_width);
return 0;
}
return XNN_LAYOUT_FLAG_COMPATIBLE_NCHW;
} else if (node->params.convolution_2d.kernel_height == 3 &&
node->params.convolution_2d.kernel_width == 3) {
if (node->params.convolution_2d.input_padding_top != 1 ||
node->params.convolution_2d.input_padding_right != 1 ||
node->params.convolution_2d.input_padding_bottom != 1 ||
node->params.convolution_2d.input_padding_left != 1) {
xnn_log_info("Node %s (3x3 kernel) padding (top=%" PRIu32
", right=%" PRIu32 ", bottom=%" PRIu32 ", left=%" PRIu32
") is incompatible with sparse inference",
xnn_node_type_to_string(node->type),
node->params.convolution_2d.input_padding_top,
node->params.convolution_2d.input_padding_right,
node->params.convolution_2d.input_padding_bottom,
node->params.convolution_2d.input_padding_left);
return 0;
}
if ((node->params.convolution_2d.subsampling_height |
node->params.convolution_2d.subsampling_width) != 2) {
xnn_log_info("Node %s (3x3 kernel) subsampling (height=%" PRIu32
", width=%" PRIu32
") is incompatible with sparse inference",
xnn_node_type_to_string(node->type),
node->params.convolution_2d.subsampling_height,
node->params.convolution_2d.subsampling_width);
return 0;
}
if (node->params.convolution_2d.group_input_channels != 3) {
xnn_log_info(
"Node %s (3x3 kernel) input channels (%zu) is incompatible with "
"sparse inference",
xnn_node_type_to_string(node->type),
node->params.convolution_2d.group_input_channels);
return 0;
}
return XNN_LAYOUT_FLAG_COMPATIBLE_NHWC2NCHW;
}
return 0;
case xnn_node_type_depthwise_convolution_2d:
// Supported cases:
// - 3x3 stride-1 convolution (no dilation, padding 1 on each side)
// - 3x3 stride-2 convolution (no dilation, padding 1 on each side)
// - 5x5 stride-1 convolution (no dilation, padding 2 on each side)
// - 5x5 stride-2 convolution (no dilation, padding 2 on each side)
if ((node->params.depthwise_convolution_2d.dilation_height |
node->params.depthwise_convolution_2d.dilation_width) != 1) {
xnn_log_info("Node %s dilation (height=%" PRIu32 ", width=%" PRIu32
") is incompatible with sparse inference",
xnn_node_type_to_string(node->type),
node->params.convolution_2d.dilation_height,
node->params.convolution_2d.dilation_width);
return 0;
}
if (node->flags & XNN_FLAG_TENSORFLOW_SAME_PADDING) {
xnn_log_info("Node %s flags (%" PRIu32
") has padding incompatible with sparse inference",
xnn_node_type_to_string(node->type), node->flags);
return 0;
}
if (node->params.depthwise_convolution_2d.depth_multiplier != 1) {
xnn_log_info("Node %s depth_multiplier (%" PRIu32
") is incompatible with sparse inference",
xnn_node_type_to_string(node->type),
node->params.depthwise_convolution_2d.depth_multiplier);
return 0;
}
if (node->params.depthwise_convolution_2d.subsampling_height !=
node->params.depthwise_convolution_2d.subsampling_width) {
xnn_log_info("Node %s subsampling (height=%" PRIu32 ", width=%" PRIu32
") is incompatible with sparse inference",
xnn_node_type_to_string(node->type),
node->params.depthwise_convolution_2d.subsampling_height,
node->params.depthwise_convolution_2d.subsampling_width);
return 0;
}
switch (node->params.depthwise_convolution_2d.subsampling_height) {
case 1:
case 2:
break;
default:
xnn_log_info(
"Node %s subsampling_height (%" PRIu32
") is incompatible with sparse inference",
xnn_node_type_to_string(node->type),
node->params.depthwise_convolution_2d.subsampling_height);
return 0;
}
if (node->params.depthwise_convolution_2d.kernel_height !=
node->params.depthwise_convolution_2d.kernel_width) {
xnn_log_info("Node %s kernel (height=%" PRIu32 ", width=%" PRIu32
") is incompatible with sparse inference",
xnn_node_type_to_string(node->type),
node->params.depthwise_convolution_2d.kernel_height,
node->params.depthwise_convolution_2d.kernel_width);
return 0;
}
switch (node->params.depthwise_convolution_2d.kernel_height) {
case 3:
if (node->params.depthwise_convolution_2d.input_padding_top == 1 &&
node->params.depthwise_convolution_2d.input_padding_right == 1 &&
node->params.depthwise_convolution_2d.input_padding_bottom == 1 &&
node->params.depthwise_convolution_2d.input_padding_left == 1) {
return XNN_LAYOUT_FLAG_COMPATIBLE_NCHW;
} else {
xnn_log_info(
"Node %s (3x3 kernel) padding (top=%" PRIu32 ", right=%" PRIu32
", bottom=%" PRIu32 ", left=%" PRIu32
") is incompatible with sparse inference",
xnn_node_type_to_string(node->type),
node->params.depthwise_convolution_2d.input_padding_top,
node->params.depthwise_convolution_2d.input_padding_right,
node->params.depthwise_convolution_2d.input_padding_bottom,
node->params.depthwise_convolution_2d.input_padding_left);
return 0;
}
case 5:
if (node->params.depthwise_convolution_2d.input_padding_top == 2 &&
node->params.depthwise_convolution_2d.input_padding_right == 2 &&
node->params.depthwise_convolution_2d.input_padding_bottom == 2 &&
node->params.depthwise_convolution_2d.input_padding_left == 2) {
return XNN_LAYOUT_FLAG_COMPATIBLE_NCHW;
} else {
xnn_log_info(
"Node %s (5x5 kernel) padding (top=%" PRIu32 ", right=%" PRIu32
", bottom=%" PRIu32 ", left=%" PRIu32
") is incompatible with sparse inference",
xnn_node_type_to_string(node->type),
node->params.depthwise_convolution_2d.input_padding_top,
node->params.depthwise_convolution_2d.input_padding_right,
node->params.depthwise_convolution_2d.input_padding_bottom,
node->params.depthwise_convolution_2d.input_padding_left);
return 0;
}
default:
return 0;
}
case xnn_node_type_depth_to_space_2d:
return XNN_LAYOUT_FLAG_COMPATIBLE_NCHW2NHWC;
case xnn_node_type_global_average_pooling_2d:
return XNN_LAYOUT_FLAG_COMPATIBLE_NCHW |
XNN_LAYOUT_FLAG_COMPATIBLE_NCHW2NHWC;
case xnn_node_type_binary_elementwise:
if (node->binary_operator != xnn_binary_add &&
node->binary_operator != xnn_binary_multiply) {
// TODO(unassigned): We can probably handle any binary operator here?
return false;
}
assert(node->num_inputs == 2);
assert(node->num_outputs == 1);
if (subgraph->values[node->inputs[0]].shape.num_dims != 4 ||
subgraph->values[node->inputs[1]].shape.num_dims != 4) {
xnn_log_info(
"Node %s inputs shape is incompatible with sparse inference",
xnn_node_type_to_string(node->type));
return 0;
}
if (subgraph->values[node->inputs[0]].data != NULL) {
// Check that the first input is representable as either a scalar, or a
// vector
size_t num_nonunit_dims = 0;
for (uint32_t i = 0;
i < subgraph->values[node->inputs[0]].shape.num_dims; i++) {
if (subgraph->values[node->inputs[0]].shape.dim[i] != 1) {
num_nonunit_dims += 1;
}
}
if (num_nonunit_dims > 1) {
return 0;
}
}
if (subgraph->values[node->inputs[1]].data != NULL) {
// Check that the second input is representable as either a scalar, or a
// vector
size_t num_nonunit_dims = 0;
for (uint32_t i = 0;
i < subgraph->values[node->inputs[0]].shape.num_dims; i++) {
if (subgraph->values[node->inputs[0]].shape.dim[i] != 1) {
num_nonunit_dims += 1;
}
}
if (num_nonunit_dims > 1) {
return 0;
}
}
return XNN_LAYOUT_FLAG_COMPATIBLE_NCHW;
case xnn_node_type_static_resize_bilinear_2d:
if (subgraph->values[node->inputs[0]].shape.dim[1] > 1 &&
subgraph->values[node->inputs[0]].shape.dim[2] > 1) {
return XNN_LAYOUT_FLAG_COMPATIBLE_NCHW;
} else {
xnn_log_info(
"Node %s inputs shape is incompatible with sparse inference",
xnn_node_type_to_string(node->type));
return 0;
}
case xnn_node_type_unary_elementwise:
assert(node->num_inputs == 1);
assert(node->num_outputs == 1);
if (subgraph->values[node->inputs[0]].shape.num_dims == 4) {
return XNN_LAYOUT_FLAG_COMPATIBLE_NCHW;
} else {
xnn_log_info(
"Node %s inputs shape is incompatible with sparse inference",
xnn_node_type_to_string(node->type));
return 0;
}
case xnn_node_type_static_mean_squared:
case xnn_node_type_static_mean:
case xnn_node_type_static_sum_squared:
case xnn_node_type_static_sum:
if (subgraph->values[node->inputs[0]].shape.num_dims == 4) {
return XNN_LAYOUT_FLAG_COMPATIBLE_NCHW |
XNN_LAYOUT_FLAG_COMPATIBLE_NCHW2NHWC;
} else {
xnn_log_info(
"Node %s inputs shape is incompatible with sparse inference",
xnn_node_type_to_string(node->type));
return 0;
}
default:
return false;
}
}
void xnn_subgraph_rewrite_for_nchw(xnn_subgraph_t subgraph) {
// Convert parts of the subgraph to NCHW for sparse inference
// Step 1: detect NCHW-compatible Nodes
// Step 2: detect NCHW-compatible clusters (run connected components graph
// algorithm) Step 3: check that all NCHW-compatible Values are consumed only
// by NCHW-compatible Nodes Step 4: switch Values' layout to NCHW
for (uint32_t n = 0; n < subgraph->num_nodes; n++) {
struct xnn_node* node = &subgraph->nodes[n];
node->layout_flags = xnn_check_nchw_compatibility(subgraph, node);
xnn_log_debug(
"Node #%" PRIu32 ": %s (NCHW: %s, NHWC->NCHW: %s, NCHW->NHWC: %s)", n,
xnn_node_type_to_string(node->type),
node->layout_flags & XNN_LAYOUT_FLAG_COMPATIBLE_NCHW ? "yes" : "no",
node->layout_flags & XNN_LAYOUT_FLAG_COMPATIBLE_NHWC2NCHW ? "yes"
: "no",
node->layout_flags & XNN_LAYOUT_FLAG_COMPATIBLE_NCHW2NHWC ? "yes"
: "no");
}
// Run Shiloach-Vishkin connected components algorithm i.e. find all
// XNN_LAYOUT_FLAG_COMPATIBLE_NCHW2NHWC nodes and set them as cluster leaders
// to all the producer nodes
bool update = false;
for (uint32_t n = 0; n < subgraph->num_nodes; n++) {
struct xnn_node* node = &subgraph->nodes[n];
node->cluster_leader = n;
if (node->layout_flags & XNN_LAYOUT_FLAG_COMPATIBLE_NCHW2NHWC) {
for (uint32_t i = 0; i < node->num_inputs; i++) {
const uint32_t input_id = node->inputs[i];
assert(input_id != XNN_INVALID_VALUE_ID);
const struct xnn_value* value = &subgraph->values[input_id];
if (value->data != NULL) {
// Static data, skip this input value. Compatibility of this static
// input with NCHW layout was validated during the initial NCHW
// compatibility check for the Node.
continue;
}
if (xnn_value_is_external(value->flags)) {
// External value, invalid cluster
node->layout_flags |= XNN_LAYOUT_FLAG_INCOMPATIBLE_CLUSTER;
continue;
}
const uint32_t producer_id = value->producer;
assert(producer_id != XNN_INVALID_NODE_ID);
assert(producer_id < n);
struct xnn_node* producer_node = &subgraph->nodes[producer_id];
if ((producer_node->layout_flags &
(XNN_LAYOUT_FLAG_COMPATIBLE_NHWC2NCHW |
XNN_LAYOUT_FLAG_COMPATIBLE_NCHW)) != 0 &&
(producer_node->layout_flags &
XNN_LAYOUT_FLAG_INCOMPATIBLE_CLUSTER) == 0) {
producer_node->layout_flags &= ~XNN_LAYOUT_FLAG_COMPATIBLE_NCHW2NHWC;
if (producer_node->cluster_leader != node->cluster_leader) {
producer_node->cluster_leader = node->cluster_leader = math_max_u32(
producer_node->cluster_leader, node->cluster_leader);
update = true;
}
} else {
node->layout_flags |= XNN_LAYOUT_FLAG_INCOMPATIBLE_CLUSTER;
}
}
}
}
// No NCHW2NHWC compatible nodes have been found thus the graph rewriting
// practically cannot happen.
if (!update) {
return;
}
// Propagate the cluster leader to other nodes in the graph until all the
// nodes in the cluster is not updated
while (update) {
update = false;
for (uint32_t n = 0; n < subgraph->num_nodes; n++) {
struct xnn_node* node = &subgraph->nodes[n];
if (node->layout_flags & XNN_LAYOUT_FLAG_INCOMPATIBLE_CLUSTER) {
continue;
}
if ((node->layout_flags & (XNN_LAYOUT_FLAG_COMPATIBLE_NCHW |
XNN_LAYOUT_FLAG_COMPATIBLE_NCHW2NHWC)) == 0) {
continue;
}
for (uint32_t i = 0; i < node->num_inputs; i++) {
const uint32_t input_id = node->inputs[i];
assert(input_id != XNN_INVALID_VALUE_ID);
const struct xnn_value* value = &subgraph->values[input_id];
if (value->data != NULL) {
// Static data, skip this input value. Compatibility of this static
// input with NCHW layout was validated during the initial NCHW
// compatibility check for the Node.
continue;
}
if (xnn_value_is_external(value->flags)) {
// External value, invalid cluster
node->layout_flags |= XNN_LAYOUT_FLAG_INCOMPATIBLE_CLUSTER;
continue;
}
const uint32_t producer_id = value->producer;
assert(producer_id != XNN_INVALID_NODE_ID);
assert(producer_id < n);
struct xnn_node* producer_node = &subgraph->nodes[producer_id];
if ((producer_node->layout_flags &
(XNN_LAYOUT_FLAG_COMPATIBLE_NHWC2NCHW |
XNN_LAYOUT_FLAG_COMPATIBLE_NCHW)) != 0 &&
(producer_node->layout_flags &
XNN_LAYOUT_FLAG_INCOMPATIBLE_CLUSTER) == 0) {
producer_node->layout_flags &= ~XNN_LAYOUT_FLAG_COMPATIBLE_NCHW2NHWC;
if (producer_node->cluster_leader != node->cluster_leader) {
producer_node->cluster_leader = node->cluster_leader = math_max_u32(
producer_node->cluster_leader, node->cluster_leader);
update = true;
}
} else {
node->layout_flags |= XNN_LAYOUT_FLAG_INCOMPATIBLE_CLUSTER;
}
}
}
}
// Propagate XNN_LAYOUT_FLAG_INCOMPATIBLE_CLUSTER flags up to the cluster
// leaders
for (uint32_t n = 0; n < subgraph->num_nodes; n++) {
struct xnn_node* node = &subgraph->nodes[n];
subgraph->nodes[node->cluster_leader].layout_flags |=
node->layout_flags & XNN_LAYOUT_FLAG_INCOMPATIBLE_CLUSTER;
}
// Check that all Values consumed by NCHW-compatible cluster don't have
// NCHW-incompatible consumers
for (uint32_t n = 0; n < subgraph->num_nodes; n++) {
struct xnn_node* node = &subgraph->nodes[n];
if ((subgraph->nodes[node->cluster_leader].layout_flags &
XNN_LAYOUT_FLAG_INCOMPATIBLE_CLUSTER) != 0) {
continue;
}
if ((node->layout_flags & (XNN_LAYOUT_FLAG_COMPATIBLE_NCHW2NHWC |
XNN_LAYOUT_FLAG_COMPATIBLE_NCHW)) == 0) {
continue;
}
for (uint32_t i = 0; i < node->num_inputs; i++) {
const uint32_t input_id = node->inputs[i];
assert(input_id != XNN_INVALID_VALUE_ID);
struct xnn_value* value = &subgraph->values[input_id];
if (value->data != NULL) {
// Static data, skip this input value because it doesn't have a producer
// Node.
continue;
}
assert(!xnn_value_is_external(value->flags));
value->num_nchw_compatible_consumers += 1;
}
}
for (uint32_t n = 0; n < subgraph->num_nodes; n++) {
struct xnn_node* node = &subgraph->nodes[n];
if ((subgraph->nodes[node->cluster_leader].layout_flags &
XNN_LAYOUT_FLAG_INCOMPATIBLE_CLUSTER) != 0) {
continue;
}
if ((node->layout_flags & (XNN_LAYOUT_FLAG_COMPATIBLE_NCHW2NHWC |
XNN_LAYOUT_FLAG_COMPATIBLE_NCHW)) == 0) {
continue;
}
for (uint32_t i = 0; i < node->num_inputs; i++) {
const uint32_t input_id = node->inputs[i];
assert(input_id != XNN_INVALID_VALUE_ID);
const struct xnn_value* value = &subgraph->values[input_id];
if (value->data != NULL) {
// Static data, skip this input value because it doesn't have a producer
// Node.
continue;
}
assert(!xnn_value_is_external(value->flags));
assert(value->num_nchw_compatible_consumers > 0);
if (value->num_nchw_compatible_consumers != value->num_consumers) {
subgraph->nodes[node->cluster_leader].layout_flags |=
XNN_LAYOUT_FLAG_INCOMPATIBLE_CLUSTER;
}
}
}
// Evaluate if it is profitable to run the model as sparse:
// - Compute the number of parameters and zeroes in 1x1 Convolution weights
// - Disable sparse rewriting for clusters without 1x1 Convolutions
// (num_params == 0)
// or with less than 2/3rd of zeroes in 1x1 Convolution filters
for (uint32_t n = 0; n < subgraph->num_nodes; n++) {
struct xnn_node* node = &subgraph->nodes[n];
if ((subgraph->nodes[node->cluster_leader].layout_flags &
XNN_LAYOUT_FLAG_INCOMPATIBLE_CLUSTER) != 0) {
continue;
}
if ((node->type == xnn_node_type_convolution_2d &&
max(node->params.convolution_2d.kernel_height,
node->params.convolution_2d.kernel_width) == 1) ||
node->type == xnn_node_type_fully_connected) {
assert(node->num_inputs >= 2);
const struct xnn_value* filter = &subgraph->values[node->inputs[1]];
assert(filter->data != NULL);
const size_t num_params =
filter->shape.dim[0] * filter->shape.dim[filter->shape.num_dims - 1];
subgraph->nodes[node->cluster_leader].num_params += num_params;
size_t num_zeroes = 0;
switch (filter->datatype) {
case xnn_datatype_fp32: {
const float* data = (const float*)filter->data;
for (size_t i = 0; i < num_params; i++) {
num_zeroes += (size_t)(data[i] == 0.0f);
}
break;
}
case xnn_datatype_fp16: {
const xnn_float16* data = (const xnn_float16*)filter->data;
for (size_t i = 0; i < num_params; i++) {
num_zeroes += (size_t)(xnn_float16_is_zero(data[i]));
}
break;
}
default:
XNN_UNREACHABLE;
}
xnn_log_debug("1x1 Convolution 2D Node #%" PRIu32 ": %zu / %zu sparsity",
n, num_zeroes, num_params);
subgraph->nodes[node->cluster_leader].num_zeroes += num_zeroes;
}
}
bool use_nchw_layout = false;
for (uint32_t n = 0; n < subgraph->num_nodes; n++) {
struct xnn_node* node = &subgraph->nodes[n];
if ((subgraph->nodes[node->cluster_leader].layout_flags &
XNN_LAYOUT_FLAG_INCOMPATIBLE_CLUSTER) != 0) {
continue;
}
if ((node->layout_flags & (XNN_LAYOUT_FLAG_COMPATIBLE_NCHW2NHWC |
XNN_LAYOUT_FLAG_COMPATIBLE_NCHW)) == 0) {
continue;
}
if (subgraph->nodes[node->cluster_leader].num_zeroes * 3 <=
subgraph->nodes[node->cluster_leader].num_params * 2) {
xnn_log_info("Node #%" PRIu32
": sparse inference disabled: 1x1 Convolutions contain %zu "
"/ %zu zero weights",
n, subgraph->nodes[node->cluster_leader].num_zeroes,
subgraph->nodes[node->cluster_leader].num_params);
continue;
}
for (uint32_t i = 0; i < node->num_inputs; i++) {
const uint32_t input_id = node->inputs[i];
assert(input_id != XNN_INVALID_VALUE_ID);
struct xnn_value* value = &subgraph->values[input_id];
if (value->data != NULL) {
// Static data, skip this input value because it doesn't have a producer
// Node.
continue;
}
assert(!xnn_value_is_external(value->flags));
assert(value->num_nchw_compatible_consumers > 0);
assert(value->num_nchw_compatible_consumers == value->num_consumers);
if (value->layout != xnn_layout_type_nchw) {
value->layout = xnn_layout_type_nchw;
xnn_log_info("set Value #%" PRIu32 " layout to NCHW", node->inputs[i]);
use_nchw_layout = true;
}
}
}
if (use_nchw_layout) {
xnn_log_info("XNNPACK has switched to sparse inference mode!");
}
}
static bool any_values_fp32(xnn_subgraph_t subgraph,
const struct xnn_node* node) {
for (uint32_t i = 0; i < node->num_inputs; i++) {
const uint32_t input_id = node->inputs[i];
assert(input_id != XNN_INVALID_VALUE_ID);
if (subgraph->values[input_id].datatype == xnn_datatype_fp32) {
return true;
}
}
for (uint32_t i = 0; i < node->num_outputs; i++) {
const uint32_t output_id = node->outputs[i];
assert(output_id != XNN_INVALID_VALUE_ID);
if (subgraph->values[output_id].datatype == xnn_datatype_fp32) {
return true;
}
}
return false;
}
static bool all_values_fp32_or_pfp32(xnn_subgraph_t subgraph,
const struct xnn_node* node) {
for (uint32_t i = 0; i < node->num_inputs; i++) {
const uint32_t input_id = node->inputs[i];
assert(input_id != XNN_INVALID_VALUE_ID);
if (subgraph->values[input_id].datatype != xnn_datatype_fp32 &&
subgraph->values[input_id].datatype != xnn_datatype_pfp32) {
return false;
}
}
for (uint32_t i = 0; i < node->num_outputs; i++) {
const uint32_t output_id = node->outputs[i];
assert(output_id != XNN_INVALID_VALUE_ID);
if (subgraph->values[output_id].datatype != xnn_datatype_fp32 &&
subgraph->values[output_id].datatype != xnn_datatype_pfp32) {
return false;
}
}
return true;
}
bool xnn_subgraph_rewrite_for_fp16(xnn_subgraph_t subgraph) {
xnn_log_info("Analyzing subgraph for FP16 compatibility");
// Count the number of consumers for each value.
xnn_subgraph_analyze_consumers_and_producers(subgraph);
// Convert tensors and operators in the subgraph to FP16
// 1. Check that all operators in the subgraph are supported in FP16.
// 2. Indicate values that must be converted to FP16.
// 3. Replace FP32 Values with FP16 Values as Nodes' inputs/outputs.
// 4. Insert FP32->FP16 Convert Nodes for external FP32 inputs and FP16->FP32
// Convert Nodes for external outputs.
const uint32_t num_original_values = subgraph->num_values;
// Check that all operators in the subgraph are supported in FP16, bail out on
// any unsupported one.
for (uint32_t n = 0; n < subgraph->num_nodes; n++) {
struct xnn_node* node = &subgraph->nodes[n];
if (node->type == xnn_node_type_invalid) {
// Node was fused away, skip.
continue;
}
if (!any_values_fp32(subgraph, node)) {
xnn_log_warning("FP16 rewrite aborted: node #%" PRIu32
" (%s) is not FP32",
n, xnn_node_type_to_string(node->type));
return false;
}
switch (node->type) {
case xnn_node_type_average_pooling_2d:
case xnn_node_type_batch_matrix_multiply:
case xnn_node_type_binary_elementwise:
case xnn_node_type_concatenate:
case xnn_node_type_convert:
case xnn_node_type_convolution_2d:
case xnn_node_type_copy:
case xnn_node_type_deconvolution_2d:
case xnn_node_type_depth_to_space_2d:
case xnn_node_type_depthwise_convolution_2d:
case xnn_node_type_even_split:
case xnn_node_type_fully_connected:
case xnn_node_type_global_average_pooling_2d:
case xnn_node_type_global_sum_pooling_2d:
case xnn_node_type_max_pooling_2d:
case xnn_node_type_rope:
case xnn_node_type_softmax:
case xnn_node_type_space_to_depth_2d:
case xnn_node_type_static_constant_pad:
case xnn_node_type_static_mean:
case xnn_node_type_static_mean_squared:
case xnn_node_type_static_reduce_max:
case xnn_node_type_static_reduce_min:
case xnn_node_type_static_reshape:
case xnn_node_type_static_resize_bilinear_2d:
case xnn_node_type_static_slice:
case xnn_node_type_static_sum:
case xnn_node_type_static_sum_squared:
case xnn_node_type_static_transpose:
case xnn_node_type_unary_elementwise:
break;
case xnn_node_type_pack_lh:
if (xnn_init_x16_pack_lh_config() != NULL) {
break;
}
XNN_FALLTHROUGH
default:
xnn_log_warning("FP16 rewrite aborted: node #%" PRIu32
" (%s) is not supported for FP16 inference",
n, xnn_node_type_to_string(node->type));
return false;
}
}
// Annotate Values to be converted to FP16 as FP16-compatible.
// Note that static weights in [Depthwise] Convolution, Fully Connected Nodes
// remain FP32, they will be converted to FP16 during weight repacking when
// the operator is created.
for (uint32_t n = 0; n < subgraph->num_nodes; n++) {
struct xnn_node* node = &subgraph->nodes[n];
switch (node->type) {
case xnn_node_type_deconvolution_2d:
case xnn_node_type_depthwise_convolution_2d:
if (subgraph->values[node->inputs[0]].datatype == xnn_datatype_fp32) {
subgraph->values[node->inputs[0]].fp16_rewrite.fp16_compatible = true;
}
subgraph->values[node->outputs[0]].fp16_rewrite.fp16_compatible = true;
break;
case xnn_node_type_convolution_2d:
if (subgraph->values[node->inputs[0]].datatype == xnn_datatype_qdint8) {
subgraph->values[node->outputs[0]].fp16_rewrite.fp16_compatible = true;
} else {
subgraph->values[node->inputs[0]].fp16_rewrite.fp16_compatible = true;
subgraph->values[node->outputs[0]].fp16_rewrite.fp16_compatible = true;
}
break;
case xnn_node_type_fully_connected:
if (subgraph->values[node->inputs[0]].datatype == xnn_datatype_qdint8 ||
subgraph->values[node->inputs[0]].datatype ==
xnn_datatype_qduint8 ||
subgraph->values[node->inputs[0]].datatype == xnn_datatype_qpint8) {
subgraph->values[node->outputs[0]].fp16_rewrite.fp16_compatible = true;
} else if (subgraph->values[node->inputs[0]].datatype ==
xnn_datatype_fp32 &&
(node->packed_input_datatype == xnn_datatype_qdint8 ||
node->packed_input_datatype == xnn_datatype_qduint8 ||
node->packed_input_datatype == xnn_datatype_qpint8)) {
subgraph->values[node->inputs[0]].fp16_rewrite.fp16_compatible = true;
subgraph->values[node->outputs[0]].fp16_rewrite.fp16_compatible = true;
} else if ((subgraph->values[node->inputs[0]].datatype ==
xnn_datatype_fp32 ||
subgraph->values[node->inputs[0]].datatype ==
xnn_datatype_pfp32) &&
(subgraph->values[node->inputs[1]].datatype ==
xnn_datatype_fp16 ||
subgraph->values[node->inputs[1]].datatype ==
xnn_datatype_fp32) &&
subgraph->values[node->outputs[0]].datatype ==
xnn_datatype_fp32) {
subgraph->values[node->inputs[0]].fp16_rewrite.fp16_compatible = true;
subgraph->values[node->outputs[0]].fp16_rewrite.fp16_compatible = true;
if (subgraph->values[node->inputs[1]].datatype == xnn_datatype_fp32) {
subgraph->values[node->inputs[1]].fp16_rewrite.fp16_compatible = true;
}
if (node->num_inputs > 2) {
assert(node->inputs[2] != XNN_INVALID_VALUE_ID);
if (subgraph->values[node->inputs[2]].datatype == xnn_datatype_fp32) {
subgraph->values[node->inputs[2]].fp16_rewrite.fp16_compatible = true;
}
}
} else if (all_values_fp32_or_pfp32(subgraph, node)) {
subgraph->values[node->inputs[0]].fp16_rewrite.fp16_compatible = true;
subgraph->values[node->outputs[0]].fp16_rewrite.fp16_compatible = true;
} else {
xnn_log_warning(
"FP16 rewrite aborted: node #%" PRIu32
" (%s). Invalid compute type (input=%s, weights=%s, output=%s)",
n, xnn_node_type_to_string(node->type),
xnn_datatype_to_string(
subgraph->values[node->inputs[0]].datatype),
xnn_datatype_to_string(
subgraph->values[node->inputs[1]].datatype),
xnn_datatype_to_string(
subgraph->values[node->outputs[0]].datatype));
return false;
}
break;
case xnn_node_type_convert:
if (subgraph->values[node->inputs[0]].datatype == xnn_datatype_fp32) {
subgraph->values[node->inputs[0]].fp16_rewrite.fp16_compatible = true;
}
if (subgraph->values[node->outputs[0]].datatype == xnn_datatype_fp32) {
subgraph->values[node->outputs[0]].fp16_rewrite.fp16_compatible = true;
}
break;
case xnn_node_type_pack_lh:
if (subgraph->values[node->inputs[0]].datatype == xnn_datatype_fp32) {
subgraph->values[node->inputs[0]].fp16_rewrite.fp16_compatible = true;
}
break;
default:
for (uint32_t i = 0; i < node->num_inputs; i++) {
const uint32_t input_id = node->inputs[i];
assert(input_id != XNN_INVALID_VALUE_ID);
switch (subgraph->values[input_id].datatype) {
case xnn_datatype_fp32:
case xnn_datatype_pfp32:
subgraph->values[input_id].fp16_rewrite.fp16_compatible = true;
break;
default:
break;
}
}
for (uint32_t o = 0; o < node->num_outputs; o++) {
switch (subgraph->values[node->outputs[o]].datatype) {
case xnn_datatype_fp32:
case xnn_datatype_pfp32:
subgraph->values[node->outputs[o]].fp16_rewrite.fp16_compatible = true;
break;
default:
break;
}
}
break;
}
}
// Attempt to allocate memory for static values and external input/outputs.
// The FP16 rewrite is cleanly aborted on failure.
for (uint32_t n = 0; n < num_original_values; n++) {
struct xnn_value* value = &subgraph->values[n];
value->fp16_rewrite.fp16_id = XNN_INVALID_VALUE_ID;
value->fp16_rewrite.fp32_id = XNN_INVALID_VALUE_ID;
if (value->fp16_rewrite.fp16_compatible) {
assert(value->datatype == xnn_datatype_fp32 ||
value->datatype == xnn_datatype_pfp32);
if (xnn_value_is_static(value->allocation_type)) {
assert(value->producer == XNN_INVALID_NODE_ID);
const size_t fp16_size =
xnn_tensor_get_size(value) / 2 + XNN_EXTRA_BYTES;
value->fp16_rewrite.fp16_temp_data = xnn_allocate_zero_memory(fp16_size);
if (value->fp16_rewrite.fp16_temp_data == NULL) {
xnn_log_error("failed to allocate %zu bytes for fp16 tensor data",
(size_t)fp16_size);
goto error;
}
} else if (xnn_value_is_external(value->flags)) {
struct xnn_value* fp16_value =
xnn_subgraph_new_internal_value(subgraph);
if (fp16_value == NULL) {
xnn_log_error(
"FP16 rewrite aborted: failed to allocate value for external "
"input/output");
goto error;
} else {
// Recompute value due to potential reallocation in
// xnn_subgraph_new_internal_value
value = &subgraph->values[n];
xnn_value_copy(fp16_value, value);
switch (value->datatype) {
case xnn_datatype_fp32:
fp16_value->datatype = xnn_datatype_fp16;
break;
case xnn_datatype_pfp32:
fp16_value->datatype = xnn_datatype_pfp16;
break;
default:
XNN_UNREACHABLE;
}
// Clear external input/output flags
fp16_value->flags = 0;
fp16_value->producer = XNN_INVALID_NODE_ID;
fp16_value->first_consumer = XNN_INVALID_NODE_ID;
fp16_value->num_consumers = 0;
fp16_value->fp16_rewrite.fp16_id = XNN_INVALID_VALUE_ID;
fp16_value->fp16_rewrite.fp32_id = value->id;
fp16_value->allocation_type = xnn_allocation_type_workspace;
value->fp16_rewrite.fp16_id = fp16_value->id;
}
} else if (xnn_value_is_internal(value)) {
// fp16 tensors only need half the memory of fp32 tensors.
value->size /= 2;
}
}
}
// Count the number of external inputs and outputs which require Convert nodes
uint32_t num_external_inputs = 0;
uint32_t num_external_outputs = 0;
for (uint32_t n = 0; n < subgraph->num_nodes; n++) {
const struct xnn_node* node = &subgraph->nodes[n];
for (uint32_t i = 0; i < node->num_inputs; i++) {
if (is_repeated_input(node, i)) {
continue;
}
const uint32_t input_id = node->inputs[i];
assert(input_id != XNN_INVALID_VALUE_ID);
const struct xnn_value* value = &subgraph->values[input_id];
if (value->fp16_rewrite.fp16_id != XNN_INVALID_VALUE_ID &&
value->first_consumer == n) {
assert(value->data == NULL);
assert(value->datatype == xnn_datatype_fp32);
assert(subgraph->values[value->fp16_rewrite.fp16_id].datatype == xnn_datatype_fp16);
// This value isn't always an external input, it could be an external
// output of the current subgraph (due to partition), and be
// simultaneously consumed by the current node.
if (xnn_value_is_external_input(value->flags)) {
num_external_inputs += 1;
}
}
}
for (uint32_t o = 0; o < node->num_outputs; o++) {
const struct xnn_value* value = &subgraph->values[node->outputs[o]];
if (value->fp16_rewrite.fp16_id != XNN_INVALID_VALUE_ID) {
assert(value->datatype == xnn_datatype_fp32);
assert(subgraph->values[value->fp16_rewrite.fp16_id].datatype == xnn_datatype_fp16);
assert(xnn_value_is_external_output(value->flags));
num_external_outputs += 1;
}
}
}
xnn_log_debug("Discovered %" PRIu32 " external inputs and %" PRIu32
" external outputs",
num_external_inputs, num_external_outputs);
// Attempt to allocate memory for the Convert nodes.
const uint32_t num_original_nodes = subgraph->num_nodes;
if (xnn_subgraph_add_nodes(subgraph,
num_external_inputs + num_external_outputs) !=
xnn_status_success) {
xnn_log_error(
"FP16 rewrite aborted: failed to allocate node for external "
"input/output");
goto error;
}
// From this point the subgraph and tensor data get mutated, clean failure is
// no longer an option.
// Replace FP32 Values in Nodes' inputs/outputs with FP16 Values.
// - FP32 values of static tensors get converted in a new data buffer.
// - For external inputs and outputs we create same-shaped FP16 Values and use
// those instead.
// - Values that are neither static nor external are converted to FP16
// in-place
for (uint32_t n = 0; n < num_original_values; n++) {
struct xnn_value* value = &subgraph->values[n];
if (value->fp16_rewrite.fp16_compatible) {
if (xnn_value_is_static(value->allocation_type)) {
assert(value->datatype == xnn_datatype_fp32);
const size_t num_elements = xnn_shape_multiply_all_dims(&value->shape);
xnn_run_unary_elementwise_nc(
xnn_unary_convert, xnn_datatype_fp32, xnn_datatype_fp16,
/*params=*/NULL, /*input_quantization=*/NULL,
/*output_quantization=*/NULL, 0, num_elements, 1, 1, 1, NULL,
value->data, value->fp16_rewrite.fp16_temp_data);
// Remember pointer to the original fp32 data, nodes like convolution
// need fp32 weights/biases.
value->fp32_data = value->data;
value->data = value->fp16_rewrite.fp16_temp_data;
value->fp16_rewrite.fp16_temp_data = NULL;
value->datatype = xnn_datatype_fp16;
xnn_log_debug("FP16 rewrite: converted static FP32 tensor #%" PRIu32
" to FP16 in new buffer",
n);
} else if (xnn_value_is_external(value->flags)) {
assert(value->datatype == xnn_datatype_fp32);
assert(value->fp16_rewrite.fp16_id != XNN_INVALID_VALUE_ID);
value->producer = XNN_INVALID_NODE_ID;
value->num_consumers = 0;
xnn_log_debug("FP16 rewrite: created FP16 tensor #%" PRIu32
" for external FP32 tensor #%" PRIu32,
subgraph->values[value->fp16_rewrite.fp16_id].id, n);
} else {
switch (value->datatype) {
case xnn_datatype_fp32:
xnn_log_debug(
"FP16 rewrite: converted FP32 tensor #%" PRIu32 " to FP16", n);
value->datatype = xnn_datatype_fp16;
break;
case xnn_datatype_pfp32:
xnn_log_debug("FP16 rewrite: converted PFP32 tensor #%" PRIu32
" to PFP16",
n);
value->datatype = xnn_datatype_pfp16;
break;
default:
XNN_UNREACHABLE;
}
}
}
}
// Switch the nodes consuming/generated converted `fp32` inputs/outputs to
// their `fp16` values.
for (uint32_t n = 0; n < subgraph->num_nodes; n++) {
struct xnn_node* node = &subgraph->nodes[n];
if (node->type == xnn_node_type_invalid) {
// Node was fused away, skip.
continue;
}
// Fix up anything node-type specific.
switch (node->type) {
case xnn_node_type_static_constant_pad:
node->params.static_pad.padding_value = fp16_ieee_from_fp32_value(
uint32_as_float(node->params.static_pad.padding_value));
break;
case xnn_node_type_batch_matrix_multiply:
case xnn_node_type_fully_connected: {
// Patch up any LHS packing of fully-connected nodes, if needed.
if (node->flags & XNN_FLAG_INLINE_LHS_PACKING) {
switch (node->packed_input_datatype) {
case xnn_datatype_pfp32:
// Switch from packed `fp32` to packed `fp16`.
node->packed_input_datatype = xnn_datatype_pfp16;
break;
case xnn_datatype_qpint8:
// Convert from `qpint8` back to `qdint8` since we don't have a
// `qpint8` packing function for `f16` inputs.
node->packed_input_datatype = xnn_datatype_qdint8;
break;
default:
break;
}
} else if (subgraph->values[node->inputs[0]].datatype ==
xnn_datatype_qpint8) {
subgraph->values[node->inputs[0]].datatype = xnn_datatype_qdint8;
}
} break;
default:
break;
}
for (uint32_t i = 0; i < node->num_inputs; i++) {
const uint32_t input_id = node->inputs[i];
assert(input_id != XNN_INVALID_VALUE_ID);
const uint32_t fp16_id = subgraph->values[input_id].fp16_rewrite.fp16_id;
if (fp16_id != XNN_INVALID_VALUE_ID) {
struct xnn_value* fp16_value = &subgraph->values[fp16_id];
assert(fp16_value->fp16_rewrite.fp32_id == node->inputs[i]);
if (fp16_value->first_consumer == XNN_INVALID_NODE_ID) {
fp16_value->first_consumer = n;
}
node->inputs[i] = fp16_id;
if (!is_repeated_input(node, i)) {
fp16_value->num_consumers++;
}
}
}
for (uint32_t o = 0; o < node->num_outputs; o++) {
const uint32_t fp32_id = node->outputs[o];
const uint32_t fp16_id = subgraph->values[fp32_id].fp16_rewrite.fp16_id;
if (fp16_id != XNN_INVALID_VALUE_ID) {
struct xnn_value* fp16_value = &subgraph->values[fp16_id];
if (fp16_value->first_consumer == XNN_INVALID_NODE_ID &&
fp16_value->producer == XNN_INVALID_NODE_ID) {
assert(fp16_value->fp16_rewrite.fp32_id == fp32_id);
} else {
// Prevent double assignments by creating a new copy of the output
// value if it has already been written to.
fp16_value = xnn_subgraph_new_internal_value(subgraph);
xnn_value_copy(fp16_value, &subgraph->values[fp16_id]);
fp16_value->first_consumer = XNN_INVALID_NODE_ID;
fp16_value->num_consumers = 0;
subgraph->values[fp32_id].fp16_rewrite.fp16_id = fp16_value->id;
}
node->outputs[o] = fp16_value->id;
fp16_value->producer = n;
}
}
}
struct xnn_node* output_node = &subgraph->nodes[subgraph->num_nodes - 1];
for (uint32_t n = num_original_nodes; n != 0; n--) {
const struct xnn_node* node = &subgraph->nodes[n - 1];
// Insert Convert nodes for outputs
for (uint32_t o = 0; o < node->num_outputs; o++) {
const struct xnn_value* value = &subgraph->values[node->outputs[o]];
const uint32_t fp32_id = value->fp16_rewrite.fp32_id;
if (fp32_id != XNN_INVALID_VALUE_ID &&
subgraph->values[fp32_id].fp16_rewrite.fp16_id == value->id) {
xnn_log_debug("Inserted FP16->FP32 Convert Node from tensor #%" PRIu32
" to output tensor #%" PRIu32,
value->id, fp32_id);
const uint32_t output_node_id = output_node->id;
assert(output_node >= subgraph->nodes);
xnn_node_clear(output_node);
output_node->id = output_node_id;
xnn_init_convert_node(output_node, value->id, fp32_id, 0 /* flags */);
output_node -= 1;
}
}
// Move the Node to the new location
if (output_node != node) {
const uint32_t output_node_id = output_node->id;
assert(output_node >= subgraph->nodes);
memcpy(output_node, node, sizeof(struct xnn_node));
output_node->id = output_node_id;
output_node -= 1;
}
// Insert Convert nodes for inputs
for (uint32_t i = 0; i < node->num_inputs; i++) {
if (is_repeated_input(node, i)) {
continue;
}
const struct xnn_value* value = &subgraph->values[node->inputs[i]];
const uint32_t fp32_id = value->fp16_rewrite.fp32_id;
if (fp32_id != XNN_INVALID_VALUE_ID &&
subgraph->values[fp32_id].first_consumer == n - 1) {
// Only insert convert nodes if the value actually is an external input.
// This value could be an external output, if that's the case, we have
// already inserted a convert node in loop above for outputs.
if (xnn_value_is_external_input(subgraph->values[fp32_id].flags)) {
xnn_log_debug("Inserted FP32->FP16 Convert Node from tensor #%" PRIu32
" to tensor #%" PRIu32,
fp32_id, value->id);
const uint32_t output_node_id = output_node->id;
assert(output_node >= subgraph->nodes);
xnn_node_clear(output_node);
output_node->id = output_node_id;
xnn_init_convert_node(output_node, fp32_id, value->id, 0 /* flags */);
output_node -= 1;
}
}
}
}
xnn_log_info("XNNPACK has switched to FP16 inference mode!");
return true;
error:
for (uint32_t n = 0; n < subgraph->num_values; n++) {
struct xnn_value* value = &subgraph->values[n];
// Deallocate extra memory used during static tensor rewrite.
if (value->fp16_rewrite.fp16_temp_data != NULL) {
xnn_release_memory(value->fp16_rewrite.fp16_temp_data);
}
// Revert marking values as FP16-compatible, as xnn_delete_subgraph() may
// assume ownership of those that are.
value->fp16_rewrite.fp16_compatible = false;
}
// Clear the fp16 values created for external inputs and outputs.
for (uint32_t n = num_original_values; n < subgraph->num_values; n++) {
xnn_value_clear(&subgraph->values[n]);
}
return false;
}
static void xnn_node_replace_output(struct xnn_node* node,
uint32_t old_output_id,
uint32_t new_output_id) {
for (size_t i = 0; i < node->num_outputs; i++) {
if (node->outputs[i] == old_output_id) {
node->outputs[i] = new_output_id;
}
}
}
static bool is_clamp(const struct xnn_node* node) {
return node->type == xnn_node_type_unary_elementwise &&
node->unary_operator == xnn_unary_clamp;
}
static bool has_clamp(const struct xnn_node* node) {
if (is_clamp(node)) {
return true;
}
switch (node->type) {
case xnn_node_type_average_pooling_2d:
case xnn_node_type_convolution_2d:
case xnn_node_type_deconvolution_2d:
case xnn_node_type_depthwise_convolution_2d:
case xnn_node_type_fully_connected:
case xnn_node_type_max_pooling_2d:
return true;
default:
return false;
}
}
// Can we reorder the use of a value from the producer to the consumer?
// We can if no nodes between the producer and the consumer use the value.
static bool can_reorder_use(xnn_subgraph_t subgraph, uint32_t value_id,
uint32_t producer_id, uint32_t consumer_id) {
assert(producer_id < consumer_id);
for (uint32_t i = producer_id + 1; i < consumer_id; i++) {
const struct xnn_node* node = &subgraph->nodes[i];
for (uint32_t j = 0; j < node->num_inputs; j++) {
if (node->inputs[j] == value_id) return false;
}
for (uint32_t j = 0; j < node->num_outputs; j++) {
if (node->outputs[j] == value_id) return false;
}
}
return true;
}
enum xnn_status xnn_subgraph_fusion(xnn_subgraph_t subgraph) {
// Fuse Nodes where possible
for (uint32_t i = 0; i < subgraph->num_values; i++) {
struct xnn_value* value = &subgraph->values[i];
if (value->num_consumers == 1) {
const uint32_t producer_id = value->producer;
if (producer_id == XNN_INVALID_NODE_ID) {
continue;
}
assert(producer_id < subgraph->num_nodes);
const uint32_t consumer_id = value->first_consumer;
if (consumer_id == XNN_INVALID_NODE_ID) {
continue;
}
assert(consumer_id < subgraph->num_nodes);
struct xnn_node* producer = &subgraph->nodes[producer_id];
assert(producer->type != xnn_node_type_invalid);
struct xnn_node* consumer = &subgraph->nodes[consumer_id];
if (consumer->type == xnn_node_type_invalid) {
xnn_log_fatal(
"Node %u (produced by %s node %u) has no consumers. Should an "
"external output have been set?",
consumer_id, xnn_node_type_to_string(producer->type), producer_id);
return xnn_status_invalid_state;
}
// Try to fuse Constant Pad node downstream into [Depthwise] Convolution
// 2D Node
if (producer->type == xnn_node_type_static_constant_pad) {
assert(producer->num_inputs == 1);
assert(producer->num_outputs == 1);
const bool is_spatial_2d_padding =
value->shape.num_dims == 4 &&
(producer->params.static_pad.pre_paddings[0] |
producer->params.static_pad.post_paddings[0] |
producer->params.static_pad.pre_paddings[3] |
producer->params.static_pad.post_paddings[3]) == 0;
const enum xnn_datatype padding_datatype =
subgraph->values[producer->outputs[0]].datatype;
const uint32_t padding_value =
producer->params.static_pad.padding_value;
const bool is_zero_padding =
(padding_datatype == xnn_datatype_fp32 && padding_value == 0) ||
((padding_datatype == xnn_datatype_qint8 ||
padding_datatype == xnn_datatype_quint8) &&
padding_value == (uint32_t)subgraph->values[producer->outputs[0]]
.quantization.zero_point);
switch (consumer->type) {
case xnn_node_type_convolution_2d:
if (is_spatial_2d_padding && is_zero_padding &&
!(consumer->flags & XNN_FLAG_TENSORFLOW_SAME_PADDING)) {
xnn_log_info("fuse Constant Pad Node #%" PRIu32
" into Convolution 2D Node #%" PRIu32,
consumer_id, producer_id);
assert(consumer->num_inputs >= 1);
assert(consumer->inputs[0] == producer->outputs[0]);
consumer->params.convolution_2d.input_padding_top +=
producer->params.static_pad.pre_paddings[1];
consumer->params.convolution_2d.input_padding_right +=
producer->params.static_pad.post_paddings[2];
consumer->params.convolution_2d.input_padding_bottom +=
producer->params.static_pad.post_paddings[1];
consumer->params.convolution_2d.input_padding_left +=
producer->params.static_pad.pre_paddings[2];
consumer->inputs[0] = producer->inputs[0];
const uint32_t fused_input_id = producer->inputs[0];
assert(fused_input_id < subgraph->num_values);
if (subgraph->values[fused_input_id].first_consumer ==
producer_id) {
subgraph->values[fused_input_id].first_consumer = consumer_id;
}
xnn_node_clear(producer);
xnn_value_clear(value);
}
break;
case xnn_node_type_depthwise_convolution_2d:
if (is_spatial_2d_padding && is_zero_padding &&
!(consumer->flags & XNN_FLAG_TENSORFLOW_SAME_PADDING)) {
xnn_log_info("fuse Constant Pad Node #%" PRIu32
" into Depthwise Convolution 2D Node #%" PRIu32,
consumer_id, producer_id);
assert(consumer->num_inputs >= 1);
assert(consumer->inputs[0] == producer->outputs[0]);
consumer->params.depthwise_convolution_2d.input_padding_top +=
producer->params.static_pad.pre_paddings[1];
consumer->params.depthwise_convolution_2d.input_padding_right +=
producer->params.static_pad.post_paddings[2];
consumer->params.depthwise_convolution_2d.input_padding_bottom +=
producer->params.static_pad.post_paddings[1];
consumer->params.depthwise_convolution_2d.input_padding_left +=
producer->params.static_pad.pre_paddings[2];
consumer->inputs[0] = producer->inputs[0];
const uint32_t fused_input_id = producer->inputs[0];
assert(fused_input_id < subgraph->num_values);
if (subgraph->values[fused_input_id].first_consumer ==
producer_id) {
subgraph->values[fused_input_id].first_consumer = consumer_id;
}
xnn_node_clear(producer);
xnn_value_clear(value);
}
break;
default:
break;
}
}
// Try to fuse copy upstream. Copy can be fused upstream as long as this
// value is internal. E.g. ---> (N1) --- value ---> (Copy) ---> v1 If
// value is persistent or external, fusing copy upstream into N1 will skip
// the write to value, N1 will write to v1 instead, which is wrong.
if (consumer->type == xnn_node_type_copy &&
xnn_value_is_valid(value->type) && xnn_value_is_internal(value) &&
can_reorder_use(subgraph, consumer->outputs[0], producer_id,
consumer_id)) {
xnn_log_info("value %d fuse Copy Node #%" PRIu32
" into upstream %s Node #%" PRIu32,
value->id, consumer->id,
xnn_node_type_to_string(producer->type), producer->id);
assert(consumer->num_inputs == 1);
assert(consumer->num_outputs == 1);
const uint32_t fused_output_id = consumer->outputs[0];
assert(fused_output_id < subgraph->num_values);
subgraph->values[fused_output_id].producer = producer_id;
xnn_node_replace_output(producer, value->id, fused_output_id);
xnn_node_clear(consumer);
xnn_value_clear(value);
}
// Try to fuse copy downstream.
// E.g. --- v1 ---> (copy) --- value ---> (n2)
// If value is external or persistent, we cannot simply remove the copy,
// since we need to write to value.
if (producer->type == xnn_node_type_copy &&
xnn_value_is_valid(value->type) && xnn_value_is_internal(value) &&
can_reorder_use(subgraph, producer->inputs[0], producer_id,
consumer_id)) {
// We need to check that value is valid here because value could have
// been cleared by a previous optimization, this can happen if we have a
// chain of Copy(s), e.g.:
// ---v1--> (Copy1) ---v2--> (Copy2) ---v3--> (Copy3) ---v4-->
// v2 could have been cleared when we fused Copy2 upstream into Copy1,
// so v2 isn't valid anymore, but since v2's producer is also a Copy, we
// will incorrectly try to fuse Copy1 downstream into Copy2 (again).
xnn_log_info("value %d fuse Copy Node #%" PRIu32
" into downstream %s Node #%" PRIu32,
value->id, producer->id,
xnn_node_type_to_string(consumer->type), consumer->id);
assert(producer->num_outputs == 1);
assert(producer->num_inputs == 1);
const uint32_t copy_input_id = producer->inputs[0];
const uint32_t copy_output_id = producer->outputs[0];
bool found_consumer_input = false;
for (size_t i = 0; i < consumer->num_inputs; i++) {
if (consumer->inputs[i] == copy_output_id) {
consumer->inputs[i] = copy_input_id;
;
found_consumer_input = true;
// TODO(b/254734644): A consumer can only consume this value once,
// since we asserted earlier that value has only 1 consumer, so we
// can break here as there will be no other consumer inputs that has
// the same id.
break;
}
}
(void)found_consumer_input; // Silence unused variable warning in
// non-debug.
assert(found_consumer_input);
if (subgraph->values[copy_input_id].first_consumer == producer_id) {
subgraph->values[copy_input_id].first_consumer = consumer_id;
}
xnn_node_clear(producer);
xnn_value_clear(value);
}
}
}
return xnn_status_success;
}
// Returns true if `value` is a broadcast of a single constant.
static bool is_broadcasted_static(const struct xnn_value* value) {
// It really shouldn't be possible for a value with static data to also be
// an external input or have a producer. But some graphs do this (somehow), so
// we need to defend against these malformed cases.
return value->data != NULL &&
value->allocation_type == xnn_allocation_type_static &&
!xnn_value_is_external_input(value->flags) &&
value->producer == XNN_INVALID_NODE_ID &&
xnn_shape_multiply_all_dims(&value->shape) == 1;
}
static bool set_contains(const uint32_t* set, uint32_t set_size, uint32_t x) {
for (uint32_t i = 0; i < set_size; i++) {
if (set[i] == x) {
return true;
}
}
return false;
}
// Returns the Value ID of a unary input. Binary operators with one constant
// operand are considered unary operators by this function, and return the non-
// constant input. `unary_values` is a set of values that have already been
// determined to be part of a unary elementwise function.
static uint32_t is_pure_unary_elementwise(xnn_subgraph_t subgraph,
const struct xnn_node* node,
const uint32_t* unary_values,
uint32_t num_unary_values) {
switch (node->type) {
case xnn_node_type_unary_elementwise:
// A pure unary elementwise function takes exactly one tensor input and
// must have a valid operator. A node already fused into a LUT by this
// same pass (see `xnn_define_unary_elementwise_lut_in_place`) has 2
// inputs (the tensor and the LUT table) and
// `unary_operator == xnn_unary_invalid`; it isn't safe to treat that as
// a plain unary function and fuse it into another LUT.
if (node->num_inputs != 1 || node->unary_operator == xnn_unary_invalid) {
return XNN_INVALID_VALUE_ID;
}
assert(node->num_inputs >= 1);
return node->inputs[0];
case xnn_node_type_binary_elementwise: {
const uint32_t input_0_id = node->inputs[0];
const uint32_t input_1_id = node->inputs[1];
const struct xnn_value* input_0 = &subgraph->values[input_0_id];
const struct xnn_value* input_1 = &subgraph->values[input_1_id];
assert(node->num_inputs == 2);
if (is_broadcasted_static(input_0) &&
!xnn_value_is_static(input_1->allocation_type)) {
return input_1_id;
} else if (is_broadcasted_static(input_1) &&
!xnn_value_is_static(input_0->allocation_type)) {
return input_0_id;
} else if (set_contains(unary_values, num_unary_values, input_0_id) &&
set_contains(unary_values, num_unary_values, input_1_id)) {
// This is a unary elementwise operator if we've determined that both
// inputs are part of the same unary elementwise function. It doesn't
// matter which operand we return as long as it is in the `unary_values`
// set (which both are).
return input_0_id;
} else {
return XNN_INVALID_VALUE_ID;
}
}
default:
return XNN_INVALID_VALUE_ID;
}
}
// We will not fuse more than this many unary elementwise ops into a LUT in the
// function below.
#define XNN_MAX_UNARY_FUSION_NODES 10
#define XNN_MAX_UNARY_FUSION_VALUES (2 * XNN_MAX_UNARY_FUSION_NODES + 1)
// Find the index of `src_id` in `value_map`. Sets `is_new` to `true` if the
// value was inserted into the map, `false` if it was found in the map.
static uint32_t map_value_id(uint32_t* value_map, uint32_t src_id,
bool* is_new) {
for (uint32_t dst_id = 0;; ++dst_id) {
if (value_map[dst_id] == src_id) {
if (is_new) {
*is_new = false;
}
return dst_id;
} else if (value_map[dst_id] == XNN_INVALID_VALUE_ID) {
value_map[dst_id] = src_id;
if (is_new) {
*is_new = true;
}
return dst_id;
}
}
XNN_UNREACHABLE;
}
// Copy a value to a new subgraph, if it doesn't exist already.
static uint32_t copy_value_to_static_subgraph(xnn_subgraph_t src_subgraph,
const struct xnn_value* src_value,
uint32_t* value_map,
xnn_subgraph_t dst_subgraph) {
bool is_new = false;
uint32_t dst_id = map_value_id(value_map, src_value->id, &is_new);
if (is_new) {
assert(dst_id == dst_subgraph->num_values);
assert(dst_id < dst_subgraph->num_reserved_values);
struct xnn_value* dst_value = &dst_subgraph->values[dst_id];
dst_subgraph->num_values++;
*dst_value = *src_value;
dst_value->id = dst_id;
dst_value->producer = XNN_INVALID_NODE_ID;
dst_value->first_consumer = XNN_INVALID_NODE_ID;
// Clear the cleanup flag to avoid double-free: the source subgraph retains
// ownership of the data pointer.
dst_value->flags &= ~XNN_VALUE_FLAG_NEEDS_CLEANUP;
}
return dst_id;
}
// Copy a Node and its Values to a new subgraph, maintaining a map of Value IDs
// as it goes.
static struct xnn_node* copy_node_to_static_subgraph(
xnn_subgraph_t src_subgraph, const struct xnn_node* src_node,
uint32_t* value_map, xnn_subgraph_t dst_subgraph) {
assert(dst_subgraph->num_nodes < dst_subgraph->num_reserved_nodes);
struct xnn_node* dst_node = &dst_subgraph->nodes[dst_subgraph->num_nodes++];
xnn_node_copy(dst_node, src_node);
for (size_t i = 0; i < src_node->num_inputs; i++) {
const struct xnn_value* value = &src_subgraph->values[src_node->inputs[i]];
const uint32_t dst_id = copy_value_to_static_subgraph(
src_subgraph, value, value_map, dst_subgraph);
dst_node->inputs[i] = dst_id;
}
for (size_t i = 0; i < src_node->num_outputs; i++) {
const struct xnn_value* value = &src_subgraph->values[src_node->outputs[i]];
const uint32_t dst_id = copy_value_to_static_subgraph(
src_subgraph, value, value_map, dst_subgraph);
dst_node->outputs[i] = dst_id;
}
return dst_node;
}
// Pass 0:1:256 to the input of the subgraph representing a unary pure
// elementwise function, storing the result in `lut`.
static enum xnn_status run_subgraph_to_make_lut(xnn_subgraph_t subgraph,
uint32_t input_id,
uint32_t output_id,
uint8_t* lut) {
xnn_log_debug("Running unary subgraph to make LUT");
xnn_runtime_t runtime;
XNN_RETURN_IF_ERROR(xnn_create_runtime_v4(
subgraph, NULL, NULL, NULL, XNN_FLAG_NO_OPERATOR_FUSION, &runtime));
const size_t ramp_size = 256;
uint8_t ramp[256];
for (size_t i = 0; i < 256; i++) {
ramp[i] = i;
}
enum xnn_status status =
xnn_reshape_external_value(runtime, input_id, 1, &ramp_size);
if (status != xnn_status_success) {
goto fail;
}
status = xnn_reshape_runtime(runtime);
if (status != xnn_status_success) {
goto fail;
}
struct xnn_external_value externals[2];
externals[0].id = input_id;
externals[0].data = ramp;
externals[1].id = output_id;
externals[1].data = lut;
status = xnn_setup_runtime(runtime, 2, externals);
if (status != xnn_status_success) {
goto fail;
}
status = xnn_invoke_runtime(runtime);
fail:
xnn_delete_runtime(runtime);
return status;
}
static enum xnn_status replace_node_with_lut(xnn_subgraph_t subgraph,
struct xnn_node* node,
uint32_t input_id,
uint32_t unary_input_id,
xnn_subgraph_t unary_subgraph) {
XNN_RETURN_IF_ERROR(xnn_subgraph_reserve_values(subgraph, 1));
const uint32_t unary_output_id =
unary_subgraph->nodes[unary_subgraph->num_nodes - 1].outputs[0];
assert(unary_input_id != XNN_INVALID_VALUE_ID);
assert(unary_output_id != XNN_INVALID_VALUE_ID);
unary_subgraph->values[unary_input_id].flags |= XNN_VALUE_FLAG_EXTERNAL_INPUT;
unary_subgraph->values[unary_output_id].flags |=
XNN_VALUE_FLAG_EXTERNAL_OUTPUT;
unary_subgraph->values[unary_input_id].allocation_type =
xnn_allocation_type_external;
unary_subgraph->values[unary_output_id].allocation_type =
xnn_allocation_type_external;
uint8_t* lut = xnn_allocate_memory(256 * sizeof(uint8_t));
if (lut == NULL) {
xnn_log_error("failed to allocate LUT");
return xnn_status_out_of_memory;
}
enum xnn_status status = run_subgraph_to_make_lut(
unary_subgraph, unary_input_id, unary_output_id, lut);
if (status != xnn_status_success) {
// Failed to generate the LUT, abandon this fusion.
xnn_release_memory(lut);
return status;
}
// We don't have any other way to store a dynamic allocation in a subgraph
// except in a value.
struct xnn_value* lut_value = xnn_subgraph_new_internal_value(subgraph);
if (lut_value == NULL) {
xnn_release_memory(lut);
return xnn_status_out_of_memory;
}
lut_value->flags |= XNN_VALUE_FLAG_NEEDS_CLEANUP;
lut_value->data = lut;
lut_value->datatype = xnn_datatype_quint8;
lut_value->allocation_type = xnn_allocation_type_static;
lut_value->shape.num_dims = 1;
lut_value->shape.dim[0] = 256;
lut_value->type = xnn_value_type_dense_tensor;
// Clear the inputs that were replaced by a fused op.
for (uint32_t i = 0; i < node->num_inputs; i++) {
struct xnn_value* input = &subgraph->values[node->inputs[i]];
if (input->num_consumers == 1 && input->first_consumer == node->id) {
xnn_value_clear(input);
}
}
return xnn_define_unary_elementwise_lut_in_place(
node, input_id, node->outputs[0], lut_value->id);
}
void reshape_for_lut(struct xnn_value* value) {
value->shape.num_dims = 1;
value->shape.dim[0] = 256;
}
void xnn_subgraph_fuse_unary_quantized_into_lut(xnn_subgraph_t subgraph) {
// Find sequences of operators that are unary, quantized, elementwise, and
// pure functions. These can be fused into a single LUT op. Examples:
// - softsign(x) = x/(1 + abs(x))
// - softplus(x) = log(1 + exp(x))
// We allow intermediate values to be datatypes other than quantized (and
// allow convert ops), but the input and output values of the sequence must be
// quantized.
for (uint32_t n = 0; n < subgraph->num_nodes; n++) {
struct xnn_node* node = &subgraph->nodes[n];
if (node->type == xnn_node_type_invalid) {
// Node was fused away, skip.
continue;
}
const uint32_t input_id =
is_pure_unary_elementwise(subgraph, node, NULL, 0);
if (input_id == XNN_INVALID_VALUE_ID) {
continue;
}
const struct xnn_value* input_value = &subgraph->values[input_id];
if (input_value->datatype == xnn_datatype_invalid ||
xnn_datatype_size_bits(input_value->datatype) != 8) {
// This value is not a quantized 8-bit value, it can't be the input to a
// LUT.
continue;
}
// This node is a pure unary elementwise op with a quantized input. Can we
// fuse more ops with this one? Here, we assume the ops are in order, so
// this op must be the first in a chain we could fuse.
// We're going to build a subgraph for all the nodes we want to fuse.
struct xnn_subgraph unary_subgraph;
memset(&unary_subgraph, 0, sizeof(unary_subgraph));
struct xnn_value unary_values[XNN_MAX_UNARY_FUSION_VALUES];
uint32_t value_map[XNN_MAX_UNARY_FUSION_VALUES];
for (size_t i = 0; i < XNN_MAX_UNARY_FUSION_VALUES; i++) {
value_map[i] = XNN_INVALID_VALUE_ID;
}
struct xnn_node unary_nodes[XNN_MAX_UNARY_FUSION_NODES];
for (size_t i = 0; i < XNN_MAX_UNARY_FUSION_NODES; i++) {
unary_nodes[i].id = i;
}
unary_subgraph.values = &unary_values[0];
unary_subgraph.num_reserved_values = XNN_MAX_UNARY_FUSION_VALUES;
unary_subgraph.nodes = &unary_nodes[0];
unary_subgraph.num_reserved_nodes = XNN_MAX_UNARY_FUSION_NODES;
// Remember the nodes we put in the unary subgraph.
struct xnn_node* nodes_to_fuse[XNN_MAX_UNARY_FUSION_NODES];
for (size_t i = 0; i < XNN_MAX_UNARY_FUSION_NODES; i++) {
nodes_to_fuse[i] = NULL;
}
do {
// Add the node we have to the unary subgraph.
nodes_to_fuse[unary_subgraph.num_nodes] = node;
const struct xnn_node* new_node = copy_node_to_static_subgraph(
subgraph, node, value_map, &unary_subgraph);
assert(node->num_outputs == 1);
const struct xnn_value* output = &subgraph->values[node->outputs[0]];
if (output->num_consumers != 1 ||
output->first_consumer == XNN_INVALID_NODE_ID) {
// Don't try to fuse nodes that don't have exactly one valid consumer.
break;
}
assert(!xnn_value_is_external_output(output->flags));
reshape_for_lut(&unary_subgraph.values[new_node->outputs[0]]);
// Include the consumer in the unary subgraph.
node = &subgraph->nodes[output->first_consumer];
} while (is_pure_unary_elementwise(subgraph, node, value_map,
XNN_MAX_UNARY_FUSION_VALUES) !=
XNN_INVALID_VALUE_ID &&
unary_subgraph.num_nodes < XNN_MAX_UNARY_FUSION_NODES);
// We need the output to be an 8-bit LUT element. Go back through the unary
// subgraph until we find one.
while (unary_subgraph.num_nodes > 1) {
const struct xnn_node* unary_node =
&unary_subgraph.nodes[unary_subgraph.num_nodes - 1];
assert(unary_node->num_outputs == 1);
const uint32_t unary_output_id = unary_node->outputs[0];
struct xnn_value* unary_output = &unary_subgraph.values[unary_output_id];
if (unary_output->datatype != xnn_datatype_invalid &&
xnn_datatype_size_bits(unary_output->datatype) == 8) {
break;
}
// Remove this node from the subgraph.
xnn_value_clear(unary_output);
unary_subgraph.num_nodes--;
}
if (unary_subgraph.num_nodes > 1) {
// Update the last node of the fusion (the node we replace).
node = nodes_to_fuse[unary_subgraph.num_nodes - 1];
// Replace the fused nodes with a LUT op.
const uint32_t unary_input_id = map_value_id(value_map, input_id, NULL);
reshape_for_lut(&unary_subgraph.values[unary_input_id]);
if (replace_node_with_lut(subgraph, node, input_id, unary_input_id,
&unary_subgraph) == xnn_status_success) {
// We replaced this subgraph with a LUT, clear out the old values and
// nodes.
for (uint32_t i = 0; i + 1 < unary_subgraph.num_nodes; i++) {
struct xnn_node* fused_node = nodes_to_fuse[i];
assert(fused_node->num_outputs == 1);
struct xnn_value* fused_output =
&subgraph->values[fused_node->outputs[0]];
xnn_log_info("Value %d fuse Node #%" PRIu32
" into downstream quantized LUT Node #%" PRIu32,
fused_output->id, fused_node->id, node->id);
if (i > 0) {
// Remove this consumer from the inputs.
for (uint32_t i = 0; i < fused_node->num_inputs; i++) {
struct xnn_value* fused_input =
&subgraph->values[fused_node->inputs[i]];
assert(!xnn_value_is_external_input(fused_input->flags));
fused_input->num_consumers--;
if (fused_input->num_consumers == 0) {
xnn_value_clear(fused_input);
}
}
}
// We only need to clear output values. Input values could be used by
// other ops, and the ones that are outputs of another node in the
// fusion will be cleared here.
xnn_value_clear(fused_output);
xnn_node_clear(fused_node);
}
}
}
}
}
static void recursive_remove_node(xnn_subgraph_t subgraph, uint32_t node_id) {
struct xnn_node* node = &subgraph->nodes[node_id];
// Decrease the number of consumers on the inputs.
for (uint32_t input_id = 0; input_id < node->num_inputs; input_id++) {
if (is_repeated_input(node, input_id)) {
continue;
}
struct xnn_value* input_value = &subgraph->values[node->inputs[input_id]];
if (!xnn_value_is_external_input(input_value->flags) &&
--input_value->num_consumers == 0) {
if (input_value->producer != XNN_INVALID_NODE_ID) {
struct xnn_node* producer = &subgraph->nodes[input_value->producer];
if (producer->num_outputs == 1) {
recursive_remove_node(subgraph, producer->id);
}
}
xnn_value_clear(input_value);
}
}
xnn_node_clear(node);
}
void xnn_subgraph_clean_up(xnn_subgraph_t subgraph) {
// Start by removing any unreferenced values and the nodes that generate them.
bool changes;
do {
// Count the number of consumers for each value.
xnn_subgraph_analyze_consumers_and_producers(subgraph);
// Clear unreferenced values.
changes = false;
for (uint32_t i = 0; i < subgraph->num_values; i++) {
struct xnn_value* value = &subgraph->values[i];
if (value->type == xnn_value_type_invalid) {
continue;
}
if (!xnn_value_is_external_input(value->flags) &&
value->num_consumers == 0) {
if (value->producer != XNN_INVALID_NODE_ID) {
struct xnn_node* producer = &subgraph->nodes[value->producer];
if (producer->num_outputs == 1) {
changes = true;
recursive_remove_node(subgraph, producer->id);
}
}
xnn_value_clear(value);
}
}
} while (changes);
// Compact the nodes and sort them hierarchically (stably), if needed. The
// temporary memory needed for `nodes_map` and `values_ready` is allocated as
// a single block to reduce overheads.
uint32_t* nodes_map =
xnn_allocate_memory(sizeof(uint32_t) * subgraph->num_nodes +
sizeof(bool) * subgraph->num_values);
if (nodes_map == NULL) {
xnn_log_error("failed to allocate nodes_map scratch buffer");
return;
}
bool* values_ready = (bool*)&nodes_map[subgraph->num_nodes];
for (uint32_t i = 0; i < subgraph->num_values; i++) {
struct xnn_value* value = &subgraph->values[i];
values_ready[i] = value->producer == XNN_INVALID_NODE_ID ||
xnn_value_is_external_input(value->flags);
}
uint32_t left = 0;
uint32_t num_invalid_nodes = 0;
while (left + num_invalid_nodes < subgraph->num_nodes) {
const uint32_t old_left = left;
num_invalid_nodes = 0;
for (uint32_t i = left; i < subgraph->num_nodes; i++) {
struct xnn_node* node = &subgraph->nodes[i];
// Skip over invalid nodes.
if (node->type == xnn_node_type_invalid) {
num_invalid_nodes++;
continue;
}
// Check whether all inputs to this node have been produced.
bool all_values_avail = true;
for (uint32_t j = 0; all_values_avail && j < node->num_inputs; j++) {
uint32_t input_id = node->inputs[j];
assert(input_id != XNN_INVALID_VALUE_ID);
assert(input_id < subgraph->num_values);
all_values_avail = values_ready[input_id];
}
// If so, bubble this node down to the left end of the list of nodes.
if (all_values_avail) {
nodes_map[node->id] = left;
node->id = left;
for (uint32_t j = 0; j < node->num_outputs; j++) {
values_ready[node->outputs[j]] = true;
}
if (left < i) {
changes = true;
struct xnn_node tmp_node = *node;
if (subgraph->nodes[left].type == xnn_node_type_invalid) {
node->type = xnn_node_type_invalid;
} else {
memmove(&subgraph->nodes[left + 1], &subgraph->nodes[left],
(i - left) * sizeof(struct xnn_node));
}
subgraph->nodes[left] = tmp_node;
}
left++;
}
}
if (left == old_left) {
xnn_log_error("Failed to schedule all nodes in subgraph (possible cycle).");
break;
}
}
// Update the node IDs in the subgraph values if they have changed.
if (changes) {
for (uint32_t i = 0; i < subgraph->num_values; i++) {
struct xnn_value* value = &subgraph->values[i];
if (value->producer != XNN_INVALID_NODE_ID) {
value->producer = nodes_map[value->producer];
}
if (value->first_consumer != XNN_INVALID_NODE_ID) {
value->first_consumer = nodes_map[value->first_consumer];
}
}
}
// Always update the number of nodes just in case any trailing invalid nodes
// were clipped.
subgraph->num_nodes = left;
// Release temporarily allocated memory.
xnn_release_memory(nodes_map);
}
static bool convert_gemm_to_qduint8(
const enum xnn_datatype input_datatype,
const enum xnn_node_type consumer_type,
const enum xnn_datatype consumer_weights_type) {
// Identify the `qdint8` and `qduint8` configs for the consumers of this op.
const struct xnn_gemm_config* original_config = NULL;
const struct xnn_gemm_config* unsigned_config = NULL;
if (input_datatype == xnn_datatype_fp32) {
if (consumer_weights_type == xnn_datatype_qcint2) {
original_config = xnn_init_qd8_f32_qc2w_gemm_config();
unsigned_config = xnn_init_qdu8_f32_qc2w_gemm_config();
} else if (consumer_weights_type == xnn_datatype_qcint4) {
original_config = xnn_init_qd8_f32_qc4w_gemm_config();
unsigned_config = xnn_init_qdu8_f32_qc4w_gemm_config();
} else if (consumer_weights_type == xnn_datatype_qcint8) {
original_config = xnn_init_qd8_f32_qc8w_gemm_config();
switch (consumer_type) {
case xnn_node_type_batch_matrix_multiply:
case xnn_node_type_fully_connected:
unsigned_config = xnn_init_qdu8_f32_qc8w_gemm_config();
break;
case xnn_node_type_convolution_2d:
case xnn_node_type_deconvolution_2d:
unsigned_config = xnn_init_qdu8_f32_qc8w_igemm_config();
break;
default:
XNN_UNREACHABLE;
}
} else if (consumer_weights_type == xnn_datatype_qbint4) {
original_config = xnn_init_qd8_f32_qb4w_gemm_config();
unsigned_config = xnn_init_qdu8_f32_qb4w_gemm_config();
}
} else if (input_datatype == xnn_datatype_fp16) {
if (consumer_weights_type == xnn_datatype_qcint4) {
original_config = xnn_init_qd8_f16_qc4w_gemm_config();
unsigned_config = xnn_init_qdu8_f16_qc4w_gemm_config();
} else if (consumer_weights_type == xnn_datatype_qcint8) {
switch (consumer_type) {
case xnn_node_type_batch_matrix_multiply:
case xnn_node_type_fully_connected:
original_config = xnn_init_qd8_f16_qc8w_gemm_config();
unsigned_config = xnn_init_qdu8_f16_qc8w_gemm_config();
break;
case xnn_node_type_convolution_2d:
case xnn_node_type_deconvolution_2d:
original_config = xnn_init_qd8_f16_qc8w_igemm_config();
unsigned_config = xnn_init_qdu8_f16_qc8w_gemm_config();
break;
default:
XNN_UNREACHABLE;
}
}
}
// If the `qduint8` config is better than the `qdint8` config, use it
// instead.
bool convert_to_qu8 = false;
if (unsigned_config) {
if (original_config == NULL) {
convert_to_qu8 = true;
} else {
enum xnn_arch_flags qdu8_arch = unsigned_config->arch;
enum xnn_arch_flags qd8_arch = original_config->arch;
if (qdu8_arch > qd8_arch) {
convert_to_qu8 = true;
}
}
}
return convert_to_qu8;
}
static void swap_value_pointers(struct xnn_value** a, struct xnn_value** b) {
struct xnn_value* temp = *a;
*a = *b;
*b = temp;
}
static float get_scalar_value_as_float(struct xnn_value* val) {
union {
float fp32;
xnn_float16 fp16;
xnn_bfloat16 bf16;
int32_t int32;
int8_t int8;
uint8_t uint8;
} data;
memcpy(&data, val->data, xnn_datatype_size_bytes(val->datatype));
switch (val->datatype) {
case xnn_datatype_fp32:
return data.fp32;
case xnn_datatype_fp16:
return xnn_float16_to_float(data.fp16);
case xnn_datatype_bf16:
return xnn_bfloat16_to_float(data.bf16);
case xnn_datatype_int32:
return data.int32;
case xnn_datatype_qint8:
return ((float)data.int8 - val->quantization.zero_point) *
val->quantization.scale;
case xnn_datatype_quint8:
return ((float)data.uint8 - val->quantization.zero_point) *
val->quantization.scale;
default:
return NAN;
}
}
// Modifies the subgraph to elide the nodes between the values `input_id` and
// `output_id`, and returns the number of changes made, or zero if the
// computation cannot be elided.
//
// This either sets all the consumers of the value `output_id` to `input_id`, or
// the output of the producer of `input_id`, and all consumers of `input_id`, to
// `output_id`, depending on whether the input/output values are external or
// not.
//
// If both `input_id` and `output_id` are external, or `output_id` is external
// and `input_id` has more than one consumer, then the computation can not be
// elided and `0` is returned.
static size_t short_circuit(xnn_subgraph_t subgraph, uint32_t input_id,
uint32_t output_id) {
size_t changes = 0;
struct xnn_value* output_value = &subgraph->values[output_id];
struct xnn_value* input_value = &subgraph->values[input_id];
// If the old value is an external value, point its producer in the right
// direction.
if (xnn_value_is_external_output(output_value->flags)) {
// If the input node does not have a producer, or it's also an output, then
// there's nothing we can do, so don't change anything.
const uint32_t producer_id = input_value->producer;
if (xnn_value_is_external_output(input_value->flags) ||
producer_id == XNN_INVALID_NODE_ID) {
return 0;
}
// Change the producer's output.
struct xnn_node* producer_node = &subgraph->nodes[producer_id];
for (int k = 0; k < producer_node->num_outputs; k++) {
if (producer_node->outputs[k] == input_id) {
producer_node->outputs[k] = output_id;
changes++;
}
}
output_value->producer = producer_id;
// Swap the input/output values so that consumers of `input_id` point to
// `output_id` instead.
swap_value_pointers(&output_value, &input_value);
output_id = output_value->id;
input_id = input_value->id;
}
// Swap consumers of `output_id` with `input_id`.
input_value->num_consumers--;
input_value->first_consumer = XNN_INVALID_VALUE_ID;
for (int k = 0; k < subgraph->num_nodes; k++) {
for (int j = 0; j < subgraph->nodes[k].num_inputs; j++) {
if (subgraph->nodes[k].inputs[j] == output_id) {
subgraph->nodes[k].inputs[j] = input_id;
if (!is_repeated_input(&subgraph->nodes[k], j)) {
input_value->num_consumers++;
}
if (input_value->first_consumer == XNN_INVALID_VALUE_ID) {
input_value->first_consumer = k;
}
changes++;
} else if (subgraph->nodes[k].inputs[j] == input_id &&
input_value->first_consumer == XNN_INVALID_VALUE_ID) {
input_value->first_consumer = k;
}
}
}
return changes;
}
static struct xnn_node* move_last_node_to(xnn_subgraph_t subgraph,
uint32_t node_id) {
assert(node_id < subgraph->num_nodes);
struct xnn_node* node = &subgraph->nodes[node_id];
*node = subgraph->nodes[--subgraph->num_nodes];
node->id = node_id;
return node;
}
// Replace `mul(x, x)` with `sqr(x)` for consistency.
static enum xnn_status optimize_common_subgraphs_mul_to_sqr(
xnn_subgraph_t subgraph, uint32_t node_id, size_t* changes) {
XNN_RETURN_IF_ERROR(xnn_subgraph_reserve_nodes(subgraph, 1));
struct xnn_node* node = &subgraph->nodes[node_id];
if (node->type != xnn_node_type_binary_elementwise ||
node->binary_operator != xnn_binary_multiply ||
node->inputs[0] != node->inputs[1]) {
return xnn_status_success;
}
const uint32_t input_id = node->inputs[0];
const uint32_t output_id = node->outputs[0];
xnn_log_info("Converting node mul[#%u](v%03u, v%03u) to sqr[#%u](v%03u).",
node_id, input_id, input_id, node_id, input_id);
XNN_RETURN_IF_ERROR(
xnn_define_unary(subgraph, xnn_unary_square,
/*params=*/NULL, input_id, output_id, node->flags),
"Failed to create new unary-elementwise node.");
node = move_last_node_to(subgraph, node_id);
(*changes)++;
return xnn_status_success;
}
// Replace `div(x, x)` or `sub(x, x)` with a constant `1.0` or `0.0`,
// respectively.
static enum xnn_status optimize_common_subgraphs_binary_to_const(
xnn_subgraph_t subgraph, uint32_t node_id, size_t* changes) {
struct xnn_node* node = &subgraph->nodes[node_id];
if (node->type != xnn_node_type_binary_elementwise ||
node->inputs[0] != node->inputs[1]) {
return xnn_status_success;
}
struct xnn_value* output_value = &subgraph->values[node->outputs[0]];
if (node->binary_operator == xnn_binary_divide) {
if ((output_value->flags & XNN_VALUE_FLAG_IS_ONE) == 0) {
xnn_log_info(
"Marking output of node div[#%u](v%03u, v%03u) as const 1.0.",
node_id, node->inputs[0], node->inputs[0]);
output_value->flags |= XNN_VALUE_FLAG_IS_ONE;
(*changes)++;
}
} else if (node->binary_operator == xnn_binary_subtract) {
if ((output_value->flags & XNN_VALUE_FLAG_IS_ZERO) == 0) {
xnn_log_info(
"Marking output of node sub[#%u](v%03u, v%03u) as const 0.0.",
node_id, node->inputs[0], node->inputs[0]);
output_value->flags |= XNN_VALUE_FLAG_IS_ZERO;
(*changes)++;
}
}
return xnn_status_success;
}
static void convert_static_value_to_fp32(struct xnn_value* value) {
assert(xnn_value_is_static(value->allocation_type));
if (value->flags & XNN_VALUE_FLAG_NEEDS_CLEANUP) {
xnn_release_memory(value->data);
}
value->data = xnn_allocate_memory(sizeof(float));
float data = get_scalar_value_as_float(value);
memcpy(value->data, &data, sizeof(float));
value->flags |= XNN_VALUE_FLAG_NEEDS_CLEANUP;
value->datatype = xnn_datatype_fp32;
}
// Replace `mul(reduce_sum(x), 1/n)`, `div(reduce_sum(x), n)` or
// `mul(reduce_sum_squared(x), 1/n)`, `div(reduce_sum_squared(x), n)`
// with `reduce_mean(x)` or `reduce_mean_squared(x)`, respectively.
static enum xnn_status widen_fp16_accumulators(xnn_subgraph_t subgraph,
uint32_t node_id,
size_t* changes) {
XNN_RETURN_IF_ERROR(xnn_subgraph_reserve_nodes(subgraph, 1));
XNN_RETURN_IF_ERROR(xnn_subgraph_reserve_values(subgraph, 1));
struct xnn_node* node = &subgraph->nodes[node_id];
if (node->type != xnn_node_type_binary_elementwise ||
(node->binary_operator != xnn_binary_multiply &&
node->binary_operator != xnn_binary_divide)) {
return xnn_status_success;
}
struct xnn_value* reduced_value = &subgraph->values[node->inputs[0]];
struct xnn_value* arg_value = &subgraph->values[node->inputs[1]];
if (xnn_shape_multiply_all_dims(&arg_value->shape) != 1 ||
!xnn_value_is_static(arg_value->allocation_type)) {
if (xnn_shape_multiply_all_dims(&reduced_value->shape) == 1 &&
xnn_value_is_static(reduced_value->allocation_type)) {
swap_value_pointers(&reduced_value, &arg_value);
} else {
return xnn_status_success;
}
}
// Check that one of the args is a sum or sum2 reduction.
if (reduced_value->datatype != xnn_datatype_fp16 ||
reduced_value->producer == XNN_INVALID_NODE_ID) {
return xnn_status_success;
}
if (reduced_value->num_consumers > 1 || arg_value->num_consumers > 1) {
// Don't rewrite if we might modify an unrelated consumer.
return xnn_status_success;
}
struct xnn_node* reduce_node = &subgraph->nodes[reduced_value->producer];
const enum xnn_node_type reduce_node_type = reduce_node->type;
if (!(reduce_node_type == xnn_node_type_static_sum ||
reduce_node_type == xnn_node_type_static_sum_squared)) {
return xnn_status_success;
}
// Rewrite the internal values to this subgraph to be fp32.
reduced_value->datatype = xnn_datatype_fp32;
convert_static_value_to_fp32(arg_value);
uint32_t output_id = node->outputs[0];
struct xnn_value* output_value = &subgraph->values[output_id];
uint32_t output_fp32_id;
enum xnn_status status = xnn_define_tensor_value(
subgraph, xnn_datatype_fp32, output_value->shape.num_dims,
output_value->shape.dim,
/*data=*/NULL, XNN_INVALID_VALUE_ID,
/*flags=*/0, &output_fp32_id);
if (status != xnn_status_success) {
return status;
}
node->outputs[0] = output_fp32_id;
status = xnn_define_unary(subgraph, xnn_unary_convert, /*params=*/NULL,
output_fp32_id, output_id, /*flags=*/0);
if (status != xnn_status_success) {
return status;
}
move_last_node_to(subgraph, node_id);
xnn_log_info(
"Converted %s[#%u](reduce_sum%s[#%u](v%03u), v%03u) to "
"fp32.",
(node->binary_operator == xnn_binary_multiply) ? "mul" : "div", node_id,
reduce_node_type == xnn_node_type_static_sum_squared ? "_squared" : "",
reduced_value->producer, reduce_node->inputs[0], arg_value->id);
(*changes)++;
return xnn_status_success;
}
// Convert `reduce_sum(sqr(a))` or `reduce_mean(sqr(a))` to
// `reduce_sum_squared(a)` or `reduce_mean_squared(a)`, respectively.
static enum xnn_status optimize_common_subgraphs_reduce_sum_to_square(
xnn_subgraph_t subgraph, uint32_t node_id, size_t* changes) {
XNN_RETURN_IF_ERROR(xnn_subgraph_reserve_nodes(subgraph, 1));
struct xnn_node* node = &subgraph->nodes[node_id];
if (node->type != xnn_node_type_static_sum &&
node->type != xnn_node_type_static_mean) {
return xnn_status_success;
}
struct xnn_value* input_value = &subgraph->values[node->inputs[0]];
if (!(input_value->datatype == xnn_datatype_fp16 ||
input_value->datatype == xnn_datatype_fp32) ||
input_value->producer == XNN_INVALID_NODE_ID) {
return xnn_status_success;
}
struct xnn_node* input_producer_node =
&subgraph->nodes[input_value->producer];
if (input_producer_node->type != xnn_node_type_unary_elementwise ||
input_producer_node->unary_operator != xnn_unary_square) {
return xnn_status_success;
}
const uint32_t output_id = node->outputs[0];
const size_t num_reduction_axes = node->params.reduce.num_reduction_axes;
int64_t reduction_axes[XNN_MAX_TENSOR_DIMS];
memcpy(reduction_axes, node->params.reduce.reduction_axes,
num_reduction_axes * sizeof(int64_t));
const enum xnn_node_type node_type = node->type;
XNN_RETURN_IF_ERROR(
xnn_define_static_reduce_v2(
subgraph,
node_type == xnn_node_type_static_sum ? xnn_reduce_sum_squared
: xnn_reduce_mean_squared,
num_reduction_axes, reduction_axes, input_producer_node->inputs[0],
output_id, node->flags),
"Failed to create new `%s` node.",
xnn_node_type_to_string(node_type == xnn_node_type_static_sum
? xnn_node_type_static_sum_squared
: xnn_node_type_static_mean_squared));
node = move_last_node_to(subgraph, node_id);
xnn_log_info(
"Converted reduce_%s[#%u](sqr[#%u](v%03u)) to "
"reduce_%s_squared[#%u](v%03u).",
node_type == xnn_node_type_static_sum ? "sum" : "mean", node_id,
input_value->producer, node->inputs[0],
node_type == xnn_node_type_static_sum ? "sum" : "mean", node_id,
node->inputs[0]);
(*changes)++;
return xnn_status_success;
}
// XNNPACK doesn't need explicit braodcasting for binary and
// batch_matrix_multiply nodes, so check if it can be elided.
static enum xnn_status optimize_common_subgraphs_broadcast(
xnn_subgraph_t subgraph, uint32_t node_id, size_t* changes) {
XNN_RETURN_IF_ERROR(xnn_subgraph_reserve_nodes(subgraph, 1));
XNN_RETURN_IF_ERROR(xnn_subgraph_reserve_values(subgraph, 1));
struct xnn_node* node = &subgraph->nodes[node_id];
if (node->type != xnn_node_type_static_broadcast) {
return xnn_status_success;
}
const uint32_t input_id = node->inputs[0];
const uint32_t output_id = node->outputs[0];
const struct xnn_value* output_value = &subgraph->values[output_id];
// Find all consumers of the broadcast node's output.
uint32_t num_consumers = output_value->num_consumers;
for (uint32_t k = node_id + 1; k < subgraph->num_nodes && num_consumers;
k++) {
struct xnn_node* consumer = &subgraph->nodes[k];
for (uint32_t j = 0; j < consumer->num_inputs; j++) {
if (consumer->inputs[j] == output_id) {
// If the consumer is known to broadcast implicitly,
// short-circuit it.
if (consumer->type == xnn_node_type_binary_elementwise ||
consumer->type == xnn_node_type_batch_matrix_multiply) {
if (!is_repeated_input(consumer, j)) {
num_consumers--;
}
consumer->inputs[j] = input_id;
(*changes)++;
}
}
}
}
// If the broadcast could not be completely elided, then replace it
// with a broadcasting `add` of zero.
// TODO(b/421935339) Replace this with a more principled solution if
// needed since both the "zero" and output values could be quite
// large.
if (num_consumers) {
// Compute the shape of the right-hand operand to the binary op
// that, when added to the input, will result in the correct
// output shape.
size_t shape[XNN_MAX_TENSOR_DIMS];
size_t num_dims = node->params.static_reshape.new_shape.num_dims;
const struct xnn_value* input_value = &subgraph->values[input_id];
const size_t* old_shape = input_value->shape.dim;
const size_t* new_shape = node->params.static_reshape.new_shape.dim;
size_t num_elements = 1;
for (uint32_t k = 0; k < num_dims; k++) {
shape[k] = (new_shape[k] == 0 || new_shape[k] == old_shape[k])
? 1
: new_shape[k];
num_elements *= shape[k];
}
// Create a static right-hand side value filled with zeros.
void* data = xnn_allocate_zero_memory(
num_elements * xnn_datatype_size_bytes(input_value->datatype) +
XNN_EXTRA_BYTES);
uint32_t new_value_id;
XNN_RETURN_IF_ERROR(
xnn_datatype_is_quantized(input_value->datatype)
? xnn_define_quantized_tensor_value(
subgraph, input_value->datatype,
/*zero_point=*/0, /*quantization_scale=*/1.0, num_dims, shape,
data,
/*external_id=*/XNN_INVALID_VALUE_ID,
/*flags=*/XNN_VALUE_FLAG_NEEDS_CLEANUP |
XNN_VALUE_FLAG_IS_ZERO,
&new_value_id)
: xnn_define_tensor_value(subgraph, input_value->datatype, num_dims,
shape, data,
/*external_id=*/XNN_INVALID_VALUE_ID,
/*flags=*/XNN_VALUE_FLAG_NEEDS_CLEANUP |
XNN_VALUE_FLAG_IS_ZERO,
&new_value_id),
"Failed to create static zero tensor.");
// Replace the broadcast node with a binary add.
XNN_RETURN_IF_ERROR(
xnn_define_binary(subgraph, xnn_binary_add,
/*params=*/NULL, input_id, new_value_id, output_id,
node->flags | XNN_NODE_FLAG_DONT_ELIDE),
"Failed to create Binary Addition node.");
node = move_last_node_to(subgraph, node_id);
xnn_log_warning(
"Converted static_broadcast[#%u](v%03u) to a broadcasting "
"add[#%u](v%03u, 0).",
node_id, input_id, node_id, input_id);
(*changes)++;
} else {
xnn_log_info(
"Removed static_broadcast[#%u](v%03u) since it is handled implicitly "
"by all its consumers.",
node_id, input_id);
}
return xnn_status_success;
}
// Merge `reshape(expand_dims(...))` and `expand_dims(reshape(...))` into a
// single `reshape`.
static enum xnn_status optimize_common_subgraphs_merge_reshapes(
xnn_subgraph_t subgraph, uint32_t node_id, size_t* changes) {
struct xnn_node* node = &subgraph->nodes[node_id];
if (node->type != xnn_node_type_static_reshape &&
node->type != xnn_node_type_static_expand_dims) {
return xnn_status_success;
}
// Check that we are the only consumer of the input node.
const uint32_t input_id = node->inputs[0];
struct xnn_value* input_value = &subgraph->values[input_id];
if (input_value->producer == XNN_INVALID_NODE_ID ||
input_value->num_consumers != 1) {
return xnn_status_success;
}
struct xnn_node* input_producer = &subgraph->nodes[input_value->producer];
// Check all the interesting combinations.
// reshape(reshape(x)).
if (input_producer->type == xnn_node_type_static_reshape &&
node->type == xnn_node_type_static_reshape) {
XNN_RETURN_IF_ERROR(
xnn_shape_fill_gaps(&input_producer->params.static_reshape.new_shape,
&node->params.static_reshape.new_shape));
node->inputs[0] = input_producer->inputs[0];
node->flags |= XNN_NODE_FLAG_DONT_ELIDE;
input_value = &subgraph->values[node->inputs[0]];
if (input_value->first_consumer == input_producer->id) {
input_value->first_consumer = node->id;
}
xnn_log_info(
"Replaced static_reshape[#%u](static_reshape[#%u](v%03u)) with "
"static_reshape[#%u](v%03u).",
node_id, input_producer->id, node->inputs[0], node_id, node->inputs[0]);
xnn_node_clear(input_producer);
(*changes)++;
}
// reshape(expand_dims(x)).
else if (input_producer->type == xnn_node_type_static_expand_dims &&
node->type == xnn_node_type_static_reshape) {
node->inputs[0] = input_producer->inputs[0];
input_value = &subgraph->values[node->inputs[0]];
if (input_value->first_consumer == input_producer->id) {
input_value->first_consumer = node->id;
}
xnn_log_info(
"Replaced static_reshape[#%u](static_expand_dims[#%u](v%03u)) with "
"static_reshape[#%u](v%03u).",
node_id, input_producer->id, node->inputs[0], node_id, node->inputs[0]);
xnn_node_clear(input_producer);
(*changes)++;
}
// expand_dims(reshape(x)).
else if (input_producer->type == xnn_node_type_static_reshape &&
node->type == xnn_node_type_static_expand_dims) {
const struct xnn_shape* reshape =
&input_producer->params.static_reshape.new_shape;
const struct xnn_shape* expanded_dims =
&node->params.static_reshape.new_shape;
if (reshape->num_dims + expanded_dims->num_dims > XNN_MAX_TENSOR_DIMS) {
return xnn_status_success; // Skip optimization, let runtime validate.
}
struct xnn_shape new_shape = {
.num_dims = reshape->num_dims + expanded_dims->num_dims, .dim = {0}};
for (uint32_t idx_expanded = 0, idx_reshape = 0, k = 0;
k < new_shape.num_dims; k++) {
if (idx_expanded < expanded_dims->num_dims &&
expanded_dims->dim[idx_expanded] == k) {
new_shape.dim[k] = 1;
idx_expanded++;
} else {
new_shape.dim[k] = reshape->dim[idx_reshape++];
}
}
input_producer->params.static_reshape.new_shape = new_shape;
input_producer->outputs[0] = node->outputs[0];
struct xnn_value* output_value =
&subgraph->values[input_producer->outputs[0]];
output_value->producer = input_producer->id;
xnn_log_info(
"Replaced static_expand_dims[#%u](static_reshape[#%u](v%03u)) with "
"static_reshape[#%u](v%03u).",
input_producer->id, node_id, input_producer->inputs[0],
input_producer->id, input_producer->inputs[0]);
xnn_node_clear(node);
(*changes)++;
}
return xnn_status_success;
}
// Apply `static_reshape` and `static_expand_dims` to static values directly.
static enum xnn_status optimize_common_subgraphs_static_reshapes(
xnn_subgraph_t subgraph, uint32_t node_id, size_t* changes) {
struct xnn_node* node = &subgraph->nodes[node_id];
if (node->type != xnn_node_type_static_reshape &&
node->type != xnn_node_type_static_expand_dims) {
return xnn_status_success;
}
const uint32_t input_id = node->inputs[0];
const uint32_t output_id = node->outputs[0];
struct xnn_value* input_value = &subgraph->values[input_id];
struct xnn_value* output_value = &subgraph->values[output_id];
// Is the input's shape static?
// Is the reshape the only consumer of the input value?
// If the output is external, don't do anything.
if (!(input_value->flags & XNN_VALUE_FLAG_SHAPE_IS_STATIC) ||
input_value->num_consumers > 1 ||
xnn_value_is_external_output(output_value->flags)) {
return xnn_status_success;
}
// Set the shape of the static-shaped value.
struct xnn_shape new_shape;
if (node->type == xnn_node_type_static_reshape) {
// Replace the old shape with the new shape, filling any gaps from the input
// shape.
new_shape = node->params.static_reshape.new_shape;
XNN_RETURN_IF_ERROR(xnn_shape_fill_gaps(&input_value->shape, &new_shape),
"Could not fill gaps for reshape[#%u].", node_id);
} else if (node->type == xnn_node_type_static_expand_dims) {
const struct xnn_shape* new_dims = &node->params.static_reshape.new_shape;
new_shape.num_dims = input_value->shape.num_dims + new_dims->num_dims;
if (new_shape.num_dims > XNN_MAX_TENSOR_DIMS) {
return xnn_status_success; // Skip optimization, let runtime validate.
}
for (uint32_t idx_new = 0, idx_old = 0, k = 0; k < new_shape.num_dims;
k++) {
if (idx_new < new_dims->num_dims && new_dims->dim[idx_new] == k) {
new_shape.dim[k] = 1;
idx_new++;
} else {
new_shape.dim[k] = input_value->shape.dim[idx_old++];
}
}
}
// If the input is a static value, apply the new shape to it directly.
bool elide = true;
if (xnn_value_is_static(input_value->allocation_type)) {
input_value->shape = new_shape;
} else {
elide = xnn_shape_match(&new_shape, &input_value->shape);
}
if (elide) {
// All nodes should consume the reshaped static-shaped value directly.
*changes += short_circuit(subgraph, input_id, output_id);
xnn_log_info("Inlined %s[#%u](v%03u) of static-shaped value v%03u.",
node->type == xnn_node_type_static_reshape
? "static_reshape"
: "static_expand_dims",
node_id, input_id, input_id);
xnn_node_clear(node);
}
return xnn_status_success;
}
// Convert min/max operations with a single static value to a unary `clamp`
// node.
static enum xnn_status optimize_common_subgraphs_min_max_to_clamp(
xnn_subgraph_t subgraph, uint32_t node_id, size_t* changes) {
XNN_RETURN_IF_ERROR(xnn_subgraph_reserve_nodes(subgraph, 1));
struct xnn_node* node = &subgraph->nodes[node_id];
if (node->type != xnn_node_type_binary_elementwise ||
(node->binary_operator != xnn_binary_maximum &&
node->binary_operator != xnn_binary_minimum)) {
return xnn_status_success;
}
// `arg_value` must be a static scalar value, but we can swap the arguments to
// make that true.
struct xnn_value* input_value = &subgraph->values[node->inputs[0]];
struct xnn_value* arg_value = &subgraph->values[node->inputs[1]];
const bool input_is_static_scalar =
xnn_value_is_const(input_value->flags) ||
(xnn_shape_multiply_all_dims(&input_value->shape) == 1 &&
xnn_value_is_static(input_value->allocation_type));
const bool arg_is_static_scalar =
xnn_value_is_const(arg_value->flags) ||
(xnn_shape_multiply_all_dims(&arg_value->shape) == 1 &&
xnn_value_is_static(arg_value->allocation_type));
if (input_is_static_scalar) {
if (!arg_is_static_scalar) {
swap_value_pointers(&arg_value, &input_value);
}
} else if (!arg_is_static_scalar) {
return xnn_status_success;
}
if (arg_value->shape.num_dims > input_value->shape.num_dims) {
// The min or max operator is broadcasting the input to match the scalar's
// rank, so we can't replace the operator.
return xnn_status_success;
}
// Extract the min/max argument.
const float arg_as_float = (arg_value->flags & XNN_VALUE_FLAG_IS_ZERO) ? 0.0f
: (arg_value->flags & XNN_VALUE_FLAG_IS_ONE)
? 1.0f
: get_scalar_value_as_float(arg_value);
const bool is_min = (node->binary_operator == xnn_binary_minimum);
const bool is_max = (node->binary_operator == xnn_binary_maximum);
// Replace the binary `min`/`max` node with a unary
// `clamp` node.
union xnn_unary_params params;
params.clamp.max = is_min ? arg_as_float : INFINITY;
params.clamp.min = is_max ? arg_as_float : -INFINITY;
XNN_RETURN_IF_ERROR(
xnn_define_unary(subgraph, xnn_unary_clamp, &params, input_value->id,
node->outputs[0], node->flags),
"Failed to create new `Clamp` node.");
node = move_last_node_to(subgraph, node_id);
xnn_log_info("Converted %s[#%u](v%03u, %f) to clamp[#%u](v%03u, [%f, %f]).",
is_max ? "maximum" : "minimum", node_id, input_value->id,
arg_as_float, node_id, input_value->id,
node->params.unary.clamp.min, node->params.unary.clamp.max);
(*changes)++;
return xnn_status_success;
}
// Folds unary clamp operations into the previous operator's activation, if
// possible.
static enum xnn_status optimize_common_subgraphs_merge_clamps(
xnn_subgraph_t subgraph, uint32_t node_id, size_t* changes) {
struct xnn_node* node = &subgraph->nodes[node_id];
if (!is_clamp(node)) {
return xnn_status_success;
}
// Verify that we are the input's only consumer and that it has a producer,
// and that its producer is a (or can) clamp.
struct xnn_value* input_value = &subgraph->values[node->inputs[0]];
if (input_value->num_consumers > 1 ||
input_value->producer == XNN_INVALID_NODE_ID) {
return xnn_status_success;
}
struct xnn_node* input_producer_node =
&subgraph->nodes[input_value->producer];
if (!has_clamp(input_producer_node)) {
return xnn_status_success;
}
// Update the input producer's clamping values based on the current clamp.
// Note that the order of the min/max matches that of the `clamp` ops
// themsevels and is required to ensure correctness when merging
// non-overlapping clamps.
if (is_clamp(input_producer_node)) {
input_producer_node->params.unary.clamp.min =
math_min_f32(math_max_f32(input_producer_node->params.unary.clamp.min,
node->params.unary.clamp.min),
node->params.unary.clamp.max);
input_producer_node->params.unary.clamp.max =
math_min_f32(math_max_f32(input_producer_node->params.unary.clamp.max,
node->params.unary.clamp.min),
node->params.unary.clamp.max);
} else {
input_producer_node->activation.output_min =
math_min_f32(math_max_f32(input_producer_node->activation.output_min,
node->params.unary.clamp.min),
node->params.unary.clamp.max);
input_producer_node->activation.output_max =
math_min_f32(math_max_f32(input_producer_node->activation.output_max,
node->params.unary.clamp.min),
node->params.unary.clamp.max);
}
// Elide the `clamp` node by just skipping it.
input_producer_node->outputs[0] = node->outputs[0];
xnn_node_clear(node);
xnn_log_info("Merged clamp[#%u](%s[#%u](...)) into clamping %s[#%u](...).",
node_id, xnn_node_type_to_string(input_producer_node->type),
input_producer_node->id,
xnn_node_type_to_string(input_producer_node->type),
input_producer_node->id);
(*changes)++;
return xnn_status_success;
}
// Remove spurious unary clamp operations.
static enum xnn_status optimize_common_subgraphs_spurious_clamps(
xnn_subgraph_t subgraph, uint32_t node_id, size_t* changes) {
XNN_RETURN_IF_ERROR(xnn_subgraph_reserve_nodes(subgraph, 1));
struct xnn_node* node = &subgraph->nodes[node_id];
if (node->type != xnn_node_type_unary_elementwise ||
node->unary_operator != xnn_unary_clamp) {
return xnn_status_success;
}
// TODO: b/455537016 - We can use different bounds for different datatypes.
if (node->params.unary.clamp.min != -INFINITY ||
node->params.unary.clamp.max != INFINITY) {
return xnn_status_success;
}
const uint32_t output_id = node->outputs[0];
struct xnn_value* input_value = &subgraph->values[node->inputs[0]];
if (short_circuit(subgraph, input_value->id, output_id)) {
xnn_node_clear(node);
xnn_log_info("Elided spurious clamp[#%u](v%03u).", node_id,
input_value->id);
} else {
// If this node cannot be elided, then replace it with a `copy`.
XNN_RETURN_IF_ERROR(
xnn_define_copy(subgraph, input_value->id, output_id, node->flags),
"Failed to create new `Copy` node.");
node = move_last_node_to(subgraph, node_id);
xnn_log_info("Replaced spurious clamp[#%u](v%03u) with copy[#%u](v%03u).",
node->id, input_value->id, node_id, input_value->id);
}
(*changes)++;
return xnn_status_success;
}
// Propagate constants through constant-preserving ops.
static void propagate_constants(xnn_subgraph_t subgraph, uint32_t node_id) {
struct xnn_node* node = &subgraph->nodes[node_id];
if (node->type == xnn_node_type_invalid) {
return;
}
switch (node->type) {
// Shape and size changes don't affect the constants.
case xnn_node_type_copy:
case xnn_node_type_even_split:
case xnn_node_type_fuse_dims:
case xnn_node_type_split_dims:
case xnn_node_type_static_broadcast:
case xnn_node_type_static_expand_dims:
case xnn_node_type_static_reshape:
case xnn_node_type_static_slice:
case xnn_node_type_static_transpose:
// Operations that propagate constants.
case xnn_node_type_convert:
case xnn_node_type_static_mean_squared:
case xnn_node_type_static_mean:
case xnn_node_type_static_reduce_max:
case xnn_node_type_static_reduce_min:
case xnn_node_type_static_resize_bilinear_2d: {
const struct xnn_value* input_value = &subgraph->values[node->inputs[0]];
if (!xnn_value_is_const(input_value->flags)) {
break;
}
for (int k = 0; k < node->num_outputs; k++) {
struct xnn_value* output_value = &subgraph->values[node->outputs[k]];
output_value->flags |= (input_value->flags & (XNN_VALUE_FLAG_IS_ZERO |
XNN_VALUE_FLAG_IS_ONE));
}
} break;
// Operations that propagate some constants.
case xnn_node_type_static_sum_squared:
case xnn_node_type_static_sum:
if (subgraph->values[node->inputs[0]].flags & XNN_VALUE_FLAG_IS_ZERO) {
subgraph->values[node->outputs[0]].flags |= XNN_VALUE_FLAG_IS_ZERO;
}
break;
// Different unary ops propagate constants in different ways.
case xnn_node_type_unary_elementwise: {
const struct xnn_value* input_value = &subgraph->values[node->inputs[0]];
if (!xnn_value_is_const(input_value->flags)) {
break;
}
struct xnn_value* output_value = &subgraph->values[node->outputs[0]];
switch (node->unary_operator) {
// The following ops preserve both zeros and ones.
case xnn_unary_abs:
case xnn_unary_bankers_rounding:
case xnn_unary_ceiling:
case xnn_unary_convert:
case xnn_unary_cube_root:
case xnn_unary_elu:
case xnn_unary_floor:
case xnn_unary_leaky_relu:
case xnn_unary_popcount:
case xnn_unary_sign:
case xnn_unary_square_root:
case xnn_unary_square:
output_value->flags |= (input_value->flags & (XNN_VALUE_FLAG_IS_ZERO |
XNN_VALUE_FLAG_IS_ONE));
break;
// The following ops preserve only zeros.
case xnn_unary_hardswish:
case xnn_unary_negate:
case xnn_unary_sine:
case xnn_unary_tanh:
output_value->flags |= (input_value->flags & XNN_VALUE_FLAG_IS_ZERO);
break;
// The following ops preserve only ones.
case xnn_unary_reciprocal_square_root:
output_value->flags |= (input_value->flags & XNN_VALUE_FLAG_IS_ZERO);
break;
// The following ops flip zeros to ones.
case xnn_unary_exp:
case xnn_unary_cosine:
if (input_value->flags & XNN_VALUE_FLAG_IS_ZERO) {
output_value->flags |= XNN_VALUE_FLAG_IS_ONE;
}
break;
// The following ops flip ones to zero.
case xnn_unary_log:
if (input_value->flags & XNN_VALUE_FLAG_IS_ONE) {
output_value->flags |= XNN_VALUE_FLAG_IS_ZERO;
}
break;
// Clamps preserve zeros and/or ones if they are in the valid range.
case xnn_unary_clamp:
if (node->params.unary.clamp.min <= 0.0f &&
node->params.unary.clamp.max >= 0.0f) {
output_value->flags |=
(input_value->flags & XNN_VALUE_FLAG_IS_ZERO);
}
if (node->params.unary.clamp.min <= 1.0f &&
node->params.unary.clamp.max >= 1.0f) {
output_value->flags |= (input_value->flags & XNN_VALUE_FLAG_IS_ONE);
}
break;
// Operations that don't preserve constant zeros or ones.
case xnn_unary_approxgelu:
case xnn_unary_bitwise_not:
case xnn_unary_count_leading_zeros:
case xnn_unary_gelu:
case xnn_unary_sigmoid:
break;
case xnn_unary_invalid:
XNN_UNREACHABLE;
}
} break;
// Different binary ops propagate constants in different ways.
case xnn_node_type_binary_elementwise: {
const struct xnn_value* input_a_value =
&subgraph->values[node->inputs[0]];
const struct xnn_value* input_b_value =
&subgraph->values[node->inputs[1]];
if (!xnn_value_is_const(input_a_value->flags) &&
!xnn_value_is_const(input_b_value->flags)) {
break;
}
const bool a_is_zero = input_a_value->flags & XNN_VALUE_FLAG_IS_ZERO;
const bool a_is_one = input_a_value->flags & XNN_VALUE_FLAG_IS_ONE;
const bool b_is_zero = input_b_value->flags & XNN_VALUE_FLAG_IS_ZERO;
const bool b_is_one = input_b_value->flags & XNN_VALUE_FLAG_IS_ONE;
struct xnn_value* output_value = &subgraph->values[node->outputs[0]];
switch (node->binary_operator) {
case xnn_binary_add:
if (a_is_zero && b_is_zero) {
output_value->flags |= XNN_VALUE_FLAG_IS_ZERO;
} else if (a_is_zero && b_is_one) {
output_value->flags |= XNN_VALUE_FLAG_IS_ONE;
} else if (a_is_one && b_is_zero) {
output_value->flags |= XNN_VALUE_FLAG_IS_ONE;
}
break;
case xnn_binary_subtract:
if (input_a_value == input_b_value) {
output_value->flags |= XNN_VALUE_FLAG_IS_ZERO;
} else if (a_is_zero && b_is_zero) {
output_value->flags |= XNN_VALUE_FLAG_IS_ZERO;
} else if (a_is_one && b_is_zero) {
output_value->flags |= XNN_VALUE_FLAG_IS_ONE;
} else if (a_is_one && b_is_one) {
output_value->flags |= XNN_VALUE_FLAG_IS_ZERO;
}
break;
case xnn_binary_multiply:
if (a_is_zero || b_is_zero) {
output_value->flags |= XNN_VALUE_FLAG_IS_ZERO;
} else if (a_is_one && b_is_one) {
output_value->flags |= XNN_VALUE_FLAG_IS_ONE;
}
break;
case xnn_binary_divide:
if (input_a_value == input_b_value) {
output_value->flags |= XNN_VALUE_FLAG_IS_ONE;
} else if (a_is_zero) {
output_value->flags |= XNN_VALUE_FLAG_IS_ZERO;
} else if (a_is_one && b_is_one) {
output_value->flags |= XNN_VALUE_FLAG_IS_ONE;
}
break;
case xnn_binary_bitwise_or:
case xnn_binary_maximum:
if (a_is_zero && b_is_zero) {
output_value->flags |= XNN_VALUE_FLAG_IS_ZERO;
} else if (a_is_zero && b_is_one) {
output_value->flags |= XNN_VALUE_FLAG_IS_ONE;
} else if (a_is_one && b_is_zero) {
output_value->flags |= XNN_VALUE_FLAG_IS_ONE;
} else if (a_is_one && b_is_one) {
output_value->flags |= XNN_VALUE_FLAG_IS_ONE;
}
break;
case xnn_binary_bitwise_and:
case xnn_binary_minimum:
if (a_is_zero && b_is_zero) {
output_value->flags |= XNN_VALUE_FLAG_IS_ZERO;
} else if (a_is_zero && b_is_one) {
output_value->flags |= XNN_VALUE_FLAG_IS_ZERO;
} else if (a_is_one && b_is_zero) {
output_value->flags |= XNN_VALUE_FLAG_IS_ZERO;
} else if (a_is_one && b_is_one) {
output_value->flags |= XNN_VALUE_FLAG_IS_ONE;
}
break;
case xnn_binary_bitwise_xor:
case xnn_binary_squared_difference:
if (a_is_zero && b_is_zero) {
output_value->flags |= XNN_VALUE_FLAG_IS_ZERO;
} else if (a_is_zero && b_is_one) {
output_value->flags |= XNN_VALUE_FLAG_IS_ONE;
} else if (a_is_one && b_is_zero) {
output_value->flags |= XNN_VALUE_FLAG_IS_ONE;
} else if (a_is_one && b_is_one) {
output_value->flags |= XNN_VALUE_FLAG_IS_ZERO;
}
break;
case xnn_binary_modulus:
case xnn_binary_prelu:
if (a_is_zero) {
output_value->flags |= XNN_VALUE_FLAG_IS_ZERO;
} else if (a_is_one) {
output_value->flags |= XNN_VALUE_FLAG_IS_ONE;
}
break;
case xnn_binary_copysign:
if (a_is_zero && (b_is_zero || b_is_one)) {
output_value->flags |= XNN_VALUE_FLAG_IS_ZERO;
} else if (a_is_one && (b_is_zero || b_is_one)) {
output_value->flags |= XNN_VALUE_FLAG_IS_ONE;
}
break;
case xnn_binary_pow:
if (b_is_zero) {
output_value->flags |= XNN_VALUE_FLAG_IS_ONE;
} else if (a_is_zero) {
output_value->flags |= XNN_VALUE_FLAG_IS_ZERO;
} else if (a_is_one) {
output_value->flags |= XNN_VALUE_FLAG_IS_ONE;
}
break;
case xnn_binary_shift_right_arithmetic:
case xnn_binary_shift_right_logical:
case xnn_binary_shift_left:
if (a_is_zero) {
output_value->flags |= XNN_VALUE_FLAG_IS_ZERO;
}
break;
// Operations that don't preserve constant zeros or ones.
case xnn_binary_atan2:
break;
case xnn_binary_invalid:
XNN_UNREACHABLE;
}
} break;
// Node types that don't preserve constant zeros or ones.
case xnn_node_type_argmax_pooling_2d:
case xnn_node_type_average_pooling_2d:
case xnn_node_type_batch_matrix_multiply:
case xnn_node_type_concatenate:
case xnn_node_type_convolution_2d:
case xnn_node_type_deconvolution_2d:
case xnn_node_type_depth_to_space_2d:
case xnn_node_type_depthwise_convolution_2d:
case xnn_node_type_fully_connected_sparse:
case xnn_node_type_fully_connected:
case xnn_node_type_global_average_pooling_1d:
case xnn_node_type_global_average_pooling_2d:
case xnn_node_type_global_sum_pooling_1d:
case xnn_node_type_global_sum_pooling_2d:
case xnn_node_type_invalid:
case xnn_node_type_max_pooling_2d:
case xnn_node_type_pack_lh:
case xnn_node_type_rope:
case xnn_node_type_softmax:
case xnn_node_type_space_to_depth_2d:
case xnn_node_type_static_constant_pad:
case xnn_node_type_unpooling_2d:
break;
}
}
// Replaces `add(a, neg(b))`, `add(neg(a), b)`, or `sub(a, neg(b))`
// with `sub(a, b)`, `sub(b, a)`, or `add(a, b)`, respectively, and
// `{mul,div}(neg(a), neg(b))` with `{mul,div}(a, b)`.
static enum xnn_status optimize_common_subgraphs_simplify_binary_neg(
xnn_subgraph_t subgraph, uint32_t node_id, size_t* changes) {
XNN_RETURN_IF_ERROR(xnn_subgraph_reserve_nodes(subgraph, 1));
struct xnn_node* node = &subgraph->nodes[node_id];
if (node->type != xnn_node_type_binary_elementwise ||
node->flags & XNN_NODE_FLAG_DONT_ELIDE) {
return xnn_status_success;
}
struct xnn_value* input_a_value = &subgraph->values[node->inputs[0]];
struct xnn_value* input_b_value = &subgraph->values[node->inputs[1]];
const uint32_t producer_a_id = input_a_value->producer;
const uint32_t producer_b_id = input_b_value->producer;
struct xnn_node* producer_a_node = producer_a_id == XNN_INVALID_NODE_ID
? NULL
: &subgraph->nodes[producer_a_id];
struct xnn_node* producer_b_node = producer_b_id == XNN_INVALID_NODE_ID
? NULL
: &subgraph->nodes[producer_b_id];
const bool producer_a_is_neg =
producer_a_node &&
producer_a_node->type == xnn_node_type_unary_elementwise &&
producer_a_node->unary_operator == xnn_unary_negate;
const bool producer_b_is_neg =
producer_b_node &&
producer_b_node->type == xnn_node_type_unary_elementwise &&
producer_b_node->unary_operator == xnn_unary_negate;
// Check for `sub(a, neg(b))`.
if (node->binary_operator == xnn_binary_subtract && producer_b_is_neg) {
XNN_RETURN_IF_ERROR(
xnn_define_binary(subgraph, xnn_binary_add, /*params=*/NULL,
node->inputs[0],
subgraph->nodes[producer_b_id].inputs[0],
node->outputs[0], node->flags),
"Failed to create new `Add` node.");
node = move_last_node_to(subgraph, node_id);
xnn_log_info(
"Converted sub[#%u](v%03u, neg[#%u](v%03u)) to "
"add[#%u](v%03u, v%03u).",
node->id, node->inputs[0], producer_b_id, node->inputs[1], node->id,
node->inputs[0], node->inputs[1]);
(*changes)++;
}
// Check for `add(a, neg(b))` or `add(neg(a), b)`.
else if (node->binary_operator == xnn_binary_add &&
(producer_a_is_neg || producer_b_is_neg)) {
if (producer_a_is_neg) {
XNN_RETURN_IF_ERROR(
xnn_define_binary(subgraph, xnn_binary_subtract, /*params=*/NULL,
node->inputs[1], producer_a_node->inputs[0],
node->outputs[0], node->flags),
"Failed to create new `Subtract` node.");
xnn_log_info(
"Converted add[#%u](neg[#%u](v%03u), v%03u) to "
"sub[#%u](v%03u, v%03u).",
node->id, producer_a_id, subgraph->nodes[producer_a_id].inputs[0],
input_b_value->id, node->id, input_b_value->id,
subgraph->nodes[producer_a_id].inputs[0]);
} else if (producer_b_is_neg) {
XNN_RETURN_IF_ERROR(
xnn_define_binary(subgraph, xnn_binary_subtract, /*params=*/NULL,
node->inputs[0], producer_b_node->inputs[0],
node->outputs[0], node->flags),
"Failed to create new `Subtract` node.");
}
node = move_last_node_to(subgraph, node_id);
xnn_log_info(
"Converted add[#%u](v%03u, neg[#%u](v%03u)) to "
"sub[#%u](v%03u, v%03u).",
node->id, node->inputs[0],
producer_a_is_neg ? producer_a_id : producer_b_id, node->inputs[1],
node->id, node->inputs[0], node->inputs[1]);
(*changes)++;
}
// Check for `mul(neg(a), neg(b))` or `div(neg(a), neg(b))`.
else if ((node->binary_operator == xnn_binary_multiply ||
node->binary_operator == xnn_binary_divide) &&
producer_a_is_neg && producer_b_is_neg) {
XNN_RETURN_IF_ERROR(
xnn_define_binary(subgraph, node->binary_operator, /*params=*/NULL,
producer_a_node->inputs[0],
producer_b_node->inputs[0], node->outputs[0],
node->flags),
"Failed to create new %s node.", xnn_node_type_to_string(node->type));
node = move_last_node_to(subgraph, node_id);
xnn_log_info(
"Converted %s[#%u](neg[#%u](v%03u), neg[#%u](v%03u)) to "
"%s[#%u](v%03u, v%03u).",
node->binary_operator == xnn_binary_multiply ? "mul" : "div", node->id,
producer_a_id, node->inputs[0], producer_b_id, node->inputs[1],
node->binary_operator == xnn_binary_multiply ? "mul" : "div", node->id,
node->inputs[0], node->inputs[1]);
(*changes)++;
}
return xnn_status_success;
}
// Replace `mul(x, 1.0)`, `div(div, 1.0)`, `add(x, 0.0)`, or `sub(x, 0.0)` with
// just `x`, where possible.
static enum xnn_status optimize_common_subgraphs_binary_const_noop(
xnn_subgraph_t subgraph, uint32_t node_id, size_t* changes) {
XNN_RETURN_IF_ERROR(xnn_subgraph_reserve_nodes(subgraph, 1));
struct xnn_node* node = &subgraph->nodes[node_id];
if (node->type != xnn_node_type_binary_elementwise ||
(node->binary_operator != xnn_binary_multiply &&
node->binary_operator != xnn_binary_divide &&
node->binary_operator != xnn_binary_add &&
node->binary_operator != xnn_binary_subtract) ||
node->flags & XNN_NODE_FLAG_DONT_ELIDE) {
return xnn_status_success;
}
struct xnn_value* input_value = &subgraph->values[node->inputs[0]];
struct xnn_value* const_value = &subgraph->values[node->inputs[1]];
// If the output value is already a constant, do nothing.
struct xnn_value* output_value = &subgraph->values[node->outputs[0]];
if (xnn_value_is_const(output_value->flags)) {
return xnn_status_success;
}
// One of the two inputs should be a constant.
if (!xnn_value_is_const(const_value->flags)) {
if (xnn_value_is_const(input_value->flags)) {
swap_value_pointers(&input_value, &const_value);
} else {
return xnn_status_success;
}
}
// If the constant shape isn't static, there could be a broadcast that
// prevents us from knowing the output shape.
if ((const_value->flags & XNN_VALUE_FLAG_SHAPE_IS_STATIC) == 0) {
return xnn_status_success;
}
const enum xnn_binary_operator binary_operator = node->binary_operator;
const bool const_is_zero = (const_value->flags & XNN_VALUE_FLAG_IS_ZERO) != 0;
const bool const_is_one = (const_value->flags & XNN_VALUE_FLAG_IS_ONE) != 0;
const bool const_is_rhs = node->inputs[1] == const_value->id;
// If this is a `sub(0.0, x)`, replace it with a `neg(x)`.
if (const_is_zero && binary_operator == xnn_binary_subtract &&
node->inputs[0] == const_value->id) {
XNN_RETURN_IF_ERROR(
xnn_define_unary(subgraph, xnn_unary_negate, /*params=*/NULL,
input_value->id, node->outputs[0], node->flags),
"Failed to create new `Unary negate` node.");
node = move_last_node_to(subgraph, node_id);
xnn_log_info("Replaced sub[#%u](0.0, v%03u) with neg[#%u](v%03u).", node_id,
input_value->id, node_id, input_value->id);
(*changes)++;
}
// Otherwise, if this is a `mul(x, 1.0)`, `div(x, 1.0)`, `add(x, 0.0)`, or
// `sub(x, 0.0)`, then skip this op.
else if ((const_is_one &&
(binary_operator == xnn_binary_multiply ||
(binary_operator == xnn_binary_divide && const_is_rhs))) ||
(const_is_zero &&
(binary_operator == xnn_binary_add ||
(binary_operator == xnn_binary_subtract && const_is_rhs)))) {
// We can elide the operation if the constant is a scalar: the scalar is
// broadcasted to the shape of the input and the output has the same shape.
const bool constant_is_scalar =
xnn_shape_multiply_all_dims(&const_value->shape) == 1 &&
const_value->shape.num_dims <= input_value->shape.num_dims;
// We can elide the operation if the input value and the output value have
// the same shape.
const bool input_and_output_have_the_same_shape =
(input_value->flags & XNN_VALUE_FLAG_SHAPE_IS_STATIC) &&
xnn_shape_match(&input_value->shape, &output_value->shape);
if (constant_is_scalar || input_and_output_have_the_same_shape) {
// We try to elide the operation...
if (short_circuit(subgraph, input_value->id, node->outputs[0])) {
xnn_log_info("Elided spurious %s[#%u](v%03u, %s).",
binary_operator == xnn_binary_multiply ? "mul"
: binary_operator == xnn_binary_divide ? "div"
: binary_operator == xnn_binary_add ? "add"
: "sub",
node->id, input_value->id, const_is_zero ? "0.0" : "1.0");
xnn_node_clear(node);
(*changes)++;
} else {
// ... if it fails, input and output have the same shape so we can copy.
XNN_RETURN_IF_ERROR(xnn_define_copy(subgraph, input_value->id,
node->outputs[0], node->flags),
"Failed to create new `Copy` node.");
node = move_last_node_to(subgraph, node_id);
xnn_log_info(
"Replaced spurious %s[#%u](v%03u, %s) with copy[#%u](v%03u).",
binary_operator == xnn_binary_multiply ? "mul"
: binary_operator == xnn_binary_divide ? "div"
: binary_operator == xnn_binary_add ? "add"
: "sub",
node->id, input_value->id, const_is_zero ? "0.0" : "1.0", node_id,
input_value->id);
(*changes)++;
}
}
}
return xnn_status_success;
}
// Merge or remove transposes of the RHS of a batch-matrix-multiply or
// fully-connected op.
static enum xnn_status optimize_common_subgraphs_gemm_rhs_transpose(
xnn_subgraph_t subgraph, uint32_t node_id, size_t* changes) {
XNN_RETURN_IF_ERROR(xnn_subgraph_reserve_nodes(subgraph, 2));
struct xnn_node* node = &subgraph->nodes[node_id];
if (node->type != xnn_node_type_fully_connected &&
node->type != xnn_node_type_batch_matrix_multiply) {
return xnn_status_success;
}
// Chech whether the RHS is produced by a `transpose` node.
uint32_t weights_id = node->inputs[1];
struct xnn_value* weights_value = &subgraph->values[weights_id];
if (weights_value->producer == XNN_INVALID_NODE_ID ||
subgraph->nodes[weights_value->producer].type !=
xnn_node_type_static_transpose ||
weights_value->num_consumers != 1) {
return xnn_status_success;
}
// Is this the bad packing (unoptimized packing kernels)?
const bool is_gio = (node->type == xnn_node_type_fully_connected &&
(node->flags & XNN_FLAG_TRANSPOSE_WEIGHTS) != 0) ||
(node->type == xnn_node_type_batch_matrix_multiply &&
(node->flags & XNN_FLAG_TRANSPOSE_WEIGHTS) == 0);
// Check whether the transpose only affects the last two dimensions.
const uint32_t transpose_node_id = weights_value->producer;
struct xnn_node* transpose_node = &subgraph->nodes[transpose_node_id];
const size_t num_perm_dims = transpose_node->params.transpose.num_dims;
size_t* perm = transpose_node->params.transpose.perm;
bool is_last_dims = true;
for (int k = 0; is_last_dims && k < num_perm_dims - 2; k++) {
is_last_dims &= perm[k] == k;
}
is_last_dims &= (perm[num_perm_dims - 2] == num_perm_dims - 1) &&
(perm[num_perm_dims - 1] == num_perm_dims - 2);
// If the input is `gio` or the transpose only affects the last two dims, flip
// the last two dimensions of the transpose.
if (!is_gio && !is_last_dims) {
return xnn_status_success;
}
// If the transpose consists of only the last two dimensions, we can skip it.
if (is_last_dims) {
weights_id = transpose_node->inputs[0];
xnn_log_info("Skipping elided static_transpose[#%u](v%03u).",
transpose_node_id, weights_id);
} else {
size_t new_perm[XNN_MAX_TENSOR_DIMS] = {0};
for (int k = 0; k + 2 < num_perm_dims; k++) {
new_perm[k] = perm[k];
}
new_perm[num_perm_dims - 2] = perm[num_perm_dims - 1];
new_perm[num_perm_dims - 1] = perm[num_perm_dims - 2];
const struct xnn_shape* transpose_input_shape =
&subgraph->values[transpose_node->inputs[0]].shape;
weights_value->shape.num_dims = num_perm_dims;
for (int k = 0; k < num_perm_dims; k++) {
weights_value->shape.dim[k] = transpose_input_shape->dim[perm[k]];
}
XNN_RETURN_IF_ERROR(xnn_define_static_transpose(
subgraph, num_perm_dims, new_perm, transpose_node->inputs[0],
transpose_node->outputs[0], transpose_node->flags));
transpose_node = move_last_node_to(subgraph, transpose_node_id);
node = &subgraph->nodes[node_id];
}
// Flip the "transpose" flag of the node. Note that this flag has opposite
// meanings for fully-connected and batch-matrix-multiply nodes, so instead of
// setting the value explicitly, we just flip whatever value was previously
// set.
const uint32_t node_flags = node->flags ^ XNN_FLAG_TRANSPOSE_WEIGHTS;
XNN_RETURN_IF_ERROR(
node->type == xnn_node_type_fully_connected
? xnn_define_fully_connected(subgraph, node->activation.output_min,
node->activation.output_max,
node->inputs[0], weights_id,
node->inputs[2], node->outputs[0],
node_flags)
: xnn_define_batch_matrix_multiply(
subgraph, node->inputs[0], weights_id, node->outputs[0],
node_flags));
node = move_last_node_to(subgraph, node_id);
xnn_log_info(
"Converted %s[#%u](v%03u, static_transpose[#%u](v%03u), %s) to "
"%s[#%u](v%03u, v%03u, %s).",
node->type == xnn_node_type_fully_connected ? "fully_connected"
: "batch_matrix_multiply",
node_id, node->inputs[0], transpose_node_id, weights_id,
is_gio ? "gio" : "goi",
node->type == xnn_node_type_fully_connected ? "fully_connected"
: "batch_matrix_multiply",
node_id, node->inputs[0], node->inputs[1],
is_gio ^ !is_last_dims ? "gio" : "goi");
(*changes)++;
return xnn_status_success;
}
static enum xnn_status optimize_common_subgraphs_iter(
xnn_subgraph_t subgraph, uint32_t optimization_flags, size_t* changes) {
// Loop over the nodes in this subgraph.
for (uint32_t node_id = 0; node_id < subgraph->num_nodes; node_id++) {
struct xnn_node* node = &subgraph->nodes[node_id];
if (node->type == xnn_node_type_invalid) {
continue;
}
// Propagate static shapes across this node.
bool all_input_shapes_are_static = true;
for (int k = 0; k < node->num_inputs && all_input_shapes_are_static; k++) {
all_input_shapes_are_static &= (subgraph->values[node->inputs[k]].flags &
XNN_VALUE_FLAG_SHAPE_IS_STATIC) != 0;
}
if (all_input_shapes_are_static) {
// Propagate static shapes for nodes for which we know how to do so.
switch (node->type) {
case xnn_node_type_unary_elementwise:
// Unary elementwise ops output always have the same shape as their
// input.
subgraph->values[node->outputs[0]].shape =
subgraph->values[node->inputs[0]].shape;
subgraph->values[node->outputs[0]].flags |=
XNN_VALUE_FLAG_SHAPE_IS_STATIC;
break;
case xnn_node_type_binary_elementwise:
// Compute the broadcasted output size of binary elementwise ops.
xnn_shape_binary_broadcast(&subgraph->values[node->inputs[0]].shape,
&subgraph->values[node->inputs[1]].shape,
&subgraph->values[node->outputs[0]].shape);
subgraph->values[node->outputs[0]].flags |=
XNN_VALUE_FLAG_SHAPE_IS_STATIC;
break;
default:
// We don't reshape outputs and don't mark their shapes as static.
break;
}
// TODO: b/455537016 - expand to other know node types.
}
// Propagate constants through constant-preserving ops.
propagate_constants(subgraph, node_id);
switch (node->type) {
case xnn_node_type_binary_elementwise:
// Replace `mul(x, x)` with `sqr(x)` for consistency.
XNN_RETURN_IF_ERROR(
optimize_common_subgraphs_mul_to_sqr(subgraph, node_id, changes));
// Replace `div(x, x)` or `sub(x, x)` with a constant `1.0` or `0.0`,
// respectively.
// TODO(b/460807548): Disabled due to bugs, complexity, and we don't
// have a use case for these yet.
// XNN_RETURN_IF_ERROR(optimize_common_subgraphs_binary_to_const(
// subgraph, node_id, changes));
// Widen fp16 accumulators for `mul(reduce_sum(x), y)` or
// `div(reduce_sum(x), y)`. This is a bit of a hack to make subgraphs
// rewritten to be fp16 less likely to overflow. Especially if x is a
// squaring operation, and the data is fp16, it is very likely that the
// sum will overflow.
XNN_RETURN_IF_ERROR(
widen_fp16_accumulators(subgraph, node_id, changes));
// Convert min/max operations with a single static value to a unary
// `clamp` node.
XNN_RETURN_IF_ERROR(optimize_common_subgraphs_min_max_to_clamp(
subgraph, node_id, changes));
// Replaces `add(a, neg(b))`, `add(neg(a), b)`, or `sub(a, neg(b))`
// with `sub(a, b)`, `sub(b, a)`, or `add(a, b)`, respectively, and
// `{mul,div}(neg(a), neg(b))` with `{mul,div}(a, b)`.
// TODO(b/460807548): Disabled due to bugs, complexity, and we don't
// have a use case for these yet.
// XNN_RETURN_IF_ERROR(optimize_common_subgraphs_simplify_binary_neg(
// subgraph, node_id, changes));
// Replace `mul(x, 1.0)`, `div(div, 1.0)`, `add(x, 0.0)`, or `sub(x,
// 0.0)` with just `x`, where possible.
// TODO(b/460807548): Disabled due to bugs, complexity, and we don't
// have a use case for these yet.
// XNN_RETURN_IF_ERROR(optimize_common_subgraphs_binary_const_noop(
// subgraph, node_id, changes));
break;
case xnn_node_type_unary_elementwise:
// Fold unary clamp operations into the previous operator's
// activation, if possible.
XNN_RETURN_IF_ERROR(
optimize_common_subgraphs_merge_clamps(subgraph, node_id, changes));
// Remove spurious unary clamp operations.
XNN_RETURN_IF_ERROR(optimize_common_subgraphs_spurious_clamps(
subgraph, node_id, changes));
break;
case xnn_node_type_static_sum:
case xnn_node_type_static_mean:
// Convert `reduce_sum(sqr(a))` or `reduce_mean(sqr(a))` to
// `reduce_sum_squared(a)` or `reduce_mean_squared(a)`, respectively.
XNN_RETURN_IF_ERROR(optimize_common_subgraphs_reduce_sum_to_square(
subgraph, node_id, changes));
break;
case xnn_node_type_static_broadcast:
XNN_RETURN_IF_ERROR(
optimize_common_subgraphs_broadcast(subgraph, node_id, changes));
break;
case xnn_node_type_static_reshape:
// If the reshape is fully defined (no zeros), then the output shape
// is static.
if (xnn_shape_multiply_all_dims(
&node->params.static_reshape.new_shape) != 0) {
xnn_log_debug(
"Marking output of static_reshape[#%u](v%03u) as static shaped.",
node->id, node->inputs[0]);
subgraph->values[node->outputs[0]].shape =
node->params.static_reshape.new_shape;
subgraph->values[node->outputs[0]].flags |=
XNN_VALUE_FLAG_SHAPE_IS_STATIC;
} else {
// If the output shape isn't static, then there's nothing to optimize.
continue;
}
XNN_FALLTHROUGH
case xnn_node_type_static_expand_dims:
// Merge `reshape(expand_dims(...))` and `expand_dims(reshape(...))`
// into a single `reshape`.
XNN_RETURN_IF_ERROR(optimize_common_subgraphs_merge_reshapes(
subgraph, node_id, changes));
// Apply `static_reshape` and `static_expand_dims` to static values
// directly.
XNN_RETURN_IF_ERROR(optimize_common_subgraphs_static_reshapes(
subgraph, node_id, changes));
// TODO: b/455537016 - Eventually we should also deal with cases such as
// `static_value` -> `unary` -> `static_reshape` where the reshape can
// be pushed back to the static value.
break;
case xnn_node_type_fully_connected:
case xnn_node_type_batch_matrix_multiply:
// Merge or remove transposes of the RHS of a batch-matrix-multiply or
// fully-connected op.
XNN_RETURN_IF_ERROR(optimize_common_subgraphs_gemm_rhs_transpose(
subgraph, node_id, changes));
break;
default:
break;
}
}
return xnn_status_success;
}
enum xnn_status xnn_subgraph_optimize_common_subgraphs(
xnn_subgraph_t subgraph, uint32_t optimization_flags) {
// If we shouldn't change the numerics, then don't do anything.
if (optimization_flags & XNN_FLAG_SLOW_CONSISTENT_ARITHMETIC ||
optimization_flags & XNN_FLAG_NO_OPERATOR_FUSION) {
return xnn_status_success;
}
// Mark static values as `XNN_VALUE_SHAPE_IS_STATIC`, and scalar static values
// as `XNN_VALUE_FLAG_IS_ZERO` or `XNN_VALUE_FLAG_IS_ONE`, where appropriate.
for (uint32_t value_id = 0; value_id < subgraph->num_values; value_id++) {
struct xnn_value* value = &subgraph->values[value_id];
if (value->datatype == xnn_datatype_invalid) {
continue;
}
// Static values have static shapes.
else if (xnn_value_is_static(value->allocation_type)) {
value->flags |= XNN_VALUE_FLAG_SHAPE_IS_STATIC;
// Is this static value scalar?
if (xnn_shape_multiply_all_dims(&value->shape) == 1) {
// Get the value as a float.
const float value_as_float = get_scalar_value_as_float(value);
xnn_log_debug("v%03u is a constant: %e.", value->id, value_as_float);
// Mark the value accordingly.
value->flags |= (value_as_float == 0.0f) ? XNN_VALUE_FLAG_IS_ZERO
: (value_as_float == 1.0f) ? XNN_VALUE_FLAG_IS_ONE
: 0;
}
}
}
while (true) {
// Count the number of changes made.
size_t changes = 0;
XNN_RETURN_IF_ERROR(
optimize_common_subgraphs_iter(subgraph, optimization_flags, &changes));
// Clean up after ourselves.
if (changes) {
xnn_subgraph_clean_up(subgraph);
} else {
break;
}
}
return xnn_status_success;
}
enum xnn_status xnn_subgraph_optimize_packed_lhs(xnn_subgraph_t subgraph,
uint32_t optimization_flags) {
// Count the number of changes made.
size_t changes = 0;
// Loop over the nodes in the subgraph.
for (uint32_t node_id = 0; node_id < subgraph->num_nodes; node_id++) {
struct xnn_node* node = &subgraph->nodes[node_id];
// Skip anything that is not a fully-connected node.
switch (node->type) {
case xnn_node_type_batch_matrix_multiply:
case xnn_node_type_fully_connected: {
// Get a handle on the inputs/outputs.
const uint32_t input_id = node->inputs[0];
struct xnn_value* input_value = &subgraph->values[input_id];
struct xnn_value* kernel_value = &subgraph->values[node->inputs[1]];
struct xnn_value* output_value = &subgraph->values[node->outputs[0]];
const enum xnn_datatype input_datatype = input_value->datatype;
const enum xnn_datatype kernel_datatype = kernel_value->datatype;
const enum xnn_datatype output_datatype = output_value->datatype;
// Check if we have a packed GEMM config for the combination of
// input/kernel/output.
const struct xnn_gemm_config* gemm_config = NULL;
enum xnn_datatype assumed_datatype = xnn_datatype_invalid;
switch (input_datatype) {
case xnn_datatype_fp16:
if (input_datatype == output_datatype &&
kernel_datatype == xnn_datatype_fp16) {
if ((gemm_config = xnn_init_pf16_gemm_config())) {
assumed_datatype = xnn_datatype_pfp16;
}
}
break;
case xnn_datatype_fp32:
if (input_datatype == output_datatype &&
(kernel_datatype == xnn_datatype_fp32 ||
kernel_datatype == xnn_datatype_fp16)) {
if ((gemm_config = xnn_init_pf32_gemm_config())) {
assumed_datatype = xnn_datatype_pfp32;
}
}
break;
case xnn_datatype_qint8:
if (input_datatype == output_datatype &&
kernel_datatype == xnn_datatype_qcint8) {
if ((gemm_config = xnn_init_pqs8_qc8w_gemm_config())) {
assumed_datatype = xnn_datatype_pqint8;
}
}
break;
case xnn_datatype_qdint8:
// We may inline the `qdint8` packing regardless of whether we have
// a specialized `qpint8` kernel or not.
assumed_datatype = xnn_datatype_qdint8;
if (output_datatype == xnn_datatype_fp32) {
switch (kernel_datatype) {
case xnn_datatype_qbint4:
// The qp8_f32_qb4w kernels only support unsigned 4-bit
// weights.
if (kernel_value->quantization.zero_point == 8 &&
(gemm_config = xnn_init_qp8_f32_qb4w_gemm_config())) {
assumed_datatype = xnn_datatype_qpint8;
}
break;
case xnn_datatype_qcint4:
if ((gemm_config = xnn_init_qp8_f32_qc4w_gemm_config())) {
assumed_datatype = xnn_datatype_qpint8;
}
break;
case xnn_datatype_qcint8:
if ((gemm_config = xnn_init_qp8_f32_qc8w_gemm_config())) {
assumed_datatype = xnn_datatype_qpint8;
}
break;
default:
break;
}
}
break;
default:
// If none of the above happened, do nothing for this node.
continue;
}
if (assumed_datatype != xnn_datatype_invalid) {
if (optimization_flags & XNN_FLAG_NO_INLINED_LHS_PACKING) {
if (assumed_datatype == xnn_datatype_qdint8) {
// If the input is already `qdint8`, don't do anything different.
continue;
} else if (assumed_datatype == xnn_datatype_qpint8) {
xnn_log_debug(
// `qpint8` inputs are generated by modifying the `convert` op
// that generated the `qdint8` input, so we only have to add a
// `pack-lh` op for the other input types.
"Coercing type of input ID #%" PRIu32
" of %s node from `%s` to `%s`.",
input_id, xnn_node_type_to_string(xnn_node_type_convert),
xnn_datatype_to_string(input_datatype),
xnn_datatype_to_string(xnn_datatype_qpint8));
subgraph->values[input_id].datatype = assumed_datatype;
subgraph->values[input_id].gemm_config = gemm_config;
} else {
// Insert a node to pack the LHS.
xnn_log_debug(
"Adding %s node for input ID #%" PRIu32
" of type `%s` for %s node.",
xnn_node_type_to_string(xnn_node_type_pack_lh), input_id,
xnn_node_type_to_string(xnn_node_type_fully_connected),
xnn_datatype_to_string(xnn_datatype_qpint8));
uint32_t new_id = XNN_INVALID_VALUE_ID;
XNN_RETURN_IF_ERROR(
xnn_insert_pack_lh_node(subgraph, input_id, &new_id));
subgraph->nodes[node_id].inputs[0] = new_id;
changes++;
}
// If this is a fully-connected op, we need to coerce the shape of
// the inputs from `[B, M, K]` to `[B * M, K]` to avoid batch-wise
// packing.
if (node->type == xnn_node_type_fully_connected) {
subgraph->values[subgraph->nodes[node_id].inputs[0]].flags |=
XNN_FLAG_SQUASH_GROUPS;
}
} else {
if (input_datatype == xnn_datatype_qdint8) {
// Short-circuit the inputs of the producer of the `qdint8`
// values.
struct xnn_node* producer =
&subgraph->nodes[input_value->producer];
if (producer->type != xnn_node_type_convert) {
xnn_log_error(
"Expected producer node #%u of %s tensor #%u to be of type "
"%s, but found type %s instead.",
input_value->producer,
xnn_datatype_to_string(input_datatype), input_id,
xnn_node_type_to_string(xnn_node_type_convert),
xnn_node_type_to_string(producer->type));
return xnn_status_invalid_state;
}
// Maybe use `qduint8` instead of `qdint8`?
xnn_log_debug(
"Skipping %s node #%u for input #%" PRIu32
" of node #%u (%s).",
xnn_node_type_to_string(producer->type), producer->id,
input_id, node_id,
xnn_node_type_to_string(xnn_node_type_fully_connected));
struct xnn_value* new_input =
&subgraph->values[producer->inputs[0]];
node->inputs[0] = producer->inputs[0];
if (new_input->first_consumer == input_value->producer) {
new_input->first_consumer = node_id;
}
if (--input_value->num_consumers == 0) {
xnn_node_clear(producer);
}
if (convert_gemm_to_qduint8(new_input->datatype, node->type,
kernel_datatype)) {
assumed_datatype = xnn_datatype_qduint8;
}
changes++;
}
xnn_log_debug("Setting assumed_datatype=%s for node #%u (%s).",
xnn_datatype_to_string(assumed_datatype), node_id,
xnn_node_type_to_string(node->type));
node->packed_input_datatype = assumed_datatype;
node->flags |= XNN_FLAG_INLINE_LHS_PACKING;
}
}
} break;
case xnn_node_type_convolution_2d:
case xnn_node_type_deconvolution_2d: {
// Get a handle on the inputs/outputs.
const uint32_t input_id = node->inputs[0];
struct xnn_value* input_value = &subgraph->values[input_id];
struct xnn_value* kernel_value = &subgraph->values[node->inputs[1]];
struct xnn_value* output_value = &subgraph->values[node->outputs[0]];
const enum xnn_datatype input_datatype = input_value->datatype;
const enum xnn_datatype kernel_datatype = kernel_value->datatype;
const enum xnn_datatype output_datatype = output_value->datatype;
// Check if we can do anything special with this operation.
if (input_datatype == xnn_datatype_qint8 &&
(kernel_datatype == xnn_datatype_qcint8 ||
kernel_datatype == xnn_datatype_qint8) &&
output_datatype == xnn_datatype_qint8 &&
xnn_init_pqs8_qc8w_gemm_config() != NULL &&
!(optimization_flags & XNN_FLAG_NO_INLINED_LHS_PACKING)) {
// Note that there is currently no option to not use inlining for this
// iGEMM kernel.
xnn_log_debug("Setting assumed_datatype=%s for node #%u (%s).",
xnn_datatype_to_string(xnn_datatype_pqint8), node_id,
xnn_node_type_to_string(node->type));
node->packed_input_datatype = xnn_datatype_pqint8;
node->flags |= XNN_FLAG_INLINE_LHS_PACKING;
}
if (input_datatype == xnn_datatype_fp32 &&
(kernel_datatype == xnn_datatype_fp16 ||
kernel_datatype == xnn_datatype_fp32) &&
output_datatype == xnn_datatype_fp32 &&
xnn_init_pf32_gemm_config() != NULL &&
!(optimization_flags & XNN_FLAG_NO_INLINED_LHS_PACKING)) {
// Note that there is currently no option to not use inlining for this
// iGEMM kernel.
xnn_log_debug("Setting assumed_datatype=%s for node #%u (%s).",
xnn_datatype_to_string(xnn_datatype_pfp32), node_id,
xnn_node_type_to_string(node->type));
node->packed_input_datatype = xnn_datatype_pfp32;
if(node->type == xnn_node_type_convolution_2d) {
node->flags |= XNN_FLAG_INLINE_LHS_PACKING;
}
}
if (input_datatype == xnn_datatype_fp16 &&
(kernel_datatype == xnn_datatype_fp16 ||
kernel_datatype == xnn_datatype_fp32) &&
output_datatype == xnn_datatype_fp16 &&
xnn_init_pf16_gemm_config() != NULL &&
!(optimization_flags & XNN_FLAG_NO_INLINED_LHS_PACKING)) {
// Note that there is currently no option to not use inlining for this
// iGEMM kernel.
xnn_log_debug("Setting assumed_datatype=%s for node #%u (%s).",
xnn_datatype_to_string(xnn_datatype_pfp16), node_id,
xnn_node_type_to_string(node->type));
node->packed_input_datatype = xnn_datatype_pfp16;
node->flags |= XNN_FLAG_INLINE_LHS_PACKING;
}
}
break;
default:
break;
}
}
// Second loop over the nodes to convert any `qdint8` to `qduint8` where
// appropriate.
for (uint32_t node_id = 0; node_id < subgraph->num_nodes; node_id++) {
struct xnn_node* node = &subgraph->nodes[node_id];
// Skip anything that is not a `convert` node.
if (node->type != xnn_node_type_convert) {
continue;
}
// Get a handle on the inputs/outputs.
const uint32_t output_id = node->outputs[0];
struct xnn_value* input_value = &subgraph->values[node->inputs[0]];
struct xnn_value* output_value = &subgraph->values[output_id];
const enum xnn_datatype input_datatype = input_value->datatype;
const enum xnn_datatype output_datatype = output_value->datatype;
// Only replace `qdint8` nodes for which all consumer are of the same type.
if (!output_value->all_consumers_types_same ||
output_datatype != xnn_datatype_qdint8) {
continue;
}
// Get a handle on the consumer and its inputs/outputs.
const struct xnn_node* consumer =
&subgraph->nodes[output_value->first_consumer];
const enum xnn_datatype consumer_weights_type =
subgraph->values[consumer->inputs[1]].datatype;
// If the `qduint8` config is better than the `qdint8` config, use it
// instead.
if (convert_gemm_to_qduint8(input_datatype, consumer->type,
consumer_weights_type)) {
xnn_log_debug("Coercing type of output ID #%" PRIu32
" of %s operator from `%s` to `%s`.",
output_id, xnn_node_type_to_string(xnn_node_type_convert),
xnn_datatype_to_string(output_datatype),
xnn_datatype_to_string(xnn_datatype_qduint8));
output_value->datatype = xnn_datatype_qduint8;
}
}
// Clean up after ourselves.
if (changes) {
xnn_subgraph_clean_up(subgraph);
}
return xnn_status_success;
}
// Returns true if `node` is a per-tensor `qint8 -> fp32` dequant convert node.
static bool is_qs8_to_f32_dequant(const struct xnn_subgraph* subgraph,
const struct xnn_node* node) {
if (node->type != xnn_node_type_unary_elementwise ||
node->unary_operator != xnn_unary_convert) {
return false;
}
const struct xnn_value* in = &subgraph->values[node->inputs[0]];
const struct xnn_value* out = &subgraph->values[node->outputs[0]];
return in->datatype == xnn_datatype_qint8 &&
out->datatype == xnn_datatype_fp32;
}
static bool rewrite_dequant_bmm_at(xnn_subgraph_t subgraph, uint32_t node_id) {
// Only batch-matrix-multiply nodes are rewritten here. Check the node type
// *before* reserving any space: reserving may reallocate the subgraph's node
// and value arrays, and some callers optimize a fixed-capacity scratch
// subgraph whose `nodes`/`values` storage is not heap-allocated and must not
// be reallocated (see `xnn_subgraph_fuse_unary_quantized_into_lut`).
if (subgraph->nodes[node_id].type != xnn_node_type_batch_matrix_multiply) {
return false;
}
// Reserve space upfront (we'll add 1 value + 1 node) so the later
// `xnn_subgraph_new_internal_value` / `xnn_subgraph_new_node` calls do not
// each trigger an additional reallocation or invalidate pointers.
if (xnn_subgraph_reserve_values(subgraph, 1) != xnn_status_success ||
xnn_subgraph_reserve_nodes(subgraph, 1) != xnn_status_success) {
return false;
}
struct xnn_node* node = &subgraph->nodes[node_id];
const uint32_t input_a_id = node->inputs[0];
const uint32_t input_b_id = node->inputs[1];
const uint32_t output_id = node->outputs[0];
const struct xnn_value* a_value = &subgraph->values[input_a_id];
const struct xnn_value* b_dequant_value = &subgraph->values[input_b_id];
const struct xnn_value* output_value = &subgraph->values[output_id];
if (a_value->datatype != xnn_datatype_fp32 ||
b_dequant_value->datatype != xnn_datatype_fp32 ||
output_value->datatype != xnn_datatype_fp32) {
return false;
}
// The b input must be the unique consumer of a qint8 -> fp32 convert.
const uint32_t dequant_id = b_dequant_value->producer;
if (dequant_id == XNN_INVALID_NODE_ID || dequant_id >= subgraph->num_nodes) {
return false;
}
const struct xnn_node* dequant = &subgraph->nodes[dequant_id];
if (!is_qs8_to_f32_dequant(subgraph, dequant)) {
return false;
}
if (b_dequant_value->num_consumers != 1 ||
b_dequant_value->first_consumer != node_id) {
return false;
}
const uint32_t b_qint8_id = dequant->inputs[0];
const struct xnn_value* b_qint8_value = &subgraph->values[b_qint8_id];
if (b_qint8_value->shape.num_dims < 2) {
return false;
}
if (b_qint8_value->quantization.zero_point != 0) {
return false;
}
if (a_value->shape.num_dims < 2) {
return false;
}
// Determine channel info for b. For bmm, b is `[..., K, N]` (or `[..., N,
// K]` if `XNN_FLAG_TRANSPOSE_B`). The N dimension is the qcint8 channel
// dimension. The actual per-channel scale array is materialized by the
// `convert(qint8 -> qcint8)` operator at reshape time (it broadcasts the
// input's per-tensor scale and writes the pointer onto this value's
// quantization metadata).
const bool transpose_b = (node->flags & XNN_FLAG_TRANSPOSE_B) != 0;
const size_t b_num_dims = b_qint8_value->shape.num_dims;
const size_t channel_dim = transpose_b ? b_num_dims - 2 : b_num_dims - 1;
if (b_qint8_value->shape.dim[channel_dim] == 0) {
return false;
}
const struct xnn_shape b_shape = b_qint8_value->shape;
// Create new internal value: b in qcint8. Channelwise scale array is left
// null here; `convert(qint8 -> qcint8)` populates it at reshape.
struct xnn_value* b_qcint8_value = xnn_subgraph_new_internal_value(subgraph);
if (b_qcint8_value == NULL) {
return false;
}
b_qcint8_value->type = xnn_value_type_dense_tensor;
b_qcint8_value->datatype = xnn_datatype_qcint8;
b_qcint8_value->quantization.zero_point = 0;
b_qcint8_value->quantization.channelwise_scale = NULL;
b_qcint8_value->quantization.channelwise_zero_point = NULL;
b_qcint8_value->quantization.channel_dimension = channel_dim;
b_qcint8_value->shape = b_shape;
b_qcint8_value->size = xnn_tensor_get_size(b_qcint8_value);
b_qcint8_value->allocation_type = xnn_allocation_type_workspace;
const uint32_t b_qcint8_id = b_qcint8_value->id;
struct xnn_node* qs8_to_qc8_node = xnn_subgraph_new_node(subgraph);
if (qs8_to_qc8_node == NULL) {
return false;
}
xnn_init_convert_node(qs8_to_qc8_node, b_qint8_id, b_qcint8_id,
/*flags=*/0);
node = &subgraph->nodes[node_id];
node->inputs[1] = b_qcint8_id;
xnn_log_info("Rewrote bmm Node #%" PRIu32 ": dequant(qint8 #%" PRIu32
") -> qcint8 #%" PRIu32,
node_id, b_qint8_id, b_qcint8_id);
return true;
}
enum xnn_status xnn_subgraph_rewrite_dequant_bmm(xnn_subgraph_t subgraph) {
if (subgraph->num_nodes == 0) {
return xnn_status_success;
}
xnn_subgraph_analyze_consumers_and_producers(subgraph);
bool any_changes = false;
for (uint32_t node_id = 0; node_id < subgraph->num_nodes; ++node_id) {
if (rewrite_dequant_bmm_at(subgraph, node_id)) {
any_changes = true;
}
}
if (any_changes) {
xnn_subgraph_clean_up(subgraph);
}
return xnn_status_success;
}
enum xnn_status xnn_subgraph_rewrite_for_row_sum(xnn_subgraph_t subgraph) {
// Count the number of consumers for each value.
xnn_subgraph_analyze_consumers_and_producers(subgraph);
// Loop over the nodes in the subgraph.
for (uint32_t node_id = 0; node_id < subgraph->num_nodes; node_id++) {
struct xnn_node* node = &subgraph->nodes[node_id];
// Skip anything that is not a fully-connected node.
switch (node->type) {
case xnn_node_type_fully_connected: {
// Get a handle on the inputs/outputs.
const uint32_t input_id = node->inputs[0];
struct xnn_value* input_value = &subgraph->values[input_id];
struct xnn_value* kernel_value = &subgraph->values[node->inputs[1]];
struct xnn_value* output_value = &subgraph->values[node->outputs[0]];
const enum xnn_datatype input_datatype = input_value->datatype;
const enum xnn_datatype kernel_datatype = kernel_value->datatype;
const enum xnn_datatype output_datatype = output_value->datatype;
switch (input_datatype) {
case xnn_datatype_qdint8: {
const struct xnn_gemm_config* gemm_config = NULL;
switch (kernel_datatype) {
case xnn_datatype_qcint2: {
struct xnn_node* producer =
&subgraph->nodes[input_value->producer];
if (((output_datatype == xnn_datatype_fp32) &&
(gemm_config = xnn_init_qd8_f32_qc2w_gemm_config())) ||
((output_datatype == xnn_datatype_fp16) &&
(gemm_config = xnn_init_qd8_f16_qc2w_gemm_config()))) {
if (producer->type != xnn_node_type_convert) {
xnn_log_error(
"Expected producer node #%u of %s tensor #%u to be of"
" type %s, but found type %s instead.",
input_value->producer,
xnn_datatype_to_string(input_datatype), input_id,
xnn_node_type_to_string(xnn_node_type_convert),
xnn_node_type_to_string(producer->type));
return xnn_status_invalid_state;
}
producer->flags |= XNN_NODE_FLAG_REQUIRES_ROW_SUM;
}
}
break;
default:
break;
}
break;
}
case xnn_datatype_qduint8: {
const struct xnn_gemm_config* gemm_config = NULL;
switch (kernel_datatype) {
case xnn_datatype_qcint2: {
struct xnn_node* producer =
&subgraph->nodes[input_value->producer];
if (((output_datatype == xnn_datatype_fp32) &&
(gemm_config = xnn_init_qdu8_f32_qc2w_gemm_config())) ||
((output_datatype == xnn_datatype_fp16) &&
(gemm_config = xnn_init_qdu8_f16_qc2w_gemm_config()))) {
if (producer->type != xnn_node_type_convert) {
xnn_log_error(
"Expected producer node #%u of %s tensor #%u to be of"
" type %s, but found type %s instead.",
input_value->producer,
xnn_datatype_to_string(input_datatype), input_id,
xnn_node_type_to_string(xnn_node_type_convert),
xnn_node_type_to_string(producer->type));
return xnn_status_invalid_state;
}
producer->flags |= XNN_NODE_FLAG_REQUIRES_ROW_SUM;
}
}
break;
default:
break;
}
break;
}
default:
// If none of the above happened, do nothing for this node.
continue;
}
} break;
default:
break;
}
}
return xnn_status_success;
}
static void replace_in_set(uint32_t* set, uint32_t size, uint32_t old_value,
uint32_t new_value) {
for (uint32_t i = 0; i < size; i++) {
if (set[i] == old_value) {
set[i] = new_value;
}
}
}
// Persistent values are values that can be read or written repeatedly. This
// isn't compatible with our graph model. To work around this, we implement
// persistent values with an [SSA] approach:
// - Persistent values are passed in as inputs
// - Writes to the persistent value creates a new value
// - Reads read the last value written to (or the input if none).
// - The last written value (if any) is passed out as an output.
//
// [SSA]. https://en.wikipedia.org/wiki/Static_single-assignment_form
void xnn_subgraph_rewrite_ssa(xnn_subgraph_t subgraph) {
bool* values_written =
(bool*)xnn_allocate_memory(sizeof(bool) * subgraph->num_values);
if (values_written == NULL) {
xnn_log_error("failed to allocate values_written scratch buffer");
return;
}
for (uint32_t i = 0; i < subgraph->num_values; i++) {
values_written[i] = false;
}
for (uint32_t i = 0; i < subgraph->num_nodes; i++) {
struct xnn_node* node = &subgraph->nodes[i];
for (uint32_t j = 0; j < node->num_outputs; j++) {
const uint32_t output_id = node->outputs[j];
const struct xnn_value* value = &subgraph->values[output_id];
if (!xnn_value_is_external_output(value->flags)) {
// We only care to rewrite external outputs. Internal values should
// already be SSA.
continue;
}
if (!values_written[output_id]) {
// This is the first time we've seen this output.
} else {
// We already wrote this value. Make a new value to replace the previous
// value with (so the external output remains the last write).
struct xnn_value* new_value = xnn_subgraph_new_internal_value(subgraph);
// xnn_subgraph_new_internal_value may have invalidated `value` pointer.
value = &subgraph->values[output_id];
xnn_value_copy(new_value, value);
xnn_log_debug("Adding new value #%" PRIu32
" for already produced output #%" PRIu32,
new_value->id, output_id);
// For outputs, we want only the last value to be the original output.
// The new value is not an output, and its an internal allocation.
new_value->flags &=
~(XNN_VALUE_FLAG_EXTERNAL_OUTPUT | XNN_VALUE_FLAG_EXTERNAL_INPUT);
new_value->allocation_type = xnn_allocation_type_workspace;
// Since we want to rewrite the previous write's value, we need to go
// back and update previously visited nodes. We only want to do the
// mapping after we find the first time this value is written.
bool found = false;
for (uint32_t k = 0; k < i; ++k) {
struct xnn_node* node_k = &subgraph->nodes[k];
if (found) {
// We've written the new value, update all subsequent reads to point
// to the new value.
replace_in_set(node_k->inputs, node_k->num_inputs, output_id,
new_value->id);
replace_in_set(node_k->outputs, node_k->num_outputs, output_id,
new_value->id);
} else if (set_contains(node_k->outputs, node_k->num_outputs,
output_id)) {
// We found where this value is written. Replace only the output
// use.
replace_in_set(node_k->outputs, node_k->num_outputs, output_id,
new_value->id);
found = true;
} else {
// We're still using the old value.
}
}
// Replace all subsequent reads with the new value, until we find
// another write (because subsequent reads from there need the
// replacement for that value).
for (uint32_t k = i; k < subgraph->num_nodes; ++k) {
struct xnn_node* node_k = &subgraph->nodes[k];
replace_in_set(node_k->inputs, node_k->num_inputs, output_id,
new_value->id);
if (set_contains(node_k->outputs, node_k->num_outputs, output_id)) {
break;
}
}
}
values_written[output_id] = true;
}
}
xnn_release_memory(values_written);
}
enum xnn_status xnn_subgraph_pack_static_values_to_fp16(
xnn_subgraph_t subgraph) {
for (uint32_t n = 0; n < subgraph->num_values; n++) {
struct xnn_value* value = &subgraph->values[n];
if (xnn_value_is_static(value->allocation_type) &&
(value->flags & XNN_VALUE_FLAG_PACK_TO_FP16)) {
if (value->datatype == xnn_datatype_fp16) {
continue;
}
if (value->datatype != xnn_datatype_fp32) {
xnn_log_error(
"failed to pack value #%" PRIu32 " to FP16: unsupported datatype %s"
" (expected FP32)", n, xnn_datatype_to_string(value->datatype));
return xnn_status_invalid_parameter;
}
const size_t fp16_size = xnn_tensor_get_size(value) / 2 + XNN_EXTRA_BYTES;
void* fp16_data = xnn_allocate_zero_memory(fp16_size);
if (fp16_data == NULL) {
xnn_log_error(
"failed to allocate %zu bytes for packed fp16 tensor data",
fp16_size);
return xnn_status_out_of_memory;
}
const size_t num_elements = xnn_shape_multiply_all_dims(&value->shape);
enum xnn_status status = xnn_run_unary_elementwise_nc(
xnn_unary_convert, xnn_datatype_fp32, xnn_datatype_fp16,
/*params=*/NULL, /*input_quantization=*/NULL,
/*output_quantization=*/NULL, 0, num_elements, 1, 1, 1, NULL,
value->data, fp16_data);
if (status != xnn_status_success) {
xnn_release_memory(fp16_data);
xnn_log_error("failed to convert value #%" PRIu32 " data to FP16", n);
return status;
}
if (value->flags & XNN_VALUE_FLAG_NEEDS_CLEANUP) {
XNN_PRAGMA_CLANG("clang diagnostic push")
XNN_PRAGMA_CLANG("clang diagnostic ignored \"-Wcast-qual\"")
xnn_release_memory((void*)value->data);
XNN_PRAGMA_CLANG("clang diagnostic pop")
}
value->data = fp16_data;
value->datatype = xnn_datatype_fp16;
value->flags |= XNN_VALUE_FLAG_NEEDS_CLEANUP;
xnn_log_debug("packed value #%" PRIu32 " to FP16", n);
}
}
return xnn_status_success;
}
enum xnn_status xnn_subgraph_optimize(xnn_subgraph_t subgraph,
uint32_t optimization_flags) {
// If the subgraph has no nodes, then there is nothing for us to do here, but
// do print a notice to the user as this seems a bit unusual.
if (!subgraph->num_nodes) {
xnn_log_info("Trying to optimize subgraph with zero nodes, skipping.");
return xnn_status_success;
}
XNN_RETURN_IF_ERROR(xnn_subgraph_pack_static_values_to_fp16(subgraph));
// Start with a clean and ordered subgraph.
xnn_subgraph_clean_up(subgraph);
if (!(optimization_flags & XNN_FLAG_NO_OPERATOR_FUSION)) {
xnn_subgraph_fusion(subgraph);
xnn_subgraph_fuse_unary_quantized_into_lut(subgraph);
}
const struct xnn_hardware_config* hardware_config =
xnn_init_hardware_config();
if (hardware_config == NULL) {
xnn_log_error("failed to get hardware config");
return xnn_status_unsupported_hardware;
}
if ((optimization_flags & XNN_FLAG_FORCE_FP16_INFERENCE) &&
(!xnn_is_f16_compatible_config(hardware_config))) {
xnn_log_error(
"failed to force FP16 inference: hardware supports neither native nor "
"emulated FP16 operators");
return xnn_status_unsupported_hardware;
}
const bool try_native_fp16 =
(optimization_flags & XNN_FLAG_HINT_FP16_INFERENCE) &&
xnn_is_f16_supported_natively(hardware_config);
const bool force_fp16 = (optimization_flags & XNN_FLAG_FORCE_FP16_INFERENCE);
if (try_native_fp16 || force_fp16) {
const bool fp16_rewrite_succeeded = xnn_subgraph_rewrite_for_fp16(subgraph);
if (force_fp16 && !fp16_rewrite_succeeded) {
xnn_log_error(
"failed to force FP16 inference: subgraph is incompatible with FP16 "
"operators");
return xnn_status_unsupported_parameter;
}
if (fp16_rewrite_succeeded) {
// Re-run xnn_subgraph_analyze_consumers_and_producers since fp16 re-write
// inserts nodes and changes producers/consumers.
xnn_subgraph_analyze_consumers_and_producers(subgraph);
}
}
// Apply some common subgraph optimizations.
XNN_RETURN_IF_ERROR(
xnn_subgraph_optimize_common_subgraphs(subgraph, optimization_flags));
#if XNN_ENABLE_SPARSE
if ((optimization_flags & XNN_FLAG_HINT_SPARSE_INFERENCE) &&
(xnn_is_chw_compatible_config(hardware_config))) {
xnn_subgraph_rewrite_for_nchw(subgraph);
}
#endif
XNN_RETURN_IF_ERROR(xnn_subgraph_rewrite_dequant_bmm(subgraph));
XNN_RETURN_IF_ERROR(
xnn_subgraph_optimize_packed_lhs(subgraph, optimization_flags));
if (!force_fp16) {
XNN_RETURN_IF_ERROR(xnn_subgraph_fallback_from_fp16_to_fp32(subgraph, optimization_flags));
}
// bf16 values are independent of the fp16 force flag, so this always runs to
// lower any bf16 ops that lack a native kernel to fp32.
XNN_RETURN_IF_ERROR(
xnn_subgraph_fallback_from_bf16_to_fp32(subgraph, optimization_flags));
return xnn_status_success;
}
uint32_t xnn_subgraph_get_value_flags(xnn_subgraph_t subgraph,
uint32_t value_id) {
return subgraph->values[value_id].flags;
}
enum xnn_datatype xnn_subgraph_get_value_datatype(xnn_subgraph_t subgraph,
uint32_t value_id) {
return subgraph->values[value_id].datatype;
}
uint32_t xnn_subgraph_get_num_external_values(xnn_subgraph_t subgraph) {
return subgraph->external_value_ids;
}
uint32_t xnn_subgraph_get_num_nodes(xnn_subgraph_t subgraph) {
return subgraph->num_nodes;
}
uint32_t xnn_subgraph_get_num_values(xnn_subgraph_t subgraph) {
return subgraph->num_values;
}
enum xnn_status xnn_delete_subgraph(xnn_subgraph_t subgraph) {
if (subgraph != NULL) {
if (subgraph->nodes != NULL) {
memset(subgraph->nodes, 0, sizeof(struct xnn_node) * subgraph->num_nodes);
xnn_release_memory(subgraph->nodes);
}
if (subgraph->values != NULL) {
// Release the dynamic allocations created during FP16 rewrite, if the
// subgraph still has ownership of them.
for (uint32_t i = 0; i < subgraph->num_values; i++) {
struct xnn_value* value = &subgraph->values[i];
if ((value->fp16_rewrite.fp16_compatible ||
(value->flags & XNN_VALUE_FLAG_NEEDS_CLEANUP)) &&
value->data != NULL) {
XNN_PRAGMA_CLANG("clang diagnostic push")
XNN_PRAGMA_CLANG("clang diagnostic ignored \"-Wcast-qual\"")
xnn_release_memory((void*)value->data);
XNN_PRAGMA_CLANG("clang diagnostic pop")
}
}
memset(subgraph->values, 0,
sizeof(struct xnn_value) * subgraph->num_values);
xnn_release_memory(subgraph->values);
}
memset(subgraph, 0, sizeof(struct xnn_subgraph));
xnn_release_memory(subgraph);
}
return xnn_status_success;
}
enum xnn_node_type xnn_reduce_operator_to_node_type(
enum xnn_reduce_operator type) {
switch (type) {
case xnn_reduce_max:
return xnn_node_type_static_reduce_max;
case xnn_reduce_mean:
return xnn_node_type_static_mean;
case xnn_reduce_mean_squared:
return xnn_node_type_static_mean_squared;
case xnn_reduce_min:
return xnn_node_type_static_reduce_min;
case xnn_reduce_sum:
return xnn_node_type_static_sum;
case xnn_reduce_sum_squared:
return xnn_node_type_static_sum_squared;
default:
return xnn_node_type_invalid;
}
}
enum xnn_reduce_operator xnn_node_type_to_reduce_operator(
enum xnn_node_type type) {
switch (type) {
case xnn_node_type_static_mean:
return xnn_reduce_mean;
case xnn_node_type_static_mean_squared:
return xnn_reduce_mean_squared;
case xnn_node_type_static_reduce_max:
return xnn_reduce_max;
case xnn_node_type_static_reduce_min:
return xnn_reduce_min;
case xnn_node_type_static_sum:
return xnn_reduce_sum;
case xnn_node_type_static_sum_squared:
return xnn_reduce_sum_squared;
default:
return xnn_reduce_invalid;
}
}