Skip to content

Commit 1568c4a

Browse files
committed
#2281: RecursiveDoubling allreduce - add Collection support
1 parent 90c5668 commit 1568c4a

7 files changed

Lines changed: 178 additions & 69 deletions

File tree

src/vt/collective/reduce/allreduce/recursive_doubling.cc

Lines changed: 36 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -47,6 +47,42 @@
4747

4848
namespace vt::collective::reduce::allreduce {
4949

50+
RecursiveDoubling::RecursiveDoubling(
51+
detail::StrongVrtProxy proxy, detail::StrongGroup group, size_t num_elems)
52+
: collection_proxy_(proxy.get()),
53+
local_num_elems_(num_elems),
54+
nodes_(theGroup()->GetGroupNodes(group.get())),
55+
num_nodes_(nodes_.size()),
56+
this_node_(theContext()->getNode()),
57+
num_steps_(static_cast<uint32_t>(std::log2(num_nodes_))),
58+
nprocs_pof2_(1 << num_steps_),
59+
nprocs_rem_(num_nodes_ - nprocs_pof2_) {
60+
auto const is_default_group = theGroup()->isGroupDefault(group.get());
61+
if (not is_default_group) {
62+
auto it = std::find(nodes_.begin(), nodes_.end(), theContext()->getNode());
63+
vtAssert(it != nodes_.end(), "This node was not found in group nodes!");
64+
65+
this_node_ = it - nodes_.begin();
66+
}
67+
68+
is_even_ = this_node_ % 2 == 0;
69+
is_part_of_adjustment_group_ = this_node_ < (2 * nprocs_rem_);
70+
if (is_part_of_adjustment_group_) {
71+
if (is_even_) {
72+
vrt_node_ = this_node_ / 2;
73+
} else {
74+
vrt_node_ = -1;
75+
}
76+
} else {
77+
vrt_node_ = this_node_ - nprocs_rem_;
78+
}
79+
80+
vt_debug_print(
81+
terse, allreduce,
82+
"RecursiveDoubling (this={}): proxy={:x} proxy_={} local_num_elems={}\n",
83+
print_ptr(this), proxy.get(), proxy_.getProxy(), local_num_elems_);
84+
}
85+
5086
RecursiveDoubling::RecursiveDoubling(
5187
detail::StrongObjGroup objgroup)
5288
: objgroup_proxy_(objgroup.get()),

src/vt/collective/reduce/allreduce/recursive_doubling.h

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -76,6 +76,7 @@ namespace vt::collective::reduce::allreduce {
7676
*/
7777

7878
struct RecursiveDoubling {
79+
RecursiveDoubling(detail::StrongVrtProxy proxy, detail::StrongGroup group, size_t num_elems);
7980
/**
8081
* \brief Constructor for RecursiveDoubling class.
8182
*

src/vt/vrt/collection/manager.h

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -744,7 +744,7 @@ struct CollectionManager
744744
);
745745

746746

747-
template <auto f, typename ColT, template <typename Arg> class Op, typename ...Args>
747+
template <typename ReducerT, auto f, typename ColT, template <typename Arg> class Op, typename ...Args>
748748
messaging::PendingSend reduceLocal(
749749
CollectionProxyWrapType<ColT> const& proxy, Args &&... args
750750
);
@@ -1796,6 +1796,7 @@ struct CollectionManager
17961796

17971797
// Allreduce stuff, probably should be moved elsewhere
17981798
std::unordered_map<VirtualProxyType, ObjGroupProxyType> rabenseifner_reducers_;
1799+
std::unordered_map<VirtualProxyType, ObjGroupProxyType> recursive_doubling_reducers_;
17991800
};
18001801

18011802
}}} /* end namespace vt::vrt::collection */

src/vt/vrt/collection/manager.impl.h

Lines changed: 71 additions & 30 deletions
Original file line numberDiff line numberDiff line change
@@ -41,6 +41,9 @@
4141
//@HEADER
4242
*/
4343

44+
#include "vt/collective/reduce/allreduce/recursive_doubling.h"
45+
#include "vt/collective/reduce/allreduce/type.h"
46+
#include <type_traits>
4447
#if !defined INCLUDED_VT_VRT_COLLECTION_MANAGER_IMPL_H
4548
#define INCLUDED_VT_VRT_COLLECTION_MANAGER_IMPL_H
4649

@@ -894,7 +897,7 @@ messaging::PendingSend CollectionManager::broadcastMsgUntypedHandler(
894897
}
895898

896899
template <
897-
auto f, typename ColT, template <typename Arg> class Op, typename... Args>
900+
typename ReducerT, auto f, typename ColT, template <typename Arg> class Op, typename... Args>
898901
messaging::PendingSend CollectionManager::reduceLocal(
899902
CollectionProxyWrapType<ColT> const& proxy, Args&&... args) {
900903
using namespace collective::reduce::allreduce;
@@ -913,44 +916,82 @@ messaging::PendingSend CollectionManager::reduceLocal(
913916
auto const group = elm_holder->group();
914917
bool const use_group = group_ready && send_group;
915918

916-
using Reducer = collective::reduce::allreduce::Rabenseifner;
917919
auto stamp = proxy(idx).tryGetLocalPtr()->getNextAllreduceStamp();
918920
auto const id = std::get<collective::reduce::detail::StrongSeq>(stamp).get();
919921

920922
auto cb = vt::theCB()->makeCallbackBcastCollectiveProxy<f>(proxy);
921923

922-
// Incorrect! will yield same reducer for different Op/payload size/final handler etc.
923-
if (auto reducer = rabenseifner_reducers_.find(col_proxy);
924-
reducer == rabenseifner_reducers_.end()) {
925-
if (use_group) {
926-
// theGroup()->allreduce<f, Op>(group, );
924+
if constexpr (std::is_same_v<ReducerT, RabenseifnerT>) {
925+
using Reducer = collective::reduce::allreduce::Rabenseifner;
926+
if (auto reducer = rabenseifner_reducers_.find(col_proxy);
927+
reducer == rabenseifner_reducers_.end()) {
928+
if (use_group) {
929+
// theGroup()->allreduce<f, Op>(group, );
930+
} else {
931+
vt_debug_print(
932+
terse, allreduce, "Creating Reducer on idx={} with id={}\n", idx, id);
933+
auto obj_proxy = theObjGroup()->makeCollective<Reducer>(
934+
"reducer", collective::reduce::detail::StrongVrtProxy{col_proxy},
935+
collective::reduce::detail::StrongGroup{group}, num_elms);
936+
937+
rabenseifner_reducers_[col_proxy] = obj_proxy.getProxy();
938+
auto* obj = obj_proxy[theContext()->getNode()].get();
939+
obj->proxy_ = obj_proxy;
940+
941+
obj->template setFinalHandler<DataT>(cb, id);
942+
obj->template localReduce<DataT, Op>(id, std::forward<Args>(args)...);
943+
}
927944
} else {
928-
929-
vt_debug_print(terse, allreduce, "Creating Reducer on idx={} with id={}\n", idx, id);
930-
auto obj_proxy = theObjGroup()->makeCollective<Reducer>(
931-
"reducer", collective::reduce::detail::StrongVrtProxy{col_proxy},
932-
collective::reduce::detail::StrongGroup{group}, num_elms
933-
);
934-
935-
rabenseifner_reducers_[col_proxy] = obj_proxy.getProxy();
936-
auto* obj = obj_proxy[theContext()->getNode()].get();
937-
obj->proxy_ = obj_proxy;
938-
939-
obj->template setFinalHandler<DataT>(cb, id);
940-
obj->template localReduce<DataT, Op>(id, std::forward<Args>(args)...);
945+
if (use_group) {
946+
// theGroup()->allreduce<f, Op>(group, );
947+
} else {
948+
vt_debug_print(
949+
terse, allreduce, "Reusing Reducer on idx={} with id={}\n", idx, id);
950+
auto obj_proxy =
951+
reducer->second; // rabenseifner_reducers_.at(col_proxy);
952+
auto typed_proxy =
953+
static_cast<vt::objgroup::proxy::Proxy<Reducer>>(obj_proxy);
954+
auto* obj = typed_proxy[theContext()->getNode()].get();
955+
956+
obj->template setFinalHandler<DataT>(cb, id);
957+
obj->template localReduce<DataT, Op>(id, std::forward<Args>(args)...);
958+
}
941959
}
942960
} else {
943-
if (use_group) {
944-
// theGroup()->allreduce<f, Op>(group, );
961+
using Reducer = collective::reduce::allreduce::RecursiveDoubling;
962+
if (auto reducer = recursive_doubling_reducers_.find(col_proxy);
963+
reducer == recursive_doubling_reducers_.end()) {
964+
if (use_group) {
965+
// theGroup()->allreduce<f, Op>(group, );
966+
} else {
967+
vt_debug_print(
968+
terse, allreduce, "Creating Reducer on idx={} with id={}\n", idx, id);
969+
auto obj_proxy = theObjGroup()->makeCollective<Reducer>(
970+
"reducer", collective::reduce::detail::StrongVrtProxy{col_proxy},
971+
collective::reduce::detail::StrongGroup{group}, num_elms);
972+
973+
recursive_doubling_reducers_[col_proxy] = obj_proxy.getProxy();
974+
auto* obj = obj_proxy[theContext()->getNode()].get();
975+
obj->proxy_ = obj_proxy;
976+
977+
obj->template setFinalHandler<DataT>(cb, id);
978+
obj->template localReduce<DataT, Op>(id, std::forward<Args>(args)...);
979+
}
945980
} else {
946-
vt_debug_print(terse, allreduce, "Reusing Reducer on idx={} with id={}\n", idx, id);
947-
auto obj_proxy = reducer->second; // rabenseifner_reducers_.at(col_proxy);
948-
auto typed_proxy =
949-
static_cast<vt::objgroup::proxy::Proxy<Reducer>>(obj_proxy);
950-
auto* obj = typed_proxy[theContext()->getNode()].get();
951-
952-
obj->template setFinalHandler<DataT>(cb, id);
953-
obj->template localReduce<DataT, Op>(id, std::forward<Args>(args)...);
981+
if (use_group) {
982+
// theGroup()->allreduce<f, Op>(group, );
983+
} else {
984+
vt_debug_print(
985+
terse, allreduce, "Reusing Reducer on idx={} with id={}\n", idx, id);
986+
auto obj_proxy =
987+
reducer->second; // rabenseifner_reducers_.at(col_proxy);
988+
auto typed_proxy =
989+
static_cast<vt::objgroup::proxy::Proxy<Reducer>>(obj_proxy);
990+
auto* obj = typed_proxy[theContext()->getNode()].get();
991+
992+
obj->template setFinalHandler<DataT>(cb, id);
993+
obj->template localReduce<DataT, Op>(id, std::forward<Args>(args)...);
994+
}
954995
}
955996
}
956997

src/vt/vrt/collection/reducable/reducable.h

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -86,11 +86,12 @@ struct Reducable : BaseProxyT {
8686
) const;
8787

8888
template <
89+
typename ReducerT,
8990
auto f,
9091
template <typename Arg> class Op = collective::NoneOp,
9192
typename... Args
9293
>
93-
messaging::PendingSend allreduce_h(
94+
messaging::PendingSend allreduce(
9495
Args&&... args
9596
) const;
9697

src/vt/vrt/collection/reducable/reducable.impl.h

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -81,12 +81,12 @@ messaging::PendingSend Reducable<ColT,IndexT,BaseProxyT>::allreduce(
8181
}
8282

8383
template <typename ColT, typename IndexT, typename BaseProxyT>
84-
template <auto f, template <typename Arg> class Op, typename... Args>
85-
messaging::PendingSend Reducable<ColT,IndexT,BaseProxyT>::allreduce_h(
84+
template <typename ReducerT, auto f, template <typename Arg> class Op, typename... Args>
85+
messaging::PendingSend Reducable<ColT,IndexT,BaseProxyT>::allreduce(
8686
Args&&... args
8787
) const {
8888
auto const proxy = this->getProxy();
89-
return theCollection()->reduceLocal<f, ColT, Op>(
89+
return theCollection()->reduceLocal<ReducerT, f, ColT, Op>(
9090
proxy, std::forward<Args>(args)...);
9191
}
9292

tests/perf/allreduce.cc

Lines changed: 63 additions & 34 deletions
Original file line numberDiff line numberDiff line change
@@ -291,60 +291,92 @@ VT_PERF_TEST(MyTest, test_allreduce_group_rabenseifner) {
291291
}
292292
}
293293

294-
struct Hello : vt::Collection<Hello, vt::Index1D> {
295-
Hello() {
294+
struct RabensifnerColl : vt::Collection<RabensifnerColl, vt::Index1D> {
295+
RabensifnerColl() {
296296
for (auto const payload_size : payloadSizes) {
297-
timer_names_[payload_size] = fmt::format("Collection {}", payload_size);
297+
timer_names_[payload_size] = fmt::format("Collection Rabenseifner {}", payload_size);
298298
}
299299
}
300300

301-
void finalMaxHan(std::vector<int32_t> result) {
302-
std::string result_s = "";
303-
for (auto val : result) {
304-
result_s.append(fmt::format("{} ", val));
305-
}
306-
fmt::print(
307-
"[{}]: Allreduce finalMaxHan (Values=[{}]), idx={}\n",
308-
theContext()->getNode(), result_s, getIndex().x());
301+
void allreduceHan(std::vector<int32_t> result) {
302+
col_send_done_ = true;
303+
parent_->StopTimer(timer_names_.at(result.size()));
304+
}
305+
306+
void executeAllreduce(size_t payload_size) {
307+
auto proxy = this->getCollectionProxy();
308+
309+
std::vector<int32_t> payload(payload_size, getIndex().x());
310+
parent_->StartTimer(timer_names_.at(payload_size));
311+
proxy.allreduce<
312+
collective::reduce::allreduce::RabenseifnerT, &RabensifnerColl::allreduceHan,
313+
collective::PlusOp
314+
>(payload);
315+
}
316+
317+
bool col_send_done_ = false;
318+
std::unordered_map<size_t, std::string> timer_names_ = {};
319+
MyTest* parent_ = {};
320+
};
321+
322+
VT_PERF_TEST(MyTest, test_allreduce_collection_rabenseifner) {
323+
auto const num_elms_per_node = 1;
324+
auto range = vt::Index1D(int32_t{num_nodes_ * num_elms_per_node});
325+
auto proxy = vt::makeCollection<RabensifnerColl>("test_collection_allreduce")
326+
.bounds(range)
327+
.bulkInsert()
328+
.wait();
309329

310-
// col_send_done_ = true;
311-
// parent_->StopTimer(timer_names_.at(result.size()));
330+
auto const thisNode = vt::theContext()->getNode();
331+
auto const nextNode = (thisNode + 1) % num_nodes_;
332+
333+
theCollective()->barrier();
334+
335+
auto const elm = thisNode * num_elms_per_node;
336+
proxy[elm].tryGetLocalPtr()->parent_ = this;
337+
338+
for (auto payload_size : payloadSizes) {
339+
proxy.broadcastCollective<&RabensifnerColl::executeAllreduce>(payload_size);
340+
341+
// We run 1 coll elem per node, so it should be ok
342+
theSched()->runSchedulerWhile(
343+
[&] { return !proxy[elm].tryGetLocalPtr()->col_send_done_; });
344+
proxy[elm].tryGetLocalPtr()->col_send_done_ = false;
312345
}
346+
}
313347

314-
void finalHan(std::vector<int32_t> result) {
315-
// std::string result_s = "";
316-
// for(auto val : result){
317-
// result_s.append(fmt::format("{} ", val));
318-
// }
319-
// fmt::print(
320-
// "[{}]: Allreduce handler (Values=[{}]), idx={}\n",
321-
// theContext()->getNode(), result_s, getIndex().x()
322-
// );
348+
struct RecursiveDoublingColl : vt::Collection<RecursiveDoublingColl, vt::Index1D> {
349+
RecursiveDoublingColl() {
350+
for (auto const payload_size : payloadSizes) {
351+
timer_names_[payload_size] = fmt::format("Collection RecursiveDoubling {}", payload_size);
352+
}
353+
}
323354

355+
void allreduceHan(std::vector<int32_t> result) {
324356
col_send_done_ = true;
325357
parent_->StopTimer(timer_names_.at(result.size()));
326358
}
327359

328-
void handler(size_t payload_size) {
360+
void executeAllreduce(size_t payload_size) {
329361
auto proxy = this->getCollectionProxy();
330362

331363
std::vector<int32_t> payload(payload_size, getIndex().x());
332364
parent_->StartTimer(timer_names_.at(payload_size));
333-
proxy.allreduce_h<&Hello::finalHan, collective::PlusOp>(payload);
334-
335-
// proxy.allreduce_h<&Hello::finalMaxHan, collective::MaxOp>(
336-
// std::move(payload));
365+
proxy.allreduce<
366+
collective::reduce::allreduce::RecursiveDoublingT, &RecursiveDoublingColl::allreduceHan,
367+
collective::PlusOp
368+
>(payload);
337369
}
338370

339371
bool col_send_done_ = false;
340372
std::unordered_map<size_t, std::string> timer_names_ = {};
341373
MyTest* parent_ = {};
342374
};
343375

344-
VT_PERF_TEST(MyTest, test_allreduce_collection_rabenseifner) {
376+
VT_PERF_TEST(MyTest, test_allreduce_collection_racursive_doubling) {
345377
auto const num_elms_per_node = 1;
346378
auto range = vt::Index1D(int32_t{num_nodes_ * num_elms_per_node});
347-
auto proxy = vt::makeCollection<Hello>("test_collection_send")
379+
auto proxy = vt::makeCollection<RecursiveDoublingColl>("test_collection_allreduce")
348380
.bounds(range)
349381
.bulkInsert()
350382
.wait();
@@ -355,13 +387,10 @@ VT_PERF_TEST(MyTest, test_allreduce_collection_rabenseifner) {
355387
theCollective()->barrier();
356388

357389
auto const elm = thisNode * num_elms_per_node;
358-
359390
proxy[elm].tryGetLocalPtr()->parent_ = this;
360-
proxy.broadcastCollective<&Hello::handler>(payloadSizes.front());
361-
theSched()->runSchedulerWhile(
362-
[&] { return !proxy[elm].tryGetLocalPtr()->col_send_done_; });
391+
363392
for (auto payload_size : payloadSizes) {
364-
proxy.broadcastCollective<&Hello::handler>(payload_size);
393+
proxy.broadcastCollective<&RecursiveDoublingColl::executeAllreduce>(payload_size);
365394

366395
// We run 1 coll elem per node, so it should be ok
367396
theSched()->runSchedulerWhile(

0 commit comments

Comments
 (0)