CUTLASS 3.3.0 (#1167)
* Release 3.3.0 Adds support for mixed precision GEMMs On Hopper and Ampere Adds support for < 16B aligned GEMMs on Hopper Enhancements to EVT Enhancements to Python interface Enhancements to Sub-byte type handling in CuTe Several other bug-fixes and performance improvements. * minor doc update
This commit is contained in:
@@ -72,16 +72,10 @@ struct bit_field
|
||||
// Number of bits in data_[idx] used for NumBits if straddling, else 0
|
||||
static constexpr uint32_t bit_hi = (idx + 1 < N) ? (storage_type_bits - bit_lo) : 0;
|
||||
|
||||
private:
|
||||
// MSVC issues warning C4293 ("shift count negative or too big, undefined behavior")
|
||||
// if we use NumBits directly in the shift expression, even if the shift occurs
|
||||
// in the branch of a ternary expression where NumBits is known to be less than
|
||||
// the number of bits of the value being shifted.
|
||||
static constexpr uint32_t MollifiedNumBits = NumBits > 63u ? 63u : NumBits;
|
||||
public:
|
||||
|
||||
// NumBits mask
|
||||
static constexpr value_type mask = (NumBits < 64u) ? ((uint64_t(1) << MollifiedNumBits) - 1) : uint64_t(-1);
|
||||
static constexpr value_type mask = value_type(uint64_t(-1) >> (64u - NumBits));
|
||||
// NumBits mask for BitStart
|
||||
static constexpr storage_type mask_lo = storage_type(mask) << bit_lo;
|
||||
// NumBits mask for leftover bits in data_[idx+1] if straddling, else 0
|
||||
@@ -93,7 +87,7 @@ public:
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
value_type get() const {
|
||||
storage_type result = (data_[idx] & mask_lo) >> bit_lo;
|
||||
if constexpr (bit_hi) {
|
||||
if constexpr (bit_hi != 0) {
|
||||
result |= (data_[idx+1] & mask_hi) << bit_hi;
|
||||
}
|
||||
return static_cast<value_type>(result);
|
||||
@@ -104,7 +98,7 @@ public:
|
||||
void set(value_type x) {
|
||||
storage_type item = static_cast<storage_type>(x & mask);
|
||||
data_[idx] = static_cast<storage_type>((data_[idx] & ~mask_lo) | (item << bit_lo));
|
||||
if constexpr (bit_hi) {
|
||||
if constexpr (bit_hi != 0) {
|
||||
data_[idx+1] = static_cast<storage_type>((data_[idx+1] & ~mask_hi) | (item >> bit_hi));
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user