CUTLASS 2.2 (#96)
Adds support for NVIDIA Ampere Architecture features. CUDA 11 Toolkit recommended.
This commit is contained in:
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* Copyright (c) 2017-2020, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without
|
||||
*modification, are permitted provided that the following conditions are met:
|
||||
@@ -216,6 +216,7 @@ class PredicatedTileAccessIterator<Shape_, Element_, layout::PitchLinear,
|
||||
predicates_[i] = 0u;
|
||||
}
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int access_idx = 0; access_idx < ThreadMap::Iterations::kCount * kAccessesPerVector; ++access_idx) {
|
||||
|
||||
int s = access_idx / (ThreadMap::Iterations::kContiguous * kAccessesPerVector);
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* Copyright (c) 2017-2020, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without
|
||||
*modification, are permitted provided that the following conditions are met:
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* Copyright (c) 2017-2020, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without modification, are permitted
|
||||
* provided that the following conditions are met:
|
||||
@@ -319,9 +319,11 @@ class PredicatedTileIterator<Shape_, Element_, layout::PitchLinear, AdvanceRank,
|
||||
|
||||
AccessType const *access_ptr = reinterpret_cast<AccessType const *>(byte_ptr);
|
||||
|
||||
if (address_iterator_.valid()) {
|
||||
frag_ptr[idx] = *access_ptr;
|
||||
}
|
||||
cutlass::arch::global_load<AccessType,
|
||||
sizeof(AccessType)
|
||||
>(
|
||||
frag_ptr[idx], access_ptr, address_iterator_.valid());
|
||||
|
||||
++address_iterator_;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* Copyright (c) 2017-2020, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without modification, are permitted
|
||||
* provided that the following conditions are met:
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* Copyright (c) 2017-2020, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without
|
||||
*modification, are permitted provided that the following conditions are met:
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* Copyright (c) 2017-2020, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without modification, are permitted
|
||||
* provided that the following conditions are met:
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* Copyright (c) 2017-2020, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without modification, are permitted
|
||||
* provided that the following conditions are met:
|
||||
@@ -798,6 +798,269 @@ class RegularTileAccessIterator<Shape_, Element_,
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Tile iterator specialized for k interleaved arrangements for TensorOps
|
||||
///
|
||||
///
|
||||
/// Satisfies: ForwardTileIteratorConcept |
|
||||
/// ReadableContiguousTileIteratorConcept |
|
||||
/// WriteableContiguousTileIteratorConcept
|
||||
///
|
||||
template <typename Shape_, typename Element_, int AdvanceRank, typename ThreadMap_, int InterleavedK, int Alignment>
|
||||
class RegularTileAccessIterator<
|
||||
Shape_, Element_,
|
||||
layout::TensorOpMultiplicandRowMajorInterleaved<sizeof_bits<Element_>::value,
|
||||
InterleavedK>,
|
||||
AdvanceRank, ThreadMap_, Alignment> {
|
||||
public:
|
||||
static_assert(
|
||||
AdvanceRank == 0 || AdvanceRank == 1,
|
||||
"Specialization for pitch-linear iterator may along advance along the "
|
||||
"contiguous(rank=0) or strided(rank=1) dimension.");
|
||||
|
||||
using Shape = Shape_;
|
||||
using Element = Element_;
|
||||
using Layout =
|
||||
layout::TensorOpMultiplicandRowMajorInterleaved<sizeof_bits<Element_>::value,
|
||||
InterleavedK>;
|
||||
static int const kAdvanceRank = AdvanceRank;
|
||||
static int const kAlignment = Alignment;
|
||||
|
||||
using Index = typename Layout::Index;
|
||||
using LongIndex = typename Layout::LongIndex;
|
||||
|
||||
using TensorRef = TensorRef<Element, Layout>;
|
||||
using TensorCoord = typename Layout::TensorCoord;
|
||||
|
||||
using ThreadMap = ThreadMap_;
|
||||
|
||||
/// Internal details made public to facilitate introspection
|
||||
struct Detail {
|
||||
/// This iterator is specialized for an access size that is 128 bits in
|
||||
/// length.
|
||||
static int const kAccessSizeInBits = 128;
|
||||
|
||||
static_assert(sizeof_bits<Element_>::value * ThreadMap::kElementsPerAccess ==
|
||||
kAccessSizeInBits,
|
||||
"This iterator requires a policy whose access size is 128bs");
|
||||
};
|
||||
|
||||
private:
|
||||
|
||||
/// Element type per access
|
||||
using AccessType = Array<Element, Layout::kElementsPerAccess>;
|
||||
|
||||
private:
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
/// Internal pointer to first access of tile
|
||||
AccessType *pointer_;
|
||||
|
||||
/// Internal byte offset
|
||||
Index byte_offset_;
|
||||
|
||||
/// Iteration in the contiguous dimension
|
||||
int iteration_contiguous_;
|
||||
|
||||
/// Iteration in the strided dimension
|
||||
int iteration_strided_;
|
||||
|
||||
public:
|
||||
/// Construct a TileIterator with zero threadblock offset
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileAccessIterator(TensorRef ref, ///< Pointer to start of tensor
|
||||
int thread_id ///< ID of each participating thread
|
||||
)
|
||||
: byte_offset_(0) {
|
||||
layout::PitchLinearCoord thread_offset_base =
|
||||
ThreadMap::initial_offset(thread_id);
|
||||
|
||||
// initialize pointer
|
||||
pointer_ = reinterpret_cast<AccessType *>(
|
||||
ref.data() + ref.offset(thread_offset_base));
|
||||
|
||||
set_iteration_index(0);
|
||||
}
|
||||
|
||||
/// Overrides the internal iteration index
|
||||
CUTLASS_HOST_DEVICE
|
||||
void set_iteration_index(int index) {
|
||||
iteration_contiguous_ = index % ThreadMap::Iterations::kContiguous;
|
||||
iteration_strided_ = index / ThreadMap::Iterations::kContiguous;
|
||||
}
|
||||
|
||||
/// Adds a pointer offset in units of Element
|
||||
CUTLASS_HOST_DEVICE
|
||||
void add_pointer_offset(LongIndex pointer_offset) {
|
||||
byte_offset_ += pointer_offset * sizeof(Element);
|
||||
}
|
||||
|
||||
/// Returns a pointer
|
||||
CUTLASS_HOST_DEVICE
|
||||
AccessType *get() const {
|
||||
AccessType *access_ptr = pointer_;
|
||||
|
||||
int access_offset =
|
||||
(iteration_strided_ * ThreadMap::Delta::kStrided * Layout::kInterleavedK +
|
||||
iteration_contiguous_ * ThreadMap::Delta::kContiguous) / ThreadMap::kElementsPerAccess;
|
||||
|
||||
char *access_byte_ptr =
|
||||
reinterpret_cast<char *>(access_ptr + access_offset);
|
||||
|
||||
return reinterpret_cast<AccessType *>(access_byte_ptr + byte_offset_);
|
||||
}
|
||||
|
||||
/// Advances to the next tile in memory.
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileAccessIterator &operator++() {
|
||||
++iteration_contiguous_;
|
||||
|
||||
if (iteration_contiguous_ < ThreadMap::Iterations::kContiguous)
|
||||
return *this;
|
||||
|
||||
// Enter here only if (iteration_contiguous_ ==
|
||||
// ThreadMap::Iteration::kContiguous)
|
||||
iteration_contiguous_ = 0;
|
||||
++iteration_strided_;
|
||||
|
||||
if (iteration_strided_ < ThreadMap::Iterations::kStrided) {
|
||||
return *this;
|
||||
}
|
||||
|
||||
// Enter here only if (iteration_strided_ == ThreadMap::Iteration::kStrided)
|
||||
// which means we enter the next tile.
|
||||
iteration_strided_ = 0;
|
||||
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances to the next tile in memory.
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileAccessIterator operator++(int) {
|
||||
RegularTileAccessIterator prev(*this);
|
||||
this->operator++();
|
||||
|
||||
return prev;
|
||||
}
|
||||
|
||||
/// Adds a tile offset
|
||||
CUTLASS_DEVICE
|
||||
void add_tile_offset(TensorCoord const &coord) {
|
||||
add_pointer_offset(coord.contiguous() * Shape::kCount);
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Tile iterator specialized for k interleaved arrangements for TensorOps
|
||||
///
|
||||
///
|
||||
/// Satisfies: ForwardTileIteratorConcept |
|
||||
/// ReadableContiguousTileIteratorConcept |
|
||||
/// WriteableContiguousTileIteratorConcept
|
||||
///
|
||||
|
||||
template <typename Shape_, typename Element_, int AdvanceRank, typename ThreadMap_, int InterleavedK, int Alignment>
|
||||
class RegularTileAccessIterator<
|
||||
Shape_, Element_,
|
||||
layout::TensorOpMultiplicandColumnMajorInterleaved<sizeof_bits<Element_>::value,
|
||||
InterleavedK>,
|
||||
AdvanceRank, ThreadMap_, Alignment> {
|
||||
|
||||
public:
|
||||
static_assert(
|
||||
AdvanceRank == 0 || AdvanceRank == 1,
|
||||
"Specialization for pitch-linear iterator may along advance along the "
|
||||
"contiguous(rank=0) or strided(rank=1) dimension.");
|
||||
|
||||
using Shape = Shape_;
|
||||
using Element = Element_;
|
||||
using Layout =
|
||||
layout::TensorOpMultiplicandColumnMajorInterleaved<sizeof_bits<Element_>::value,
|
||||
InterleavedK>;
|
||||
static int const kAdvanceRank = AdvanceRank;
|
||||
static int const kAlignment = Alignment;
|
||||
|
||||
using Index = typename Layout::Index;
|
||||
using LongIndex = typename Layout::LongIndex;
|
||||
|
||||
using TensorRef = TensorRef<Element, Layout>;
|
||||
using TensorCoord = typename Layout::TensorCoord;
|
||||
|
||||
using ThreadMap = ThreadMap_;
|
||||
|
||||
/// Underlying iterator type
|
||||
using UnderlyingIterator = RegularTileAccessIterator<
|
||||
cutlass::MatrixShape<Shape::kColumn, Shape::kRow>,
|
||||
Element,
|
||||
layout::TensorOpMultiplicandRowMajorInterleaved<sizeof_bits<Element_>::value, InterleavedK>,
|
||||
(kAdvanceRank == 1 ? 0 : 1),
|
||||
ThreadMap
|
||||
>;
|
||||
|
||||
private:
|
||||
|
||||
/// Element type per access
|
||||
using AccessType = Array<Element, Layout::kElementsPerAccess>;
|
||||
|
||||
private:
|
||||
|
||||
/// Underlying iterator
|
||||
UnderlyingIterator iterator_;
|
||||
|
||||
public:
|
||||
/// Construct a TileIterator with zero threadblock offset
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileAccessIterator(TensorRef ref, ///< Pointer to start of tensor
|
||||
int thread_id ///< ID of each participating thread
|
||||
)
|
||||
: iterator_({ref.data(), ref.stride()}, thread_id) {}
|
||||
|
||||
/// Overrides the internal iteration index
|
||||
CUTLASS_HOST_DEVICE
|
||||
void set_iteration_index(int index) {
|
||||
iterator_.set_iteration_index(index);
|
||||
}
|
||||
|
||||
/// Adds a pointer offset in units of Element
|
||||
CUTLASS_HOST_DEVICE
|
||||
void add_pointer_offset(LongIndex pointer_offset) {
|
||||
iterator_.add_pointer_offset(pointer_offset);
|
||||
}
|
||||
|
||||
/// Returns a pointer
|
||||
CUTLASS_HOST_DEVICE
|
||||
AccessType *get() const {
|
||||
return iterator_.get();
|
||||
}
|
||||
|
||||
/// Advances to the next tile in memory.
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileAccessIterator &operator++() {
|
||||
++iterator_;
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances to the next tile in memory.
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileAccessIterator operator++(int) {
|
||||
RegularTileAccessIterator prev(*this);
|
||||
++iterator_;
|
||||
|
||||
return prev;
|
||||
}
|
||||
|
||||
/// Adds a tile offset
|
||||
CUTLASS_DEVICE
|
||||
void add_tile_offset(TensorCoord const &coord) {
|
||||
iterator_.add_tile_offset({coord.strided(), coord.contiguous()});
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace threadblock
|
||||
} // namespace transform
|
||||
} // namespace cutlass
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* Copyright (c) 2017-2020, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without modification, are permitted
|
||||
* provided that the following conditions are met:
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* Copyright (c) 2017-2020, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without modification, are permitted
|
||||
* provided that the following conditions are met:
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* Copyright (c) 2017-2020, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without modification, are permitted
|
||||
* provided that the following conditions are met:
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* Copyright (c) 2017-2020, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without modification, are permitted
|
||||
* provided that the following conditions are met:
|
||||
@@ -831,6 +831,269 @@ class RegularTileIterator<Shape_, Element_,
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Tile iterator specialized for k interleaved arrangements for TensorOps
|
||||
///
|
||||
///
|
||||
/// Satisfies: ForwardTileIteratorConcept |
|
||||
/// ReadableContiguousTileIteratorConcept |
|
||||
/// WriteableContiguousTileIteratorConcept
|
||||
///
|
||||
template <typename Shape_, typename Element_, int AdvanceRank, typename ThreadMap_, int InterleavedK, int Alignment>
|
||||
class RegularTileIterator<
|
||||
Shape_, Element_,
|
||||
layout::TensorOpMultiplicandRowMajorInterleaved<sizeof_bits<Element_>::value,
|
||||
InterleavedK>,
|
||||
AdvanceRank, ThreadMap_, Alignment> {
|
||||
public:
|
||||
static_assert(
|
||||
AdvanceRank == 0 || AdvanceRank == 1,
|
||||
"Specialization for pitch-linear iterator may along advance along the "
|
||||
"contiguous(rank=0) or strided(rank=1) dimension.");
|
||||
|
||||
using Shape = Shape_;
|
||||
using Element = Element_;
|
||||
using Layout =
|
||||
layout::TensorOpMultiplicandRowMajorInterleaved<sizeof_bits<Element_>::value,
|
||||
InterleavedK>;
|
||||
static int const kAdvanceRank = AdvanceRank;
|
||||
static int const kAlignment = Alignment;
|
||||
|
||||
using Index = typename Layout::Index;
|
||||
using LongIndex = typename Layout::LongIndex;
|
||||
|
||||
using TensorRef = TensorRef<Element, Layout>;
|
||||
using TensorCoord = typename Layout::TensorCoord;
|
||||
|
||||
using ThreadMap = ThreadMap_;
|
||||
|
||||
/// Internal details made public to facilitate introspection
|
||||
struct Detail {
|
||||
/// This iterator is specialized for an access size that is 128 bits in
|
||||
/// length.
|
||||
static int const kAccessSizeInBits = 128;
|
||||
|
||||
static_assert(sizeof_bits<Element_>::value * ThreadMap::kElementsPerAccess ==
|
||||
kAccessSizeInBits,
|
||||
"This iterator requires a policy whose access size is 128bs");
|
||||
};
|
||||
|
||||
private:
|
||||
|
||||
/// Element type per access
|
||||
using AccessType = Array<Element, Layout::kElementsPerAccess>;
|
||||
|
||||
public:
|
||||
/// Fragment object to be loaded or stored
|
||||
using Fragment =
|
||||
Array<Element, ThreadMap::Iterations::kCount * Layout::kElementsPerAccess>;
|
||||
|
||||
/// Underlying iterator to compute the addresses
|
||||
using TileAccessIterator = RegularTileAccessIterator<Shape, Element, Layout,
|
||||
kAdvanceRank, ThreadMap>;
|
||||
|
||||
private:
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
/// Data member to the tile access iterator
|
||||
TileAccessIterator address_iterator_;
|
||||
|
||||
public:
|
||||
/// Construct a TileIterator with zero threadblock offset
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileIterator(TensorRef ref, ///< Pointer to start of tensor
|
||||
int thread_id ///< ID of each participating thread
|
||||
)
|
||||
: address_iterator_(ref, thread_id) {}
|
||||
|
||||
/// Adds a pointer offset in units of Element
|
||||
CUTLASS_HOST_DEVICE
|
||||
void add_pointer_offset(LongIndex pointer_offset) {
|
||||
address_iterator_.add_pointer_offset(pointer_offset);
|
||||
}
|
||||
|
||||
/// Advances to the next tile in memory.
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileIterator &operator++() {
|
||||
address_iterator_.add_pointer_offset(Shape::kCount);
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances to the next tile in memory.
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileIterator operator++(int) {
|
||||
RegularTileIterator prev(*this);
|
||||
this->operator++();
|
||||
|
||||
return prev;
|
||||
}
|
||||
|
||||
/// Adds a tile offset
|
||||
CUTLASS_DEVICE
|
||||
void add_tile_offset(TensorCoord const &coord) {
|
||||
address_iterator_.add_pointer_offset(coord.contiguous() * Shape::kCount);
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory
|
||||
CUTLASS_DEVICE
|
||||
void load_with_pointer_offset(Fragment &frag, Index pointer_offset) {
|
||||
address_iterator_.set_iteration_index(0);
|
||||
AccessType *frag_ptr = reinterpret_cast<AccessType *>(&frag);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int s = 0; s < ThreadMap::Iterations::kStrided; ++s) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int c = 0; c < ThreadMap::Iterations::kContiguous; ++c) {
|
||||
int access_idx = c + s * ThreadMap::Iterations::kContiguous;
|
||||
frag_ptr[access_idx] = *(address_iterator_.get() + pointer_offset);
|
||||
++address_iterator_;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory
|
||||
CUTLASS_DEVICE
|
||||
void load(Fragment &frag) { load_with_pointer_offset(frag, 0); }
|
||||
|
||||
/// Store a fragment to memory
|
||||
CUTLASS_DEVICE
|
||||
void store_with_pointer_offset(Fragment const &frag, Index pointer_offset) {
|
||||
AccessType const *frag_ptr = reinterpret_cast<AccessType const *>(&frag);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int s = 0; s < ThreadMap::Iterations::kStrided; ++s) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int c = 0; c < ThreadMap::Iterations::kContiguous; ++c) {
|
||||
int access_idx = c + s * ThreadMap::Iterations::kContiguous;
|
||||
*(address_iterator_.get() + pointer_offset) = frag_ptr[access_idx];
|
||||
++address_iterator_;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Store a fragment to memory
|
||||
CUTLASS_DEVICE
|
||||
void store(Fragment const &frag) { store_with_pointer_offset(frag, 0); }
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Tile iterator specialized for k interleaved arrangements for TensorOps
|
||||
///
|
||||
///
|
||||
/// Satisfies: ForwardTileIteratorConcept |
|
||||
/// ReadableContiguousTileIteratorConcept |
|
||||
/// WriteableContiguousTileIteratorConcept
|
||||
///
|
||||
|
||||
template <typename Shape_, typename Element_, int AdvanceRank, typename ThreadMap_, int InterleavedK, int Alignment>
|
||||
class RegularTileIterator<
|
||||
Shape_, Element_,
|
||||
layout::TensorOpMultiplicandColumnMajorInterleaved<sizeof_bits<Element_>::value,
|
||||
InterleavedK>,
|
||||
AdvanceRank, ThreadMap_, Alignment> {
|
||||
|
||||
public:
|
||||
static_assert(
|
||||
AdvanceRank == 0 || AdvanceRank == 1,
|
||||
"Specialization for pitch-linear iterator may along advance along the "
|
||||
"contiguous(rank=0) or strided(rank=1) dimension.");
|
||||
|
||||
using Shape = Shape_;
|
||||
using Element = Element_;
|
||||
using Layout =
|
||||
layout::TensorOpMultiplicandColumnMajorInterleaved<sizeof_bits<Element_>::value,
|
||||
InterleavedK>;
|
||||
static int const kAdvanceRank = AdvanceRank;
|
||||
static int const kAlignment = Alignment;
|
||||
|
||||
using Index = typename Layout::Index;
|
||||
using LongIndex = typename Layout::LongIndex;
|
||||
|
||||
using TensorRef = TensorRef<Element, Layout>;
|
||||
using TensorCoord = typename Layout::TensorCoord;
|
||||
|
||||
using ThreadMap = ThreadMap_;
|
||||
|
||||
/// Underlying iterator type
|
||||
using UnderlyingIterator = RegularTileIterator<
|
||||
cutlass::MatrixShape<Shape::kColumn, Shape::kRow>,
|
||||
Element,
|
||||
layout::TensorOpMultiplicandRowMajorInterleaved<sizeof_bits<Element_>::value, InterleavedK>,
|
||||
(kAdvanceRank == 1 ? 0 : 1),
|
||||
ThreadMap
|
||||
>;
|
||||
|
||||
public:
|
||||
/// Fragment object to be loaded or stored
|
||||
using Fragment = Array<Element, UnderlyingIterator::Fragment::kElements>;
|
||||
|
||||
private:
|
||||
|
||||
/// Underlying iterator
|
||||
UnderlyingIterator iterator_;
|
||||
|
||||
public:
|
||||
/// Construct a TileIterator with zero threadblock offset
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileIterator(TensorRef ref, ///< Pointer to start of tensor
|
||||
int thread_id ///< ID of each participating thread
|
||||
)
|
||||
: iterator_({ref.data(), ref.stride()}, thread_id) {}
|
||||
|
||||
/// Adds a pointer offset in units of Element
|
||||
CUTLASS_HOST_DEVICE
|
||||
void add_pointer_offset(LongIndex pointer_offset) {
|
||||
iterator_.add_pointer_offset(pointer_offset);
|
||||
}
|
||||
|
||||
/// Advances to the next tile in memory.
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileIterator &operator++() {
|
||||
++iterator_;
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances to the next tile in memory.
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileIterator operator++(int) {
|
||||
RegularTileIterator prev(*this);
|
||||
++iterator_;
|
||||
|
||||
return prev;
|
||||
}
|
||||
|
||||
/// Adds a tile offset
|
||||
CUTLASS_DEVICE
|
||||
void add_tile_offset(TensorCoord const &coord) {
|
||||
iterator_.add_tile_offset({coord.strided(), coord.contiguous()});
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory
|
||||
CUTLASS_DEVICE
|
||||
void load_with_pointer_offset(Fragment &frag, Index pointer_offset) {
|
||||
iterator_.load_with_pointer_offset(frag, pointer_offset);
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory
|
||||
CUTLASS_DEVICE
|
||||
void load(Fragment &frag) { load_with_pointer_offset(frag, 0); }
|
||||
|
||||
/// Store a fragment to memory
|
||||
CUTLASS_DEVICE
|
||||
void store_with_pointer_offset(Fragment const &frag, Index pointer_offset) {
|
||||
iterator_.store_with_pointer_offset(frag, pointer_offset);
|
||||
}
|
||||
|
||||
/// Store a fragment to memory
|
||||
CUTLASS_DEVICE
|
||||
void store(Fragment const &frag) { store_with_pointer_offset(frag, 0); }
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace threadblock
|
||||
} // namespace transform
|
||||
} // namespace cutlass
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* Copyright (c) 2017-2020, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without modification, are permitted
|
||||
* provided that the following conditions are met:
|
||||
|
||||
Reference in New Issue
Block a user