CUTLASS 2.10 bug fixes and minor updates. (#626)
This commit is contained in:
@@ -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:
|
||||
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -91,11 +91,11 @@ public:
|
||||
Convert(Params const ¶ms = 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
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user