MCPcopy Create free account
hub / github.com/KnowingNothing/MatmulTutorial / OrderedSequenceBarrier

Class OrderedSequenceBarrier

include/barrier.h:228–305  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

226
227template <int SequenceDepth, int SequenceLength>
228class 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 &params)
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) {

Callers

nothing calls this directly

Calls 1

indexMethod · 0.80

Tested by

no test coverage detected