| 226 | |
| 227 | template <int SequenceDepth, int SequenceLength> |
| 228 | class OrderedSequenceBarrier { |
| 229 | public: |
| 230 | using Barrier = ClusterBarrier; |
| 231 | |
| 232 | struct SharedStorage { |
| 233 | Barrier barrier_[SequenceDepth][SequenceLength]; |
| 234 | }; |
| 235 | |
| 236 | struct Params { |
| 237 | uint32_t group_id; |
| 238 | uint32_t group_size; |
| 239 | int active_warps = 0; |
| 240 | }; |
| 241 | |
| 242 | private: |
| 243 | // In future this Params object can be replaced easily with a CG object |
| 244 | Params params_; |
| 245 | Barrier *barrier_ptr_; |
| 246 | PipelineState<SequenceDepth> stage_; |
| 247 | |
| 248 | static constexpr int Depth = SequenceDepth; |
| 249 | static constexpr int Length = SequenceLength; |
| 250 | |
| 251 | public: |
| 252 | OrderedSequenceBarrier() = delete; |
| 253 | OrderedSequenceBarrier(const OrderedSequenceBarrier &) = delete; |
| 254 | OrderedSequenceBarrier(OrderedSequenceBarrier &&) = delete; |
| 255 | OrderedSequenceBarrier &operator=(const OrderedSequenceBarrier &) = delete; |
| 256 | OrderedSequenceBarrier &operator=(OrderedSequenceBarrier &&) = delete; |
| 257 | ~OrderedSequenceBarrier() = default; |
| 258 | |
| 259 | DEVICE OrderedSequenceBarrier(SharedStorage &storage, Params const ¶ms) |
| 260 | : params_(params), |
| 261 | barrier_ptr_(&storage.barrier_[0][0]), |
| 262 | // Group 0 - starts with an opposite phase |
| 263 | stage_({0, (params.group_id == 0), 0}) { |
| 264 | int warp_idx = threadIdx.x / WARP_SIZE; |
| 265 | int lane_predicate = elect_one_sync(); |
| 266 | |
| 267 | // Barrier FULL, EMPTY init |
| 268 | // Init is done only by the one elected thread of the block |
| 269 | if (warp_idx == params.active_warps && lane_predicate == 1) { |
| 270 | for (int d = 0; d < Depth; ++d) { |
| 271 | for (int l = 0; l < Length; ++l) { |
| 272 | barrier_ptr_[d * Length + l].init(params.group_size); |
| 273 | } |
| 274 | } |
| 275 | } |
| 276 | fence_barrier_init(); |
| 277 | } |
| 278 | |
| 279 | // Wait on a stage to be unlocked |
| 280 | DEVICE void wait() { |
| 281 | get_barrier_for_current_stage(params_.group_id).wait(stage_.phase()); |
| 282 | } |
| 283 | |
| 284 | DEVICE void check_phase(int val) { |
| 285 | if (threadIdx.x % WARP_GROUP_SIZE == 0) { |