// Copyright (c) ONNX Project Contributors // // SPDX-License-Identifier: Apache-2.0 // Helper Methods for Adapters #pragma once #include #include #include "onnx/common/assertions.h" #include "onnx/common/ir.h" #include "onnx/defs/tensor_util.h" namespace ONNX_NAMESPACE::version_conversion { int check_numpy_unibroadcastable_and_require_broadcast( const std::vector& input1_sizes, const std::vector& input2_sizes); void assert_numpy_multibroadcastable( const std::vector& input1_sizes, const std::vector& input2_sizes); void assertNotParams(const std::vector& sizes); void assertInputsAvailable(const ArrayRef& inputs, const char* name, uint64_t num_inputs); // Decode an INT64 tensor; rejects mismatched element type or dims/raw byte length. inline std::vector ReadInt64Tensor(const Tensor& tensor) { ONNX_ASSERTM( tensor.elem_type() == ONNX_NAMESPACE::TensorProto_DataType_INT64, "expected INT64 tensor, got elem_type=", tensor.elem_type()) if (tensor.is_raw_data()) { const size_t raw_bytes = tensor.raw().size(); // elem_num() returns 1 for scalars, so covers dims=[]. ONNX_ASSERTM( raw_bytes == static_cast(tensor.elem_num()) * sizeof(int64_t), "INT64 tensor: ", raw_bytes, " raw bytes does not match dims (", tensor.elem_num(), " elements)") } return ParseData(&tensor); } } // namespace ONNX_NAMESPACE::version_conversion