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
896899template <
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>
898901messaging::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
0 commit comments