@@ -205,18 +205,22 @@ struct subbyte_iterator
|
||||
private:
|
||||
|
||||
template <class, class> friend struct swizzle_ptr;
|
||||
template <class U> friend CUTE_HOST_DEVICE constexpr U* raw_pointer_cast(subbyte_iterator<U> const&);
|
||||
template <class N, class U> friend CUTE_HOST_DEVICE constexpr auto recast_ptr(subbyte_iterator<U> const&);
|
||||
template <class U> friend CUTE_HOST_DEVICE void print(subbyte_iterator<U> const&);
|
||||
|
||||
// Pointer to storage element
|
||||
storage_type* ptr_ = nullptr;
|
||||
storage_type* ptr_;
|
||||
|
||||
// Bit index of value_type starting position within storage_type element.
|
||||
// RI: 0 <= idx_ < sizeof_bit<storage_type>
|
||||
uint8_t idx_ = 0;
|
||||
uint8_t idx_;
|
||||
|
||||
public:
|
||||
|
||||
// Ctor
|
||||
subbyte_iterator() = default;
|
||||
// Default Ctor
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
subbyte_iterator() : ptr_{nullptr}, idx_{0} {};
|
||||
|
||||
// Ctor
|
||||
template <class PointerType>
|
||||
@@ -286,43 +290,48 @@ public:
|
||||
return x.ptr_ == y.ptr_ && x.idx_ == y.idx_;
|
||||
}
|
||||
CUTE_HOST_DEVICE constexpr friend
|
||||
bool operator!=(subbyte_iterator const& x, subbyte_iterator const& y) { return !(x == y); }
|
||||
CUTE_HOST_DEVICE constexpr friend
|
||||
bool operator< (subbyte_iterator const& x, subbyte_iterator const& y) {
|
||||
return x.ptr_ < y.ptr_ || (x.ptr_ == y.ptr_ && x.idx_ < y.idx_);
|
||||
}
|
||||
CUTE_HOST_DEVICE constexpr friend
|
||||
bool operator!=(subbyte_iterator const& x, subbyte_iterator const& y) { return !(x == y); }
|
||||
CUTE_HOST_DEVICE constexpr friend
|
||||
bool operator<=(subbyte_iterator const& x, subbyte_iterator const& y) { return !(y < x); }
|
||||
CUTE_HOST_DEVICE constexpr friend
|
||||
bool operator> (subbyte_iterator const& x, subbyte_iterator const& y) { return (y < x); }
|
||||
CUTE_HOST_DEVICE constexpr friend
|
||||
bool operator>=(subbyte_iterator const& x, subbyte_iterator const& y) { return !(x < y); }
|
||||
|
||||
// Conversion to raw pointer with loss of subbyte index
|
||||
CUTE_HOST_DEVICE constexpr friend
|
||||
T* raw_pointer_cast(subbyte_iterator const& x) {
|
||||
assert(x.idx_ == 0);
|
||||
return reinterpret_cast<T*>(x.ptr_);
|
||||
}
|
||||
|
||||
// Conversion to NewT_ with possible loss of subbyte index
|
||||
template <class NewT_>
|
||||
CUTE_HOST_DEVICE constexpr friend
|
||||
auto recast_ptr(subbyte_iterator const& x) {
|
||||
using NewT = conditional_t<(is_const_v<T>), NewT_ const, NewT_>;
|
||||
if constexpr (cute::is_subbyte_v<NewT>) { // Making subbyte_iter, preserve the subbyte idx
|
||||
return subbyte_iterator<NewT>(x.ptr_, x.idx_);
|
||||
} else { // Not subbyte, assume/assert subbyte idx 0
|
||||
return reinterpret_cast<NewT*>(raw_pointer_cast(x));
|
||||
}
|
||||
CUTE_GCC_UNREACHABLE;
|
||||
}
|
||||
|
||||
CUTE_HOST_DEVICE friend void print(subbyte_iterator x) {
|
||||
printf("subptr[%db](%p.%u)", int(sizeof_bits_v<T>), x.ptr_, x.idx_);
|
||||
}
|
||||
};
|
||||
|
||||
// Conversion to raw pointer with loss of subbyte index
|
||||
template <class T>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
T*
|
||||
raw_pointer_cast(subbyte_iterator<T> const& x) {
|
||||
assert(x.idx_ == 0);
|
||||
return reinterpret_cast<T*>(x.ptr_);
|
||||
}
|
||||
|
||||
// Conversion to NewT_ with possible loss of subbyte index
|
||||
template <class NewT_, class T>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
recast_ptr(subbyte_iterator<T> const& x) {
|
||||
using NewT = conditional_t<(is_const_v<T>), NewT_ const, NewT_>;
|
||||
if constexpr (cute::is_subbyte_v<NewT>) { // Making subbyte_iter, preserve the subbyte idx
|
||||
return subbyte_iterator<NewT>(x.ptr_, x.idx_);
|
||||
} else { // Not subbyte, assume/assert subbyte idx 0
|
||||
return reinterpret_cast<NewT*>(raw_pointer_cast(x));
|
||||
}
|
||||
CUTE_GCC_UNREACHABLE;
|
||||
}
|
||||
|
||||
template <class T>
|
||||
CUTE_HOST_DEVICE void
|
||||
print(subbyte_iterator<T> const& x) {
|
||||
printf("subptr[%db](%p.%u)", int(sizeof_bits_v<T>), x.ptr_, x.idx_);
|
||||
}
|
||||
|
||||
//
|
||||
// array_subbyte
|
||||
// Statically sized array for non-byte-aligned data types
|
||||
@@ -365,17 +374,6 @@ private:
|
||||
|
||||
public:
|
||||
|
||||
constexpr
|
||||
array_subbyte() = default;
|
||||
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
array_subbyte(array_subbyte const& x) {
|
||||
CUTE_UNROLL
|
||||
for (size_type i = 0; i < StorageElements; ++i) {
|
||||
storage[i] = x.storage[i];
|
||||
}
|
||||
}
|
||||
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
size_type size() const {
|
||||
return N;
|
||||
@@ -448,25 +446,16 @@ public:
|
||||
return at(N-1);
|
||||
}
|
||||
|
||||
// In analogy to std::vector<bool>::data(), these functions are deleted to prevent bugs.
|
||||
// Instead, prefer
|
||||
// auto* data = raw_pointer_cast(my_subbyte_array.begin());
|
||||
// where the type of auto* is implementation-defined and
|
||||
// with the knowledge that [data, data + my_subbyte_array.size()) may not be a valid range.
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
pointer data() {
|
||||
return reinterpret_cast<pointer>(storage);
|
||||
}
|
||||
pointer data() = delete;
|
||||
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
const_pointer data() const {
|
||||
return reinterpret_cast<const_pointer>(storage);
|
||||
}
|
||||
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
storage_type* raw_data() {
|
||||
return storage;
|
||||
}
|
||||
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
storage_type const* raw_data() const {
|
||||
return storage;
|
||||
}
|
||||
const_pointer data() const = delete;
|
||||
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
iterator begin() {
|
||||
|
||||
Reference in New Issue
Block a user