CUTLASS 2.10 bug fixes and minor updates. (#626)

This commit is contained in:
Andrew Kerr
2022-09-15 16:20:33 -04:00
committed by GitHub
parent 2cc2c7ba1f
commit fc9ebc645b
5 changed files with 63 additions and 8 deletions

View File

@@ -80,9 +80,9 @@ public:
typedef value_type *pointer;
typedef value_type const * const_pointer;
using Array = Array<T, N>;
using reference = typename Array::reference;
using const_reference = typename Array::const_reference;
using ArrayType = Array<T, N>;
using reference = typename ArrayType::reference;
using const_reference = typename ArrayType::const_reference;
public:

View File

@@ -633,7 +633,7 @@ public:
CUTLASS_PRAGMA_UNROLL
for (int v_idx = 0; v_idx < kAccessesPerVector; ++v_idx) {
clear_mask(v_idx, filter_k_ >= problem_size.K);
clear_mask(v_idx, filter_k_ + v_idx * AccessType::kElements >= problem_size.K);
}
set_iteration_index(0);

View File

@@ -91,11 +91,11 @@ public:
Convert(Params const &params = Params()) {
}
/// Functionally required for serial reduction in the epilogue
CUTLASS_HOST_DEVICE
void set_k_partition(int k_partition, int k_partition_count) {
}
/// Returns true if source is needed based on state of runtime arguments

View File

@@ -671,7 +671,7 @@ public:
state_[1] = 0;
++state_[2];
byte_pointer_ += params_.advance_cluster;
store_byte_pointer_ += params_.advance_group;
store_byte_pointer_ += params_.advance_cluster;
thread_start_row_ += ThreadMap::Count::kGroup *
ThreadMap::Shape::kGroup * ThreadMap::Count::kRow * ThreadMap::Shape::kRow;
@@ -679,7 +679,7 @@ public:
if (state_[2] == ThreadMap::Count::kCluster) {
state_[2] = 0;
byte_pointer_ += params_.advance_tile;
store_byte_pointer_ += params_.advance_group;
store_byte_pointer_ += params_.advance_tile;
}
}
}