[core] Low precision element iterator and `u2, u3, u6` types (#23279)
### Details:
- Introduce new low precision types `u2`, `u3`, `u6`.
- Introduce `ov::element::Iterator` for low precision types like `u1,
u2, u3, u4, i4, u6`:
- Gives pointer like access to low precision values in Tensor,
containers etc.
- Can be used by STL algorithms to access data in unified algorithms for
data manipulation.
- Can be used in Constant, Convert operators to replace duplicate
implementations for accessing low precision data (bin-size reduction).
- Can be used for operator reference implementation or plugin if there
is no hardware specific solution.
### Tickets:
- [CVS-126998](https://jira.devtools.intel.com/browse/CVS-126998)
- Part of
[CVS-128024](https://jira.devtools.intel.com/browse/CVS-128024)
This commit is contained in:
parent
4d06afa3da
commit
fd93e3b33f
|
|
@ -18,7 +18,7 @@ VariableReference: '^\w+$'
|
|||
|
||||
EnumName: '^[A-Z][\w]+$'
|
||||
# excepts element_type
|
||||
EnumConstantName: '^([A-Z\d_]+|undefined|dynamic|boolean|bf16|f16|f32|f64|i4|i8|i16|i32|i64|u1|u4|u8|u16|u32|u64|nf4|f8e4m3|f8e5m2|string|asymmetric|align_corners|round_prefer_floor|round_prefer_ceil|floor|ceil|simple|nearest|linear|linear_onnx|cubic|area|scales|sizes|half_pixel|tf_half_pixel_for_nn|pytorch_half_pixel|asymetric)$'
|
||||
EnumConstantName: '^([A-Z\d_]+|undefined|dynamic|boolean|bf16|f16|f32|f64|i4|i8|i16|i32|i64|u1|u2|u3|u4|u6|u8|u16|u32|u64|nf4|f8e4m3|f8e5m2|string|asymmetric|align_corners|round_prefer_floor|round_prefer_ceil|floor|ceil|simple|nearest|linear|linear_onnx|cubic|area|scales|sizes|half_pixel|tf_half_pixel_for_nn|pytorch_half_pixel|asymetric)$'
|
||||
# TODO: align
|
||||
UsingDeclaration: '^.*$'
|
||||
TypedefName: '^.*$'
|
||||
|
|
|
|||
|
|
@ -0,0 +1,502 @@
|
|||
// Copyright (C) 2018-2024 Intel Corporation
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
//
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "openvino/core/type/element_type_traits.hpp"
|
||||
|
||||
namespace ov {
|
||||
namespace util {
|
||||
|
||||
/**
|
||||
* @brief Make bit mask by setting N less significant bits.
|
||||
*
|
||||
* @tparam T Type of value.
|
||||
* @param n Number of bits to set.
|
||||
* @return Bit-mask value with N bits set.
|
||||
*/
|
||||
template <class T>
|
||||
constexpr T make_n_bit_mask(const T n) {
|
||||
return (1ULL << n) - 1ULL;
|
||||
}
|
||||
} // namespace util
|
||||
|
||||
namespace element {
|
||||
|
||||
/**
|
||||
* @brief Checks if element type is N in-raw bits type.
|
||||
*
|
||||
* @param et Element type to check
|
||||
* @return True if element type is bit type otherwise false.
|
||||
*/
|
||||
constexpr bool is_bit_type(Type_t et) {
|
||||
return et == u1 || et == u2;
|
||||
}
|
||||
|
||||
/**
|
||||
* @brief Checks if element type is 4-bits type.
|
||||
*
|
||||
* @param et Element type to check
|
||||
* @return True if element type is nibble type otherwise false.
|
||||
*/
|
||||
constexpr bool is_nibble_type(Type_t et) {
|
||||
return et == u4 || et == i4 || et == nf4;
|
||||
}
|
||||
|
||||
/**
|
||||
* @brief Checks if element type is split bit type.
|
||||
*
|
||||
* The value is stored in byte(s) like [b0, b1, x, .., x, b2, b3].
|
||||
*
|
||||
* @param et Element type to check
|
||||
* @return True if element type is split bit type otherwise false.
|
||||
*/
|
||||
constexpr bool is_split_bit_type(Type_t et) {
|
||||
return et == u3 || et == u6;
|
||||
}
|
||||
|
||||
/**
|
||||
* @brief Checks element type is using only N bytes as value.
|
||||
*
|
||||
* @param et Element type to check.
|
||||
* @return True if element type use byte(s) for its value, false otherwise.
|
||||
*/
|
||||
constexpr bool is_byte_type(Type_t et) {
|
||||
return !is_bit_type(et) && !is_split_bit_type(et) && !is_nibble_type(et) && et != string;
|
||||
}
|
||||
|
||||
/**
|
||||
* @brief Gets bit width of ov::element::Type_t.
|
||||
*
|
||||
* @return Number of bits representing the Type_t.
|
||||
*/
|
||||
template <Type_t ET>
|
||||
constexpr size_t bit_width() {
|
||||
return sizeof(typename ov::fundamental_type_for<ET>());
|
||||
}
|
||||
|
||||
template <>
|
||||
constexpr size_t bit_width<Type_t::u1>() {
|
||||
return 1;
|
||||
}
|
||||
|
||||
template <>
|
||||
constexpr size_t bit_width<Type_t::u2>() {
|
||||
return 2;
|
||||
}
|
||||
|
||||
template <>
|
||||
constexpr size_t bit_width<Type_t::u3>() {
|
||||
return 3;
|
||||
}
|
||||
|
||||
template <>
|
||||
constexpr size_t bit_width<Type_t::u4>() {
|
||||
return 4;
|
||||
}
|
||||
|
||||
template <>
|
||||
constexpr size_t bit_width<Type_t::i4>() {
|
||||
return 4;
|
||||
}
|
||||
|
||||
template <>
|
||||
constexpr size_t bit_width<Type_t::u6>() {
|
||||
return 6;
|
||||
}
|
||||
|
||||
/**
|
||||
* @brief The BitProxy value class used by ov::element::Iterator to access values which has no standard byte(s) layout.
|
||||
*
|
||||
* It used by iterator to access values represented by precisions like u2, i4, u6 etc. in the way like stored
|
||||
* on bytes.
|
||||
* The R/W access is done via conversion and copy assignment operators.
|
||||
* The public members are used to work on sub-byte value like on its fundamental type defined by T.
|
||||
*
|
||||
* @tparam T Fundamental type of sub-byte value which must be same as fundamental type of element::Type_t.
|
||||
* @tparam ET OpenVINO element type.
|
||||
* @tparam Enable Type to enable/disable this class.
|
||||
*/
|
||||
template <class T, Type_t ET, class Enable = void>
|
||||
class BitProxy {};
|
||||
|
||||
/**
|
||||
* @brief The BitProxy specialization for types which are represented by N in-raw bits in byte.
|
||||
*
|
||||
* @tparam T Fundamental type of sub-byte value which must be same as fundamental type of element::Type_t.
|
||||
* @tparam ET OpenVINO element type.
|
||||
*/
|
||||
template <class T, Type_t ET>
|
||||
class BitProxy<T, ET, typename std::enable_if<is_bit_type(ET) || is_nibble_type(ET)>::type> {
|
||||
private:
|
||||
template <Type_t, class>
|
||||
friend class Iterator; //!< Iterator class is friend to access private members to manipulate pointer.
|
||||
|
||||
static constexpr size_t m_bits = bit_width<ET>(); //!< Number of bit for single value.
|
||||
static constexpr size_t m_num_values = 8 / m_bits; //!< Number values in byte.
|
||||
static constexpr size_t m_shift_init = is_nibble_type(ET) ? 0 : 8 - m_bits; //!< Initial value for bit shift.
|
||||
|
||||
T* m_ptr; //!< Pointer to T used to get value.
|
||||
size_t m_bit_shift; //!< Current bit shift to get value.
|
||||
|
||||
constexpr BitProxy(T* ptr) noexcept : m_ptr{ptr}, m_bit_shift{m_shift_init} {}
|
||||
|
||||
uint8_t get_bit_value() const {
|
||||
constexpr auto value_mask = util::make_n_bit_mask(m_bits);
|
||||
return (*m_ptr >> m_bit_shift) & value_mask;
|
||||
}
|
||||
|
||||
public:
|
||||
using value_type = typename std::decay<T>::type; //!< Fundamental type of bound to BitProxy.
|
||||
|
||||
/**
|
||||
* @brief Compare proxy value with other provided value.
|
||||
* @param rhs Value to compare.
|
||||
* @return True if equal otherwise false.
|
||||
*/
|
||||
template <class U>
|
||||
constexpr bool operator==(const U& rhs) const {
|
||||
return static_cast<value_type>(*this) == rhs;
|
||||
}
|
||||
|
||||
/**
|
||||
* @brief Compare proxy value is less than rhs.
|
||||
*
|
||||
* @tparam U Type of value to compare.
|
||||
* @param rhs Value to compare.
|
||||
* @return True if less otherwise false.
|
||||
*/
|
||||
template <class U>
|
||||
constexpr bool operator<(const U& rhs) const {
|
||||
return static_cast<value_type>(*this) < rhs;
|
||||
}
|
||||
|
||||
/**
|
||||
* @brief Converts to fundamental type.
|
||||
*
|
||||
* @return Value of BitProxy.
|
||||
*/
|
||||
template <Type_t ETT = ET, typename std::enable_if<ETT != i4>::type* = nullptr>
|
||||
operator value_type() const {
|
||||
return static_cast<value_type>(get_bit_value());
|
||||
}
|
||||
|
||||
/**
|
||||
* @brief Converts to fundamental type.
|
||||
*
|
||||
* @return Value of BitProxy.
|
||||
*/
|
||||
template <Type_t ETT = ET, typename std::enable_if<ETT == i4>::type* = nullptr>
|
||||
operator value_type() const {
|
||||
constexpr auto value_mask = util::make_n_bit_mask(m_bits);
|
||||
constexpr auto value_msb_mask = (1U << (m_bits - 1U));
|
||||
|
||||
auto v = get_bit_value();
|
||||
if (v & value_msb_mask) {
|
||||
// If N bit value MSB bit is set then value is negative.
|
||||
// As v is byte then all bits above N must be set to be two's complement.
|
||||
v |= ~value_mask;
|
||||
}
|
||||
return static_cast<value_type>(v);
|
||||
}
|
||||
|
||||
/**
|
||||
* @brief Sets current ProxyBit to value.
|
||||
* @param v Value to be set.
|
||||
*/
|
||||
BitProxy<T, ET>& operator=(const value_type v) {
|
||||
constexpr auto value_mask = util::make_n_bit_mask(m_bits);
|
||||
*m_ptr &= ~(value_mask << m_bit_shift);
|
||||
*m_ptr |= (static_cast<uint8_t>(v) & value_mask) << m_bit_shift;
|
||||
return *this;
|
||||
}
|
||||
};
|
||||
|
||||
/**
|
||||
* @brief The BitProxy specialization for u3, u6 precisions.
|
||||
*
|
||||
* @note The input pointer must point on buffer which has got 3 * n bytes.
|
||||
*
|
||||
* @tparam T Fundamental type of sub-byte value which must be same as fundamental type of element::Type_t.
|
||||
* @tparam ET OpenVINO element type.
|
||||
*/
|
||||
template <class T, Type_t ET>
|
||||
class BitProxy<T, ET, typename std::enable_if<is_split_bit_type(ET)>::type> {
|
||||
private:
|
||||
template <Type_t, class>
|
||||
friend class Iterator; //!< Iterator class is friend to access private members to manipulate pointer.
|
||||
|
||||
static constexpr size_t m_bits = bit_width<ET>(); //!< Number of bit for single value.
|
||||
static constexpr size_t m_num_values = (3 * 8) / m_bits; //!< Number values in byte.
|
||||
static constexpr size_t m_shift_init = m_num_values - 1; //!< Initial value for bit shift.
|
||||
|
||||
struct ByteValue {
|
||||
uint8_t b0;
|
||||
uint8_t b1;
|
||||
uint8_t b2;
|
||||
};
|
||||
|
||||
union {
|
||||
T* m_ptr; //!< Pointer to T buffer.
|
||||
ByteValue* m_bytes; //!< Pointer to buffer as 3 bytes representation.
|
||||
};
|
||||
|
||||
size_t m_bit_shift; //!< Current bit shift to get value.
|
||||
|
||||
constexpr BitProxy(T* ptr) noexcept : m_ptr{ptr}, m_bit_shift{m_shift_init} {}
|
||||
|
||||
public:
|
||||
using value_type = typename std::decay<T>::type; //!< Fundamental type of sub-byte.
|
||||
|
||||
/**
|
||||
* @brief Compare proxy value is equal than rhs.
|
||||
*
|
||||
* @tparam U Type of value to compare.
|
||||
* @param rhs Value to compare.
|
||||
* @return True if equal, false otherwise.
|
||||
*/
|
||||
template <class U>
|
||||
constexpr bool operator==(const U& rhs) const {
|
||||
return static_cast<value_type>(*this) == rhs;
|
||||
}
|
||||
|
||||
/**
|
||||
* @brief Compare proxy value is less than rhs.
|
||||
*
|
||||
* @tparam U Type of value to compare.
|
||||
* @param rhs Value to compare.
|
||||
* @return True if less otherwise false.
|
||||
*/
|
||||
template <class U>
|
||||
constexpr bool operator<(const U& rhs) const {
|
||||
return static_cast<value_type>(*this) < rhs;
|
||||
}
|
||||
|
||||
/**
|
||||
* @brief Converts to fundamental type.
|
||||
*
|
||||
* @return Value of BitProxy.
|
||||
*/
|
||||
operator value_type() const {
|
||||
constexpr uint16_t lower_mask_bits = 16 / m_num_values;
|
||||
constexpr uint16_t upper_mask_bits = 8 / m_num_values;
|
||||
constexpr uint16_t mask_lower = util::make_n_bit_mask(lower_mask_bits);
|
||||
constexpr uint16_t mask_upper = util::make_n_bit_mask(upper_mask_bits) << lower_mask_bits;
|
||||
|
||||
// get lower part of value
|
||||
uint16_t v = ((m_bytes->b0 << 8U) | m_bytes->b1) >> (lower_mask_bits * m_bit_shift);
|
||||
v &= mask_lower;
|
||||
// get upper part of value
|
||||
v |= ((m_bytes->b2 << lower_mask_bits) >> (upper_mask_bits * m_bit_shift)) & mask_upper;
|
||||
return static_cast<value_type>(v);
|
||||
}
|
||||
|
||||
/**
|
||||
* @brief Sets current ProxyBit to value.
|
||||
* @param v Value to be set.
|
||||
*/
|
||||
BitProxy<T, ET>& operator=(const value_type v) {
|
||||
constexpr uint16_t lower_mask_bits = 16 / m_num_values;
|
||||
constexpr uint16_t upper_mask_bits = 8 / m_num_values;
|
||||
constexpr uint16_t mask_lower = util::make_n_bit_mask(lower_mask_bits);
|
||||
constexpr uint16_t mask_upper = util::make_n_bit_mask(upper_mask_bits) << lower_mask_bits;
|
||||
|
||||
uint16_t tmp = (m_bytes->b0 << 8U) | m_bytes->b1;
|
||||
tmp &= ~(mask_lower << (lower_mask_bits * m_bit_shift));
|
||||
tmp |= (v & mask_lower) << (lower_mask_bits * m_bit_shift);
|
||||
m_bytes->b0 = tmp >> 8U;
|
||||
m_bytes->b1 = tmp & 0x00ff;
|
||||
|
||||
tmp = m_bytes->b2 & ~((mask_upper >> lower_mask_bits) << (upper_mask_bits * m_bit_shift));
|
||||
tmp |= (((v & mask_upper) >> lower_mask_bits) << (upper_mask_bits * m_bit_shift));
|
||||
m_bytes->b2 = tmp & 0x00ff;
|
||||
return *this;
|
||||
}
|
||||
};
|
||||
|
||||
/**
|
||||
* @brief Put BitProxy value to output stream.
|
||||
*
|
||||
* @param os Reference to output stream.
|
||||
* @param value Value to print.
|
||||
* @return return output stream.
|
||||
*/
|
||||
template <class T, Type_t ET>
|
||||
std::ostream& operator<<(std::ostream& os, const BitProxy<T, ET>& value) {
|
||||
os << +static_cast<T>(value);
|
||||
return os;
|
||||
}
|
||||
|
||||
/**
|
||||
* @brief Bidirectional iterator of specified precision.
|
||||
*
|
||||
* The iterator supports low precisions using BitProxy to access values via conversion.
|
||||
*
|
||||
* @tparam ET Type of OpenVINO element type (ov::element::Type_t).
|
||||
* @tparam T Must be fundamental type for specified ET.
|
||||
*/
|
||||
template <Type_t ET, class T>
|
||||
class Iterator {
|
||||
using proxy_type = BitProxy<T, ET>;
|
||||
|
||||
public:
|
||||
using iterator_category = std::bidirectional_iterator_tag;
|
||||
using difference_type = std::ptrdiff_t;
|
||||
using value_type = T;
|
||||
using reference = typename std::conditional<std::is_const<T>::value, const proxy_type&, proxy_type&>::type;
|
||||
using pointer = typename std::conditional<std::is_const<T>::value, const proxy_type*, proxy_type*>::type;
|
||||
|
||||
static_assert(std::is_same<typename std::decay<T>::type, ov::fundamental_type_for<ET>>::value,
|
||||
"Iterator value_type must be same as fundamental type of ET");
|
||||
|
||||
constexpr Iterator(T* ptr) noexcept : m_et_ptr{ptr} {}
|
||||
|
||||
// Iteration operators
|
||||
template <Type_t ETT = ET>
|
||||
typename std::enable_if<is_bit_type(ETT), Iterator<ET, T>>::type& operator++() {
|
||||
m_et_ptr.m_bit_shift -= m_et_ptr.m_bits;
|
||||
m_et_ptr.m_bit_shift = m_et_ptr.m_bit_shift % (m_et_ptr.m_num_values * m_et_ptr.m_bits);
|
||||
m_et_ptr.m_ptr += static_cast<std::ptrdiff_t>(m_et_ptr.m_bit_shift == m_et_ptr.m_shift_init);
|
||||
return *this;
|
||||
}
|
||||
|
||||
template <Type_t ETT = ET>
|
||||
typename std::enable_if<is_nibble_type(ETT), Iterator<ET, T>>::type& operator++() {
|
||||
m_et_ptr.m_bit_shift ^= m_et_ptr.m_bits;
|
||||
m_et_ptr.m_ptr += static_cast<std::ptrdiff_t>(m_et_ptr.m_bit_shift == m_et_ptr.m_shift_init);
|
||||
return *this;
|
||||
}
|
||||
|
||||
template <Type_t ETT = ET>
|
||||
typename std::enable_if<is_split_bit_type(ETT), Iterator<ET, T>>::type& operator++() {
|
||||
--m_et_ptr.m_bit_shift;
|
||||
m_et_ptr.m_bit_shift = m_et_ptr.m_bit_shift % m_et_ptr.m_num_values;
|
||||
m_et_ptr.m_ptr += (m_et_ptr.m_bit_shift == m_et_ptr.m_shift_init) ? 3 : 0;
|
||||
return *this;
|
||||
}
|
||||
|
||||
Iterator<ET, T> operator++(int) {
|
||||
auto old = *this;
|
||||
++(*this);
|
||||
return old;
|
||||
}
|
||||
|
||||
template <Type_t ETT = ET>
|
||||
typename std::enable_if<is_bit_type(ETT), Iterator<ET, T>>::type& operator+=(const difference_type& n) {
|
||||
const auto advance = n + (m_et_ptr.m_shift_init - m_et_ptr.m_bit_shift) / m_et_ptr.m_bits;
|
||||
m_et_ptr.m_bit_shift = m_et_ptr.m_shift_init - (advance % m_et_ptr.m_num_values) * m_et_ptr.m_bits;
|
||||
m_et_ptr.m_ptr += advance / m_et_ptr.m_num_values;
|
||||
return *this;
|
||||
}
|
||||
|
||||
template <Type_t ETT = ET>
|
||||
typename std::enable_if<is_nibble_type(ETT), Iterator<ET, T>>::type& operator+=(const difference_type& n) {
|
||||
m_et_ptr.m_ptr += n / m_et_ptr.m_num_values;
|
||||
return (n % m_et_ptr.m_num_values) ? ++*this : *this;
|
||||
}
|
||||
|
||||
template <Type_t ETT = ET>
|
||||
typename std::enable_if<is_split_bit_type(ETT), Iterator<ET, T>>::type& operator+=(const difference_type& n) {
|
||||
const auto advance = n + m_et_ptr.m_shift_init - m_et_ptr.m_bit_shift;
|
||||
m_et_ptr.m_bit_shift = m_et_ptr.m_shift_init - (advance % m_et_ptr.m_num_values);
|
||||
m_et_ptr.m_ptr += 3 * (advance / m_et_ptr.m_num_values);
|
||||
return *this;
|
||||
}
|
||||
|
||||
Iterator<ET, T> operator+(const difference_type& n) {
|
||||
auto tmp(*this);
|
||||
tmp += n;
|
||||
return tmp;
|
||||
}
|
||||
|
||||
template <Type_t ETT = ET>
|
||||
typename std::enable_if<is_bit_type(ETT), Iterator<ET, T>>::type& operator--() {
|
||||
m_et_ptr.m_bit_shift += m_et_ptr.m_bits;
|
||||
m_et_ptr.m_bit_shift = m_et_ptr.m_bit_shift % (m_et_ptr.m_num_values * m_et_ptr.m_bits);
|
||||
m_et_ptr.m_ptr -= static_cast<std::ptrdiff_t>(m_et_ptr.m_bit_shift == 0);
|
||||
return *this;
|
||||
}
|
||||
|
||||
template <Type_t ETT = ET>
|
||||
typename std::enable_if<is_nibble_type(ETT), Iterator<ET, T>>::type& operator--() {
|
||||
m_et_ptr.m_bit_shift ^= m_et_ptr.m_bits;
|
||||
m_et_ptr.m_ptr -= static_cast<std::ptrdiff_t>(m_et_ptr.m_bit_shift == 4);
|
||||
return *this;
|
||||
}
|
||||
|
||||
template <Type_t ETT = ET>
|
||||
typename std::enable_if<is_split_bit_type(ETT), Iterator<ET, T>>::type& operator--() {
|
||||
++m_et_ptr.m_bit_shift;
|
||||
m_et_ptr.m_bit_shift = m_et_ptr.m_bit_shift % m_et_ptr.m_num_values;
|
||||
m_et_ptr.m_ptr -= m_et_ptr.m_bit_shift == 0 ? 3 : 0;
|
||||
return *this;
|
||||
}
|
||||
|
||||
Iterator<ET, T> operator--(int) {
|
||||
auto old = *this;
|
||||
--(*this);
|
||||
return old;
|
||||
}
|
||||
|
||||
template <Type_t ETT = ET>
|
||||
typename std::enable_if<is_bit_type(ETT), Iterator<ET, T>>::type& operator-=(const difference_type& n) {
|
||||
const auto advance = m_et_ptr.m_bit_shift / m_et_ptr.m_bits + n;
|
||||
m_et_ptr.m_bit_shift = (advance % m_et_ptr.m_num_values) * m_et_ptr.m_bits;
|
||||
m_et_ptr.m_ptr -= advance / m_et_ptr.m_num_values;
|
||||
return *this;
|
||||
}
|
||||
|
||||
template <Type_t ETT = ET>
|
||||
typename std::enable_if<is_nibble_type(ETT), Iterator<ET, T>>::type& operator-=(const difference_type& n) {
|
||||
m_et_ptr.m_ptr -= n / m_et_ptr.m_num_values;
|
||||
return (n % m_et_ptr.m_num_values) ? --*this : *this;
|
||||
}
|
||||
|
||||
template <Type_t ETT = ET>
|
||||
typename std::enable_if<is_split_bit_type(ETT), Iterator<ET, T>>::type& operator-=(const difference_type& n) {
|
||||
const auto advance = m_et_ptr.m_bit_shift + n;
|
||||
m_et_ptr.m_bit_shift = advance % m_et_ptr.m_num_values;
|
||||
m_et_ptr.m_ptr -= 3 * (advance / m_et_ptr.m_num_values);
|
||||
return *this;
|
||||
}
|
||||
|
||||
Iterator<ET, T> operator-(const difference_type& n) {
|
||||
auto tmp(*this);
|
||||
tmp -= n;
|
||||
return tmp;
|
||||
}
|
||||
|
||||
// compare operators
|
||||
constexpr bool operator!=(const Iterator<ET, T>& rhs) const {
|
||||
return (m_et_ptr.m_ptr != rhs.m_et_ptr.m_ptr) || (m_et_ptr.m_bit_shift != rhs.m_et_ptr.m_bit_shift);
|
||||
}
|
||||
|
||||
// dereference operators
|
||||
constexpr const proxy_type& operator*() const {
|
||||
return m_et_ptr;
|
||||
}
|
||||
|
||||
reference operator*() {
|
||||
return m_et_ptr;
|
||||
}
|
||||
|
||||
private:
|
||||
proxy_type m_et_ptr;
|
||||
};
|
||||
|
||||
/**
|
||||
* @brief Make element iterator from pointer.
|
||||
*
|
||||
* @tparam ET Type of ov::element::Type_t.
|
||||
* @tparam T Type of pointer data. Must be fundamental type of ET.
|
||||
|
||||
* @param ptr Pointer to data.
|
||||
* @return Element iterator for type ET.
|
||||
*/
|
||||
template <Type_t ET, class T, typename std::enable_if<!is_byte_type(ET) && ET != string>::type* = nullptr>
|
||||
constexpr Iterator<ET, T> iterator(T* ptr) {
|
||||
return {ptr};
|
||||
}
|
||||
} // namespace element
|
||||
} // namespace ov
|
||||
|
|
@ -48,7 +48,10 @@ enum class Type_t {
|
|||
i32, //!< i32 element type
|
||||
i64, //!< i64 element type
|
||||
u1, //!< binary element type
|
||||
u2, //!< u2 element type
|
||||
u3, //!< u3 element type
|
||||
u4, //!< u4 element type
|
||||
u6, //!< u6 element type
|
||||
u8, //!< u8 element type
|
||||
u16, //!< u16 element type
|
||||
u32, //!< u32 element type
|
||||
|
|
@ -168,9 +171,18 @@ constexpr Type i64(Type_t::i64);
|
|||
/// \brief binary element type
|
||||
/// \ingroup ov_element_cpp_api
|
||||
constexpr Type u1(Type_t::u1);
|
||||
/// \brief u2 element type
|
||||
/// \ingroup ov_element_cpp_api
|
||||
constexpr Type u2(Type_t::u2);
|
||||
/// \brief u3 element type
|
||||
/// \ingroup ov_element_cpp_api
|
||||
constexpr Type u3(Type_t::u3);
|
||||
/// \brief u4 element type
|
||||
/// \ingroup ov_element_cpp_api
|
||||
constexpr Type u4(Type_t::u4);
|
||||
/// \brief u6 element type
|
||||
/// \ingroup ov_element_cpp_api
|
||||
constexpr Type u6(Type_t::u6);
|
||||
/// \brief u8 element type
|
||||
/// \ingroup ov_element_cpp_api
|
||||
constexpr Type u8(Type_t::u8);
|
||||
|
|
|
|||
|
|
@ -68,11 +68,26 @@ struct element_type_traits<element::Type_t::u1> {
|
|||
using value_type = int8_t;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct element_type_traits<element::Type_t::u2> {
|
||||
using value_type = int8_t;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct element_type_traits<element::Type_t::u3> {
|
||||
using value_type = int8_t;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct element_type_traits<element::Type_t::u4> {
|
||||
using value_type = int8_t;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct element_type_traits<element::Type_t::u6> {
|
||||
using value_type = int8_t;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct element_type_traits<element::Type_t::u8> {
|
||||
using value_type = uint8_t;
|
||||
|
|
|
|||
|
|
@ -146,6 +146,9 @@ public:
|
|||
case Type_t::string:
|
||||
fill_data<Type_t::string>(value);
|
||||
break;
|
||||
case Type_t::u2:
|
||||
case Type_t::u3:
|
||||
case Type_t::u6:
|
||||
case Type_t::undefined:
|
||||
case Type_t::dynamic:
|
||||
OPENVINO_THROW("unsupported type");
|
||||
|
|
@ -872,6 +875,9 @@ private:
|
|||
case Type_t::string:
|
||||
write_buffer<Type_t::string>(source);
|
||||
break;
|
||||
case element::Type_t::u2:
|
||||
case element::Type_t::u3:
|
||||
case element::Type_t::u6:
|
||||
case element::Type_t::undefined:
|
||||
case element::Type_t::dynamic:
|
||||
OPENVINO_THROW("unsupported type");
|
||||
|
|
|
|||
|
|
@ -373,7 +373,10 @@ static std::string get_value(const std::shared_ptr<ov::op::v0::Constant>& consta
|
|||
case ov::element::Type_t::undefined:
|
||||
case ov::element::Type_t::dynamic:
|
||||
case ov::element::Type_t::u1:
|
||||
case ov::element::Type_t::u2:
|
||||
case ov::element::Type_t::u3:
|
||||
case ov::element::Type_t::u4:
|
||||
case ov::element::Type_t::u6:
|
||||
case ov::element::Type_t::nf4:
|
||||
case ov::element::Type_t::i4:
|
||||
case ov::element::Type_t::f8e4m3:
|
||||
|
|
|
|||
|
|
@ -59,8 +59,14 @@ inline TypeInfo get_type_info(ov::element::Type_t type) {
|
|||
return {64, false, true, false, "int64_t", "i64"};
|
||||
case ov::element::Type_t::u1:
|
||||
return {1, false, false, false, "uint1_t", "u1"};
|
||||
case ov::element::Type_t::u2:
|
||||
return {2, false, false, false, "uint2_t", "u2"};
|
||||
case ov::element::Type_t::u3:
|
||||
return {3, false, false, false, "uint3_t", "u3"};
|
||||
case ov::element::Type_t::u4:
|
||||
return {4, false, false, false, "uint4_t", "u4"};
|
||||
case ov::element::Type_t::u6:
|
||||
return {6, false, false, false, "uint6_t", "u6"};
|
||||
case ov::element::Type_t::u8:
|
||||
return {8, false, false, true, "uint8_t", "u8"};
|
||||
case ov::element::Type_t::u16:
|
||||
|
|
@ -103,8 +109,14 @@ ov::element::Type type_from_string(const std::string& type) {
|
|||
return ::ov::element::Type(::ov::element::Type_t::i64);
|
||||
} else if (type == "u1" || type == "U1" || type == "BIN" || type == "bin") {
|
||||
return ::ov::element::Type(::ov::element::Type_t::u1);
|
||||
} else if (type == "u2" || type == "U2") {
|
||||
return ::ov::element::Type(::ov::element::Type_t::u2);
|
||||
} else if (type == "u3" || type == "U3") {
|
||||
return ::ov::element::Type(::ov::element::Type_t::u3);
|
||||
} else if (type == "u4" || type == "U4") {
|
||||
return ::ov::element::Type(::ov::element::Type_t::u4);
|
||||
} else if (type == "u6" || type == "U6") {
|
||||
return ::ov::element::Type(::ov::element::Type_t::u6);
|
||||
} else if (type == "u8" || type == "U8") {
|
||||
return ::ov::element::Type(::ov::element::Type_t::u8);
|
||||
} else if (type == "u16" || type == "U16") {
|
||||
|
|
@ -135,11 +147,11 @@ ov::element::Type type_from_string(const std::string& type) {
|
|||
|
||||
std::vector<const ov::element::Type*> ov::element::Type::get_known_types() {
|
||||
std::vector<const ov::element::Type*> rc = {
|
||||
&ov::element::dynamic, &ov::element::boolean, &ov::element::bf16, &ov::element::f16, &ov::element::f32,
|
||||
&ov::element::f64, &ov::element::i4, &ov::element::i8, &ov::element::i16, &ov::element::i32,
|
||||
&ov::element::i64, &ov::element::u1, &ov::element::u4, &ov::element::u8, &ov::element::u16,
|
||||
&ov::element::u32, &ov::element::u64, &ov::element::nf4, &ov::element::f8e4m3, &ov::element::f8e5m2,
|
||||
&ov::element::string};
|
||||
&ov::element::dynamic, &ov::element::boolean, &ov::element::bf16, &ov::element::f16, &ov::element::f32,
|
||||
&ov::element::f64, &ov::element::i4, &ov::element::i8, &ov::element::i16, &ov::element::i32,
|
||||
&ov::element::i64, &ov::element::u1, &ov::element::u2, &ov::element::u3, &ov::element::u4,
|
||||
&ov::element::u6, &ov::element::u8, &ov::element::u16, &ov::element::u32, &ov::element::u64,
|
||||
&ov::element::nf4, &ov::element::f8e4m3, &ov::element::f8e5m2, &ov::element::string};
|
||||
return rc;
|
||||
}
|
||||
|
||||
|
|
@ -163,7 +175,10 @@ ov::element::Type::Type(size_t bitwidth,
|
|||
{ov::element::Type_t::i32, {32, false, true, true, "int32_t", "i32"}},
|
||||
{ov::element::Type_t::i64, {64, false, true, false, "int64_t", "i64"}},
|
||||
{ov::element::Type_t::u1, {1, false, false, false, "uint1_t", "u1"}},
|
||||
{ov::element::Type_t::u2, {2, false, false, false, "uint2_t", "u2"}},
|
||||
{ov::element::Type_t::u3, {3, false, false, false, "uint3_t", "u3"}},
|
||||
{ov::element::Type_t::u4, {4, false, false, false, "uint4_t", "u4"}},
|
||||
{ov::element::Type_t::u6, {6, false, false, false, "uint6_t", "u6"}},
|
||||
{ov::element::Type_t::u8, {8, false, false, true, "uint8_t", "u8"}},
|
||||
{ov::element::Type_t::u16, {16, false, false, false, "uint16_t", "u16"}},
|
||||
{ov::element::Type_t::u32, {32, false, false, false, "uint32_t", "u32"}},
|
||||
|
|
@ -304,8 +319,14 @@ Type fundamental_type_for(const Type& type) {
|
|||
return from<element_type_traits<Type_t::i64>::value_type>();
|
||||
case Type_t::u1:
|
||||
return from<element_type_traits<Type_t::u1>::value_type>();
|
||||
case Type_t::u2:
|
||||
return from<element_type_traits<Type_t::u2>::value_type>();
|
||||
case Type_t::u3:
|
||||
return from<element_type_traits<Type_t::u3>::value_type>();
|
||||
case Type_t::u4:
|
||||
return from<element_type_traits<Type_t::u4>::value_type>();
|
||||
case Type_t::u6:
|
||||
return from<element_type_traits<Type_t::u6>::value_type>();
|
||||
case Type_t::u8:
|
||||
return from<element_type_traits<Type_t::u8>::value_type>();
|
||||
case Type_t::u16:
|
||||
|
|
@ -415,7 +436,10 @@ inline size_t compiler_byte_size(ov::element::Type_t et) {
|
|||
ET_CASE(i32);
|
||||
ET_CASE(i64);
|
||||
ET_CASE(u1);
|
||||
ET_CASE(u2);
|
||||
ET_CASE(u3);
|
||||
ET_CASE(u4);
|
||||
ET_CASE(u6);
|
||||
ET_CASE(u8);
|
||||
ET_CASE(u16);
|
||||
ET_CASE(u32);
|
||||
|
|
@ -451,7 +475,10 @@ OPENVINO_API EnumNames<element::Type_t>& EnumNames<element::Type_t>::get() {
|
|||
{"i32", element::Type_t::i32},
|
||||
{"i64", element::Type_t::i64},
|
||||
{"u1", element::Type_t::u1},
|
||||
{"u2", element::Type_t::u2},
|
||||
{"u3", element::Type_t::u3},
|
||||
{"u4", element::Type_t::u4},
|
||||
{"u6", element::Type_t::u6},
|
||||
{"u8", element::Type_t::u8},
|
||||
{"u16", element::Type_t::u16},
|
||||
{"u32", element::Type_t::u32},
|
||||
|
|
|
|||
|
|
@ -0,0 +1,478 @@
|
|||
// Copyright (C) 2018-2024 Intel Corporation
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
//
|
||||
|
||||
#include "openvino/core/type/element_iterator.hpp"
|
||||
|
||||
#include <gmock/gmock.h>
|
||||
|
||||
#include <array>
|
||||
|
||||
#include "openvino/runtime/tensor.hpp"
|
||||
|
||||
namespace ov {
|
||||
namespace test {
|
||||
|
||||
using testing::ElementsAre;
|
||||
using testing::ElementsAreArray;
|
||||
|
||||
namespace {
|
||||
constexpr size_t get_buffer_size(const size_t bit_width, const size_t num_of_elements) {
|
||||
return (num_of_elements * bit_width + 7) / 8;
|
||||
}
|
||||
} // namespace
|
||||
|
||||
// bits number in comments are counted [b7, b6, ..., b0]
|
||||
// ---- u1
|
||||
TEST(ElementIteratorTest, write_u1_data) {
|
||||
constexpr auto elements_count = 16;
|
||||
auto input = std::array<int8_t, elements_count>{0, 0, 0, 1, 0, 1, 1, 0, 1, 0, 1, 1, 0, 0, 1, 1};
|
||||
auto output = std::array<int8_t, get_buffer_size(1, elements_count)>{};
|
||||
auto iter = element::iterator<element::u1>(output.data());
|
||||
|
||||
std::copy(input.begin(), input.end(), iter);
|
||||
EXPECT_THAT(output, ElementsAre(0x16, 0xB3));
|
||||
}
|
||||
|
||||
TEST(ElementIteratorTest, read_const_u1_data) {
|
||||
constexpr auto elements_count = 16;
|
||||
constexpr auto input = std::array<int8_t, get_buffer_size(1, elements_count)>{0x21, static_cast<int8_t>(0xa3)};
|
||||
auto iter = element::iterator<element::u1>(input.data());
|
||||
|
||||
EXPECT_THAT(std::vector<int8_t>(iter, iter + elements_count),
|
||||
ElementsAre(0, 0, 1, 0, 0, 0, 0, 1, 1, 0, 1, 0, 0, 0, 1, 1));
|
||||
}
|
||||
|
||||
TEST(ElementIteratorTest, read_non_const_u1_data) {
|
||||
constexpr auto elements_count = 16;
|
||||
auto input = std::array<int8_t, get_buffer_size(1, elements_count)>{0x21, static_cast<int8_t>(0xa3)};
|
||||
auto iter = element::iterator<element::u1>(input.data());
|
||||
|
||||
EXPECT_THAT(std::vector<int8_t>(iter, iter + elements_count),
|
||||
ElementsAre(0, 0, 1, 0, 0, 0, 0, 1, 1, 0, 1, 0, 0, 0, 1, 1));
|
||||
}
|
||||
|
||||
TEST(ElementIteratorTest, read_u1_data_increment_decrement_iterator) {
|
||||
auto input = std::array<int8_t, 3>{0x32, static_cast<int8_t>(0xa3), 0x55};
|
||||
auto iter = element::iterator<element::u1>(input.data() + 1);
|
||||
|
||||
EXPECT_EQ(*iter--, 1); // 2nd byte bit7
|
||||
EXPECT_EQ(*iter++, 0); // 1st byte bit0
|
||||
EXPECT_EQ(*++iter, 0); // 2nd byte bit6
|
||||
EXPECT_EQ(*iter--, 0); // 2nd byte bit6
|
||||
EXPECT_EQ(*iter, 1); // 2nd byte bit7
|
||||
}
|
||||
|
||||
TEST(ElementIteratorTest, read_u1_data_iterator_with_offset) {
|
||||
auto input = std::array<int8_t, 3>{0x32, static_cast<int8_t>(0xa3), 0x41};
|
||||
auto iter = element::iterator<element::u1>(input.data() + 1);
|
||||
|
||||
EXPECT_EQ(*iter, 1); // 2nd byte bit7
|
||||
EXPECT_EQ(*(iter - 2), 1); // 1st byte bit1
|
||||
EXPECT_EQ(*(iter - 5), 1); // 1st byte bit4
|
||||
EXPECT_EQ(*(iter + 1), 0); // 2nd byte bit6
|
||||
EXPECT_EQ(*(iter + 8), 0); // 3rd byte bit7
|
||||
EXPECT_EQ(*(iter + 9), 1); // 3rd byte bit6
|
||||
EXPECT_EQ(*std::prev(iter, 1), 0); // 1st byte bit0
|
||||
EXPECT_EQ(*std::next(iter, 2), 1); // 2nd byte bit5
|
||||
}
|
||||
|
||||
TEST(ElementIteratorTest, read_u1_from_tensor) {
|
||||
auto input = std::array<int8_t, 4>{0x32, static_cast<int8_t>(0xa3), 0x41, 0x11};
|
||||
auto t = ov::Tensor(element::u1, Shape{2, 16}, input.data());
|
||||
auto iter = element::iterator<element::u1>(static_cast<int8_t*>(t.data(element::u1)));
|
||||
|
||||
EXPECT_THAT(
|
||||
std::vector<int8_t>(iter, iter + t.get_size()),
|
||||
ElementsAre(0, 0, 1, 1, 0, 0, 1, 0, 1, 0, 1, 0, 0, 0, 1, 1, 0, 1, 0, 0, 0, 0, 0, 1, 0, 0, 0, 1, 0, 0, 0, 1));
|
||||
}
|
||||
|
||||
TEST(ElementIteratorTest, u1_value_to_output_stream) {
|
||||
constexpr auto value = static_cast<int8_t>(0x80);
|
||||
auto iter = element::iterator<element::u1>(&value);
|
||||
|
||||
std::stringstream s;
|
||||
s << *iter;
|
||||
|
||||
EXPECT_EQ(s.str(), "1");
|
||||
}
|
||||
|
||||
// ---- u2
|
||||
TEST(ElementIteratorTest, write_u2_data) {
|
||||
constexpr auto elements_count = 16;
|
||||
auto input = std::array<int8_t, elements_count>{2, 0, 1, 3, 0, 0, 3, 3, 1, 2, 1, 2, 3, 2, 1, 0};
|
||||
auto output = std::array<int8_t, get_buffer_size(2, elements_count)>{};
|
||||
auto iter = element::iterator<element::u2>(output.data());
|
||||
|
||||
std::copy(input.begin(), input.end(), iter);
|
||||
|
||||
EXPECT_THAT(output, ElementsAre(0x87, 0x0f, 0x66, 0xe4));
|
||||
}
|
||||
|
||||
TEST(ElementIteratorTest, read_const_u2_data) {
|
||||
constexpr auto elements_count = 16;
|
||||
constexpr auto input = std::array<int8_t, get_buffer_size(2, elements_count)>{static_cast<int8_t>(0x87),
|
||||
0x0f,
|
||||
0x66,
|
||||
static_cast<int8_t>(0xe4)};
|
||||
auto iter = element::iterator<element::u2>(input.data());
|
||||
|
||||
EXPECT_THAT(std::vector<int8_t>(iter, iter + elements_count),
|
||||
ElementsAre(2, 0, 1, 3, 0, 0, 3, 3, 1, 2, 1, 2, 3, 2, 1, 0));
|
||||
}
|
||||
|
||||
TEST(ElementIteratorTest, read_non_const_u2_data) {
|
||||
constexpr auto elements_count = 16;
|
||||
auto input = std::array<int8_t, get_buffer_size(2, elements_count)>{static_cast<int8_t>(0x87),
|
||||
0x0f,
|
||||
0x66,
|
||||
static_cast<int8_t>(0xe4)};
|
||||
auto iter = element::iterator<element::u2>(input.data());
|
||||
|
||||
EXPECT_THAT(std::vector<int8_t>(iter, iter + elements_count),
|
||||
ElementsAre(2, 0, 1, 3, 0, 0, 3, 3, 1, 2, 1, 2, 3, 2, 1, 0));
|
||||
}
|
||||
|
||||
TEST(ElementIteratorTest, read_u2_data_increment_decrement_iterator) {
|
||||
auto input = std::array<int8_t, 2>{0x33, static_cast<int8_t>(0x93)};
|
||||
auto iter = element::iterator<element::u2>(input.data() + 1);
|
||||
|
||||
EXPECT_EQ(*iter--, 2); // 2nd byte 1st half-nibble
|
||||
EXPECT_EQ(*iter++, 3); // 1st byte 4th half-nibble
|
||||
EXPECT_EQ(*++iter, 1); // 2nd byte 2nd half-nibble
|
||||
EXPECT_EQ(*iter--, 1); // 2nd byte 2nd half-nibble
|
||||
EXPECT_EQ(*--iter, 3); // 1st byte 4th half-nibble
|
||||
}
|
||||
|
||||
TEST(ElementIteratorTest, read_u2_data_iterator_with_offset) {
|
||||
auto input = std::array<int8_t, 3>{0x43, static_cast<int8_t>(0x93), 0x41};
|
||||
auto iter = element::iterator<element::u2>(input.data() + 1);
|
||||
|
||||
EXPECT_EQ(*iter, 2); // 2nd byte 1st half-nibble
|
||||
EXPECT_EQ(*(iter - 3), 0); // 1st byte 2nd half-nibble
|
||||
EXPECT_EQ(*(iter - 4), 1); // 1st byte 1st half-nibble
|
||||
EXPECT_EQ(*(iter + 1), 1); // 2nd byte 2nd half-nibble
|
||||
EXPECT_EQ(*(iter + 7), 1); // 3rd byte 4th half-nibble
|
||||
EXPECT_EQ(*std::prev(iter, 1), 3); // 1st byte 4th half-nibble
|
||||
EXPECT_EQ(*std::next(iter, 2), 0); // 2nd byte 3rd half-nibble
|
||||
}
|
||||
|
||||
TEST(ElementIteratorTest, u2_value_to_output_stream) {
|
||||
constexpr auto value = static_cast<int8_t>(0x80);
|
||||
auto iter = element::iterator<element::u2>(&value);
|
||||
|
||||
std::stringstream s;
|
||||
s << *iter;
|
||||
|
||||
EXPECT_EQ(s.str(), "2");
|
||||
}
|
||||
|
||||
TEST(ElementIteratorTest, read_u2_from_tensor) {
|
||||
auto input = std::array<int8_t, 4>{0x32, static_cast<int8_t>(0xa3), 0x41, 0x11};
|
||||
auto t = ov::Tensor(element::u2, Shape{4, 4}, input.data());
|
||||
auto iter = element::iterator<element::u2>(static_cast<int8_t*>(t.data(element::u2)));
|
||||
|
||||
EXPECT_THAT(std::vector<int8_t>(iter, iter + t.get_size()),
|
||||
ElementsAre(0, 3, 0, 2, 2, 2, 0, 3, 1, 0, 0, 1, 0, 1, 0, 1));
|
||||
}
|
||||
|
||||
// --- u3
|
||||
TEST(ElementIteratorTest, write_u3_data) {
|
||||
constexpr auto elements_count = 8;
|
||||
auto input = std::array<int8_t, elements_count>{2, 3, 0, 1, 4, 5, 6, 7};
|
||||
auto output = std::array<int8_t, 3>{};
|
||||
auto iter = element::iterator<element::u3>(output.data());
|
||||
|
||||
std::copy(input.begin(), input.end(), iter);
|
||||
|
||||
EXPECT_THAT(output, ElementsAre(0b10110001, 0b00011011, 0b00001111));
|
||||
}
|
||||
|
||||
TEST(ElementIteratorTest, read_non_const_u3_data) {
|
||||
constexpr auto elements_count = 16;
|
||||
auto input = std::array<int8_t, 6>{0x7a, 0x6f, 0x55, static_cast<int8_t>(0xb1), 0x1b, 0x0f};
|
||||
auto iter = element::iterator<element::u3>(input.data());
|
||||
|
||||
EXPECT_THAT(std::vector<int8_t>(iter, iter + elements_count),
|
||||
ElementsAre(1, 7, 2, 6, 1, 6, 3, 7, 2, 3, 0, 1, 4, 5, 6, 7));
|
||||
}
|
||||
|
||||
TEST(ElementIteratorTest, read_const_u3_data) {
|
||||
constexpr auto elements_count = 8;
|
||||
constexpr auto input = std::array<int8_t, 3>{static_cast<int8_t>(0b10110001), 0b00011011, 0b00001111};
|
||||
auto iter = element::iterator<element::u3>(input.data());
|
||||
|
||||
EXPECT_THAT(std::vector<int8_t>(iter, iter + elements_count), ElementsAre(2, 3, 0, 1, 4, 5, 6, 7));
|
||||
}
|
||||
|
||||
TEST(ElementIteratorTest, read_u3_data_iterator_with_offset) {
|
||||
// Has values {1, 7, 2, 6, 1, 6, 3, 7, [2], 3, 0, 1, 4, 5, 6, 7}
|
||||
auto input = std::array<int8_t, 6>{0x7a, 0x6f, 0x55, static_cast<int8_t>(0xb1), 0x1b, 0x0f};
|
||||
auto iter = element::iterator<element::u3>(input.data() + 3);
|
||||
|
||||
EXPECT_EQ(*iter, 2);
|
||||
EXPECT_EQ(*(iter - 3), 6);
|
||||
EXPECT_EQ(*(iter - 4), 1);
|
||||
EXPECT_EQ(*(iter - 5), 6);
|
||||
EXPECT_EQ(*(iter + 1), 3);
|
||||
EXPECT_EQ(*(iter + 5), 5);
|
||||
EXPECT_EQ(*(iter + 7), 7);
|
||||
EXPECT_EQ(*std::prev(iter, 1), 7);
|
||||
EXPECT_EQ(*std::next(iter, 2), 0);
|
||||
}
|
||||
|
||||
TEST(ElementIteratorTest, read_u3_from_tensor) {
|
||||
// Has values {1, 7, 2, 6, 1, 6, 3, 7, [2], 3, 0, 1, 4, 5, 6, 7}
|
||||
auto input = std::array<int8_t, 6>{0x7a, 0x6f, 0x55, static_cast<int8_t>(0xb1), 0x1b, 0x0f};
|
||||
auto t = ov::Tensor(element::u3, Shape{4, 2, 2}, input.data());
|
||||
auto iter = element::iterator<element::u3>(static_cast<int8_t*>(t.data(element::u3)));
|
||||
|
||||
EXPECT_THAT(std::vector<int8_t>(iter, iter + t.get_size()),
|
||||
ElementsAre(1, 7, 2, 6, 1, 6, 3, 7, 2, 3, 0, 1, 4, 5, 6, 7));
|
||||
}
|
||||
|
||||
// --- u4
|
||||
// nibbles are counted as [n1, n0]
|
||||
TEST(ElementIteratorTest, write_u4_data) {
|
||||
constexpr auto elements_count = 16;
|
||||
auto input = std::array<int8_t, elements_count>{1, 2, 3, 10, 12, 15, 14, 4, 7, 9, 11, 13, 8, 0, 5, 6};
|
||||
auto output = std::array<int8_t, get_buffer_size(4, elements_count)>{};
|
||||
auto iter = element::iterator<element::u4>(output.data());
|
||||
|
||||
std::copy(input.begin(), input.end(), iter);
|
||||
|
||||
EXPECT_THAT(output, ElementsAre(0x21, 0xa3, 0xfc, 0x4e, 0x97, 0xdb, 0x08, 0x65));
|
||||
}
|
||||
|
||||
TEST(ElementIteratorTest, read_const_u4_data) {
|
||||
constexpr auto elements_count = 16;
|
||||
constexpr auto byte_size = get_buffer_size(4, elements_count);
|
||||
constexpr auto input = std::array<int8_t, byte_size>{0x12,
|
||||
0x3a,
|
||||
static_cast<int8_t>(0xcf),
|
||||
static_cast<int8_t>(0xe4),
|
||||
0x79,
|
||||
static_cast<int8_t>(0xbd),
|
||||
0x08,
|
||||
0x56};
|
||||
auto iter = element::iterator<element::u4>(input.data());
|
||||
|
||||
EXPECT_THAT(std::vector<int8_t>(iter, iter + elements_count),
|
||||
ElementsAre(2, 1, 10, 3, 15, 12, 4, 14, 9, 7, 13, 11, 8, 0, 6, 5));
|
||||
}
|
||||
|
||||
TEST(ElementIteratorTest, read_non_const_u4_data) {
|
||||
constexpr auto elements_count = 16;
|
||||
constexpr auto byte_size = get_buffer_size(4, elements_count);
|
||||
auto input = std::array<int8_t, byte_size>{0x12,
|
||||
0x3a,
|
||||
static_cast<int8_t>(0xcf),
|
||||
static_cast<int8_t>(0xe4),
|
||||
0x79,
|
||||
static_cast<int8_t>(0xbd),
|
||||
0x08,
|
||||
0x56};
|
||||
auto iter = element::iterator<element::u4>(input.data());
|
||||
|
||||
EXPECT_THAT(std::vector<int8_t>(iter, iter + elements_count),
|
||||
ElementsAre(2, 1, 10, 3, 15, 12, 4, 14, 9, 7, 13, 11, 8, 0, 6, 5));
|
||||
}
|
||||
|
||||
TEST(ElementIteratorTest, read_u4_data_increment_decrement_iterator) {
|
||||
auto input = std::array<int8_t, 3>{0x12, 0x3a};
|
||||
auto iter = element::iterator<element::u4>(input.data() + 1);
|
||||
|
||||
EXPECT_EQ(*iter--, 10); // 2nd byte 1st nibble
|
||||
EXPECT_EQ(*iter++, 1); // 1st byte 2nd nibble
|
||||
EXPECT_EQ(*++iter, 3); // 2nd byte 2nd nibble
|
||||
EXPECT_EQ(*iter--, 3); // 2nd byte 2nd nibble
|
||||
EXPECT_EQ(*--iter, 1); // 1st byte 2nd nibble
|
||||
}
|
||||
|
||||
TEST(ElementIteratorTest, read_u4_data_iterator_with_offset) {
|
||||
auto input = std::array<int8_t, 5>{0x42, 0x3a, 0x61, 0x79, 0x5b};
|
||||
auto iter = element::iterator<element::u4>(input.data() + 1);
|
||||
|
||||
EXPECT_EQ(*iter, 10); // 2nd byte 1st nibble
|
||||
EXPECT_EQ(*(iter - 2), 2); // 1st byte 1st nibble
|
||||
EXPECT_EQ(*(iter + 7), 5); // 5th byte 2nd nibble
|
||||
EXPECT_EQ(*(iter + 6), 11); // 2nd byte 1st nibble
|
||||
EXPECT_EQ(*(iter - 1), 4); // 1st byte 2nd nibble
|
||||
EXPECT_EQ(*std::prev(iter, 1), 4); // 1st byte 2nd nibble
|
||||
EXPECT_EQ(*std::next(iter, 2), 1); // 3rd byte 1st nibble
|
||||
}
|
||||
|
||||
TEST(ElementIteratorTest, read_u4_from_tensor) {
|
||||
auto input = std::array<int8_t, 5>{0x42, 0x3a, 0x61, 0x79, 0x5b};
|
||||
auto t = ov::Tensor(element::u4, Shape{5, 2}, input.data());
|
||||
auto iter = element::iterator<element::u4>(static_cast<int8_t*>(t.data(element::u4)));
|
||||
|
||||
EXPECT_THAT(std::vector<int8_t>(iter, iter + t.get_size()), ElementsAre(2, 4, 10, 3, 1, 6, 9, 7, 11, 5));
|
||||
}
|
||||
|
||||
// --- i4
|
||||
// nibbles are counted as [n1, n0]
|
||||
TEST(ElementIteratorTest, write_i4_data) {
|
||||
constexpr auto elements_count = 16;
|
||||
auto input = std::array<int8_t, elements_count>{1, 2, 3, -6, -4, -1, -2, 4, 7, -7, -5, -3, -8, 0, 5, 6};
|
||||
auto output = std::array<int8_t, get_buffer_size(4, elements_count)>{};
|
||||
auto iter = element::iterator<element::i4>(output.data());
|
||||
|
||||
std::copy(input.begin(), input.end(), iter);
|
||||
|
||||
EXPECT_THAT(output, ElementsAre(0x21, 0xa3, 0xfc, 0x4e, 0x97, 0xdb, 0x08, 0x65));
|
||||
}
|
||||
|
||||
TEST(ElementIteratorTest, read_const_i4_data) {
|
||||
constexpr auto elements_count = 16;
|
||||
constexpr auto byte_size = get_buffer_size(4, elements_count);
|
||||
constexpr auto input = std::array<int8_t, byte_size>{0x12,
|
||||
0x3a,
|
||||
static_cast<int8_t>(0xcf),
|
||||
static_cast<int8_t>(0xe4),
|
||||
0x79,
|
||||
static_cast<int8_t>(0xbd),
|
||||
0x08,
|
||||
0x56};
|
||||
auto iter = element::iterator<element::i4>(input.data());
|
||||
|
||||
EXPECT_THAT(std::vector<int8_t>(iter, iter + elements_count),
|
||||
ElementsAre(2, 1, -6, 3, -1, -4, 4, -2, -7, 7, -3, -5, -8, 0, 6, 5));
|
||||
}
|
||||
|
||||
TEST(ElementIteratorTest, read_non_const_i4_data) {
|
||||
constexpr auto elements_count = 16;
|
||||
constexpr auto byte_size = get_buffer_size(4, elements_count);
|
||||
auto input = std::array<int8_t, byte_size>{0x12,
|
||||
0x3a,
|
||||
static_cast<int8_t>(0xcf),
|
||||
static_cast<int8_t>(0xe4),
|
||||
0x79,
|
||||
static_cast<int8_t>(0xbd),
|
||||
0x08,
|
||||
0x56};
|
||||
auto iter = element::iterator<element::i4>(input.data());
|
||||
|
||||
EXPECT_THAT(std::vector<int8_t>(iter, iter + elements_count),
|
||||
ElementsAre(2, 1, -6, 3, -1, -4, 4, -2, -7, 7, -3, -5, -8, 0, 6, 5));
|
||||
}
|
||||
|
||||
TEST(ElementIteratorTest, read_i4_data_increment_decrement_iterator) {
|
||||
auto input = std::array<int8_t, 2>{0x12, 0x3a};
|
||||
auto iter = element::iterator<element::i4>(input.data() + 1);
|
||||
|
||||
EXPECT_EQ(*iter--, -6); // 2nd byte 1st nibble
|
||||
EXPECT_EQ(*iter++, 1); // 1st byte 2nd nibble
|
||||
EXPECT_EQ(*++iter, 3); // 2nd byte 2nd nibble
|
||||
EXPECT_EQ(*iter--, 3); // 2nd byte 2nd nibble
|
||||
EXPECT_EQ(*--iter, 1); // 1st byte 2nd nibble
|
||||
}
|
||||
|
||||
TEST(ElementIteratorTest, read_i4_data_iterator_with_offset) {
|
||||
auto input = std::array<int8_t, 5>{0x42, 0x3a, 0x61, 0x79, 0x5b};
|
||||
auto iter = element::iterator<element::i4>(input.data() + 1);
|
||||
|
||||
EXPECT_EQ(*iter, -6); // 2nd byte 1st nibble
|
||||
EXPECT_EQ(*(iter - 2), 2); // 1st byte 1st nibble
|
||||
EXPECT_EQ(*(iter + 7), 5); // 5th byte 2nd nibble
|
||||
EXPECT_EQ(*(iter + 6), -5); // 2nd byte 1st nibble
|
||||
EXPECT_EQ(*(iter - 1), 4); // 1st byte 2nd nibble
|
||||
EXPECT_EQ(*std::prev(iter, 1), 4); // 1st byte 2nd nibble
|
||||
EXPECT_EQ(*std::next(iter, 2), 1); // 3rd byte 1st nibble
|
||||
}
|
||||
|
||||
TEST(ElementIteratorTest, i4_value_to_output_stream) {
|
||||
constexpr auto value = static_cast<int8_t>(0x19);
|
||||
auto iter = element::iterator<element::i4>(&value);
|
||||
|
||||
std::stringstream s;
|
||||
s << *iter;
|
||||
|
||||
EXPECT_EQ(s.str(), "-7");
|
||||
}
|
||||
|
||||
TEST(ElementIteratorTest, read_i4_from_tensor) {
|
||||
auto input = std::array<int8_t, 5>{0x42, 0x3a, 0x61, 0x79, 0x5b};
|
||||
auto t = ov::Tensor(element::i4, Shape{10, 1, 1}, input.data());
|
||||
auto iter = element::iterator<element::i4>(static_cast<int8_t*>(t.data(element::i4)));
|
||||
|
||||
EXPECT_THAT(std::vector<int8_t>(iter, iter + t.get_size()), ElementsAre(2, 4, -6, 3, 1, 6, -7, 7, -5, 5));
|
||||
}
|
||||
|
||||
// --- u6
|
||||
TEST(ElementIteratorTest, write_u6_data) {
|
||||
constexpr auto elements_count = 8;
|
||||
auto input = std::array<int8_t, elements_count>{2, 1, 0, 3, 18, 49, 35, 16};
|
||||
auto output = std::array<int8_t, 6>{};
|
||||
auto iter = element::iterator<element::u6>(output.data());
|
||||
|
||||
std::copy(input.begin(), input.end(), iter);
|
||||
|
||||
EXPECT_THAT(output, ElementsAre(0x21, 0x03, 0x00, 0x21, 0x30, 0x79));
|
||||
}
|
||||
|
||||
TEST(ElementIteratorTest, read_non_const_u6_data) {
|
||||
constexpr auto elements_count = 8;
|
||||
auto input = std::array<int8_t, 6>{0x21, 0x03, 0x00, 0x21, 0x30, 0x79};
|
||||
auto iter = element::iterator<element::u6>(input.data());
|
||||
|
||||
EXPECT_THAT(std::vector<int8_t>(iter, iter + elements_count), ElementsAre(2, 1, 0, 3, 18, 49, 35, 16));
|
||||
}
|
||||
|
||||
TEST(ElementIteratorTest, read_const_u6_data) {
|
||||
constexpr auto elements_count = 8;
|
||||
constexpr auto input = std::array<int8_t, 6>{0x21, 0x03, 0x00, 0x21, 0x30, 0x79};
|
||||
auto iter = element::iterator<element::u6>(input.data());
|
||||
|
||||
EXPECT_THAT(std::vector<int8_t>(iter, iter + elements_count), ElementsAre(2, 1, 0, 3, 18, 49, 35, 16));
|
||||
}
|
||||
|
||||
TEST(ElementIteratorTest, read_u6_data_increment_decrement_iterator) {
|
||||
// Has values {1, 2, 3, 10, [3], 8, 7, 2}
|
||||
auto input = std::array<int8_t, 6>{0x12, 0x3a, 0x00, 0x38, 0x72, 0x00};
|
||||
auto iter = element::iterator<element::u6>(input.data() + 3);
|
||||
|
||||
EXPECT_EQ(*iter--, 3);
|
||||
EXPECT_EQ(*iter++, 10);
|
||||
EXPECT_EQ(*++iter, 8);
|
||||
EXPECT_EQ(*iter--, 8);
|
||||
EXPECT_EQ(*--iter, 10);
|
||||
}
|
||||
|
||||
TEST(ElementIteratorTest, read_u6_data_iterator_with_offset) {
|
||||
// Has values {1, 2, 3, 10, [3], 8, 7, 2, 1, 42, 4, 20}
|
||||
auto input = std::array<int8_t, 9>{0x12, 0x3a, 0x00, 0x38, 0x72, 0x00, 0x1a, 0x44, 0x21};
|
||||
auto iter = element::iterator<element::u6>(input.data() + 3);
|
||||
|
||||
EXPECT_EQ(*iter, 3);
|
||||
EXPECT_EQ(*(iter - 3), 2);
|
||||
EXPECT_EQ(*(iter - 4), 1);
|
||||
EXPECT_EQ(*(iter - 2), 3);
|
||||
EXPECT_EQ(*(iter + 1), 8);
|
||||
EXPECT_EQ(*(iter + 5), 42);
|
||||
EXPECT_EQ(*(iter + 7), 20);
|
||||
EXPECT_EQ(*std::prev(iter, 1), 10);
|
||||
EXPECT_EQ(*std::next(iter, 2), 7);
|
||||
}
|
||||
|
||||
TEST(ElementIteratorTest, u6_value_to_output_stream) {
|
||||
auto input = std::array<int8_t, 3>{0x12, 0x3a, 0x00};
|
||||
auto iter = element::iterator<element::u6>(input.data());
|
||||
|
||||
std::stringstream s;
|
||||
s << *iter;
|
||||
|
||||
EXPECT_EQ(s.str(), "1");
|
||||
}
|
||||
|
||||
TEST(ElementIteratorTest, read_u6_from_tensor) {
|
||||
// Has values {1, 2, 3, 10, 3, 8, 7, 2, 1, 42, 4, 20}
|
||||
auto input = std::array<int8_t, 9>{0x12, 0x3a, 0x00, 0x38, 0x72, 0x00, 0x1a, 0x44, 0x21};
|
||||
auto t = ov::Tensor(element::u6, Shape{4, 1, 3}, input.data());
|
||||
auto iter = element::iterator<element::u6>(static_cast<int8_t*>(t.data(element::u6)));
|
||||
|
||||
EXPECT_THAT(std::vector<int8_t>(iter, iter + t.get_size()), ElementsAre(1, 2, 3, 10, 3, 8, 7, 2, 1, 42, 4, 20));
|
||||
}
|
||||
|
||||
} // namespace test
|
||||
} // namespace ov
|
||||
|
|
@ -81,9 +81,10 @@ private:
|
|||
CASE(f8e4m3)
|
||||
CASE(f8e5m2)
|
||||
CASE(string)
|
||||
default:
|
||||
return _undefined;
|
||||
}
|
||||
#undef CASE
|
||||
return _undefined;
|
||||
}
|
||||
};
|
||||
|
||||
|
|
|
|||
|
|
@ -255,7 +255,7 @@ public:
|
|||
const auto primitive_hash = primitve->hash();
|
||||
const auto params_hash = prim_inst->get_impl_params()->hash();
|
||||
ASSERT_EQ(primitive_hash, 4135863035456568493UL);
|
||||
ASSERT_EQ(params_hash, 5990757629995899044UL);
|
||||
ASSERT_EQ(params_hash, 11563701278302723583UL);
|
||||
}
|
||||
};
|
||||
|
||||
|
|
|
|||
Loading…
Reference in New Issue