// Copyright (c) ONNX Project Contributors // // SPDX-License-Identifier: Apache-2.0 #pragma once #include #include #include #include "onnx/defs/schema.h" #include "onnx/defs/shape_inference.h" #include "onnx/defs/tensor_proto_util.h" #include "onnx/onnx_pb.h" namespace ONNX_NAMESPACE::defs::math::utils { std::function TopKOpGenerator(std::vector allowed_types); // Unary elementwise ops on float types: T input -> T output, no attrs, no function body. std::function UnaryFloatMathOpGenerator( const char* doc, const char* output_description, std::vector allowed_types = OpSchema::all_float_types_ir4()); template T GetScalarValueFromTensor(const ONNX_NAMESPACE::TensorProto* t) { if (t == nullptr) { return T{}; } auto data_type = t->data_type(); switch (data_type) { case ONNX_NAMESPACE::TensorProto::FLOAT: return static_cast(ONNX_NAMESPACE::ParseData(t).at(0)); case ONNX_NAMESPACE::TensorProto::DOUBLE: return static_cast(ONNX_NAMESPACE::ParseData(t).at(0)); case ONNX_NAMESPACE::TensorProto::INT32: return static_cast(ONNX_NAMESPACE::ParseData(t).at(0)); case ONNX_NAMESPACE::TensorProto::INT64: return static_cast(ONNX_NAMESPACE::ParseData(t).at(0)); default: fail_shape_inference("Unsupported input data type of ", data_type); } } void MatMulShapeInference(ONNX_NAMESPACE::InferenceContext& ctx, int input1Idx, int input2Idx); void QLinearMatMulShapeInference(ONNX_NAMESPACE::InferenceContext& ctx); const char* QLinearMatMulDoc(); int64_t MathOpTwoIntegers(const std::string& op_type, int64_t a, int64_t b); } // namespace ONNX_NAMESPACE::defs::math::utils