Skip to content

Commit 29af994

Browse files
committed
Pass the Execution Space as TT template parameter
Defaults to Host execution so existing code is not affected. Properly set by make_tt. We cannot query TT::derivedT for flags because at the time TT is instantiated because derivedT is incomplete at that point. For now pass the Space as template parameter. We need to find a different way if we want to have multiple implementations of a task. Signed-off-by: Joseph Schuchart <joseph.schuchart@stonybrook.edu>
1 parent b180cac commit 29af994

7 files changed

Lines changed: 44 additions & 60 deletions

File tree

examples/spmm/spmm.cc

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -575,7 +575,7 @@ class SpMM25D {
575575
/// 3-D process grid only
576576
template<ttg::ExecutionSpace Space_>
577577
class MultiplyAdd : public TT<Key<3>, std::tuple<Out<Key<2>, Blk>, Out<Key<3>, Blk>>, MultiplyAdd<Space_>,
578-
ttg::typelist<const Blk, const Blk, Blk>> {
578+
ttg::typelist<const Blk, const Blk, Blk>, Space_> {
579579
static constexpr const bool is_device_space = (Space_ != ttg::ExecutionSpace::Host);
580580
using task_return_type = std::conditional_t<is_device_space, ttg::device::Task, void>;
581581

ttg/ttg/madness/fwd.h

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -8,7 +8,9 @@
88

99
namespace ttg_madness {
1010

11-
template <typename keyT, typename output_terminalsT, typename derivedT, typename input_valueTs = ttg::typelist<>>
11+
template <typename keyT, typename output_terminalsT, typename derivedT,
12+
typename input_valueTs = ttg::typelist<>,
13+
ttg::ExecutionSpace Space = ttg::ExecutionSpace::Host>
1214
class TT;
1315

1416
/// \internal the OG name

ttg/ttg/madness/ttg.h

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -184,8 +184,9 @@ namespace ttg_madness {
184184
/// values
185185
/// flowing into this TT; a const type indicates nonmutating (read-only) use, nonconst type
186186
/// indicates mutating use (e.g. the corresponding input can be used as scratch, moved-from, etc.)
187-
template <typename keyT, typename output_terminalsT, typename derivedT, typename input_valueTs>
188-
class TT : public ttg::TTBase, public ::madness::WorldObject<TT<keyT, output_terminalsT, derivedT, input_valueTs>> {
187+
template <typename keyT, typename output_terminalsT, typename derivedT, typename input_valueTs, ttg::ExecutionSpace Space>
188+
class TT : public ttg::TTBase, public ::madness::WorldObject<TT<keyT, output_terminalsT, derivedT, input_valueTs, Space>> {
189+
static_assert(Space == ttg::ExecutionSpace::Host, "MADNESS backend only supports Host Execution Space");
189190
static_assert(ttg::meta::is_typelist_v<input_valueTs>,
190191
"The fourth template for ttg::TT must be a ttg::typelist containing the input types");
191192
using input_tuple_type = ttg::meta::typelist_to_tuple_t<input_valueTs>;

ttg/ttg/make_tt.h

Lines changed: 1 addition & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -17,7 +17,7 @@ class CallableWrapTT
1717
: public TT<
1818
keyT, output_terminalsT,
1919
CallableWrapTT<funcT, returnT, funcT_receives_input_tuple, funcT_receives_outterm_tuple, space, keyT, output_terminalsT, input_valuesT...>,
20-
ttg::typelist<input_valuesT...>> {
20+
ttg::typelist<input_valuesT...>, space> {
2121
using baseT = typename CallableWrapTT::ttT;
2222

2323
using input_values_tuple_type = typename baseT::input_values_tuple_type;
@@ -44,11 +44,6 @@ class CallableWrapTT
4444
void;
4545
#endif // TTG_HAVE_COROUTINE
4646

47-
public:
48-
static constexpr bool have_cuda_op = (space == ttg::ExecutionSpace::CUDA);
49-
static constexpr bool have_hip_op = (space == ttg::ExecutionSpace::HIP);
50-
static constexpr bool have_level_zero_op = (space == ttg::ExecutionSpace::L0);
51-
5247
protected:
5348

5449
template<typename ReturnT>

ttg/ttg/parsec/fwd.h

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -10,7 +10,9 @@ extern "C" struct parsec_context_s;
1010

1111
namespace ttg_parsec {
1212

13-
template <typename keyT, typename output_terminalsT, typename derivedT, typename input_valueTs = ttg::typelist<>>
13+
template <typename keyT, typename output_terminalsT, typename derivedT,
14+
typename input_valueTs = ttg::typelist<>,
15+
ttg::ExecutionSpace Space = ttg::ExecutionSpace::Host>
1416
class TT;
1517

1618
/// \internal the OG name

ttg/ttg/parsec/task.h

Lines changed: 6 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -252,9 +252,9 @@ namespace ttg_parsec {
252252
template<ttg::ExecutionSpace Space>
253253
parsec_hook_return_t invoke_op() {
254254
if constexpr (Space == ttg::ExecutionSpace::Host) {
255-
return TT::template static_op<Space>(&this->parsec_task);
255+
return TT::static_op(&this->parsec_task);
256256
} else {
257-
return TT::template device_static_op<Space>(&this->parsec_task);
257+
return TT::device_static_op(&this->parsec_task);
258258
}
259259
}
260260

@@ -263,7 +263,7 @@ namespace ttg_parsec {
263263
if constexpr (Space == ttg::ExecutionSpace::Host) {
264264
return PARSEC_HOOK_RETURN_DONE;
265265
} else {
266-
return TT::template device_static_evaluate<Space>(&this->parsec_task);
266+
return TT::device_static_evaluate(&this->parsec_task);
267267
}
268268
}
269269

@@ -310,9 +310,9 @@ namespace ttg_parsec {
310310
template<ttg::ExecutionSpace Space>
311311
parsec_hook_return_t invoke_op() {
312312
if constexpr (Space == ttg::ExecutionSpace::Host) {
313-
return TT::template static_op<Space>(&this->parsec_task);
313+
return TT::static_op(&this->parsec_task);
314314
} else {
315-
return TT::template device_static_op<Space>(&this->parsec_task);
315+
return TT::device_static_op(&this->parsec_task);
316316
}
317317
}
318318

@@ -321,7 +321,7 @@ namespace ttg_parsec {
321321
if constexpr (Space == ttg::ExecutionSpace::Host) {
322322
return PARSEC_HOOK_RETURN_DONE;
323323
} else {
324-
return TT::template device_static_evaluate<Space>(&this->parsec_task);
324+
return TT::device_static_evaluate(&this->parsec_task);
325325
}
326326
}
327327

ttg/ttg/parsec/ttg.h

Lines changed: 27 additions & 43 deletions
Original file line numberDiff line numberDiff line change
@@ -514,8 +514,9 @@ namespace ttg_parsec {
514514
#endif // TTG_USE_USER_TERMDET
515515
}
516516

517-
template <typename keyT, typename output_terminalsT, typename derivedT, typename input_valueTs = ttg::typelist<>>
518-
void register_tt_profiling(const TT<keyT, output_terminalsT, derivedT, input_valueTs> *t) {
517+
template <typename keyT, typename output_terminalsT, typename derivedT,
518+
typename input_valueTs = ttg::typelist<>, ttg::ExecutionSpace Space>
519+
void register_tt_profiling(const TT<keyT, output_terminalsT, derivedT, input_valueTs, Space> *t) {
519520
#if defined(PARSEC_PROF_TRACE)
520521
std::stringstream ss;
521522
build_composite_name_rec(t->ttg_ptr(), ss);
@@ -1180,7 +1181,7 @@ namespace ttg_parsec {
11801181

11811182
} // namespace detail
11821183

1183-
template <typename keyT, typename output_terminalsT, typename derivedT, typename input_valueTs>
1184+
template <typename keyT, typename output_terminalsT, typename derivedT, typename input_valueTs, ttg::ExecutionSpace Space>
11841185
class TT : public ttg::TTBase, detail::ParsecTTBase {
11851186
private:
11861187
/// preconditions
@@ -1217,29 +1218,17 @@ namespace ttg_parsec {
12171218
public:
12181219
/// @return true if derivedT::have_cuda_op exists and is defined to true
12191220
static constexpr bool derived_has_cuda_op() {
1220-
if constexpr (ttg::meta::is_detected_v<have_cuda_op_non_type_t, derivedT>) {
1221-
return derivedT::have_cuda_op;
1222-
} else {
1223-
return false;
1224-
}
1221+
return Space == ttg::ExecutionSpace::CUDA;
12251222
}
12261223

12271224
/// @return true if derivedT::have_hip_op exists and is defined to true
12281225
static constexpr bool derived_has_hip_op() {
1229-
if constexpr (ttg::meta::is_detected_v<have_hip_op_non_type_t, derivedT>) {
1230-
return derivedT::have_hip_op;
1231-
} else {
1232-
return false;
1233-
}
1226+
return Space == ttg::ExecutionSpace::HIP;
12341227
}
12351228

12361229
/// @return true if derivedT::have_hip_op exists and is defined to true
12371230
static constexpr bool derived_has_level_zero_op() {
1238-
if constexpr (ttg::meta::is_detected_v<have_level_zero_op_non_type_t, derivedT>) {
1239-
return derivedT::have_level_zero_op;
1240-
} else {
1241-
return false;
1242-
}
1231+
return Space == ttg::ExecutionSpace::L0;
12431232
}
12441233

12451234
/// @return true if the TT supports device execution
@@ -1354,18 +1343,17 @@ namespace ttg_parsec {
13541343
/// dispatches a call to derivedT::op
13551344
/// @return void if called a synchronous function, or ttg::coroutine_handle<> if called a coroutine (if non-null,
13561345
/// points to the suspended coroutine)
1357-
template <ttg::ExecutionSpace Space, typename... Args>
1346+
template <typename... Args>
13581347
auto op(Args &&...args) {
13591348
derivedT *derived = static_cast<derivedT *>(this);
1360-
//if constexpr (Space == ttg::ExecutionSpace::Host) {
1361-
using return_type = decltype(derived->op(std::forward<Args>(args)...));
1362-
if constexpr (std::is_same_v<return_type,void>) {
1363-
derived->op(std::forward<Args>(args)...);
1364-
return;
1365-
}
1366-
else {
1367-
return derived->op(std::forward<Args>(args)...);
1368-
}
1349+
using return_type = decltype(derived->op(std::forward<Args>(args)...));
1350+
if constexpr (std::is_same_v<return_type,void>) {
1351+
derived->op(std::forward<Args>(args)...);
1352+
return;
1353+
}
1354+
else {
1355+
return derived->op(std::forward<Args>(args)...);
1356+
}
13691357
}
13701358

13711359
template <std::size_t i, typename terminalT, typename Key>
@@ -1418,7 +1406,6 @@ namespace ttg_parsec {
14181406
/**
14191407
* Submit callback called by PaRSEC once all input transfers have completed.
14201408
*/
1421-
template <ttg::ExecutionSpace Space>
14221409
static int device_static_submit(parsec_device_gpu_module_t *gpu_device,
14231410
parsec_gpu_task_t *gpu_task,
14241411
parsec_gpu_exec_stream_t *gpu_stream) {
@@ -1464,7 +1451,7 @@ namespace ttg_parsec {
14641451
#endif // defined(PARSEC_HAVE_DEV_CUDA_SUPPORT) && defined(TTG_HAVE_CUDA)
14651452

14661453
/* Here we call back into the coroutine again after the transfers have completed */
1467-
static_op<Space>(&task->parsec_task);
1454+
static_op(&task->parsec_task);
14681455

14691456
ttg::device::detail::reset_current();
14701457

@@ -1494,7 +1481,6 @@ namespace ttg_parsec {
14941481
return rc;
14951482
}
14961483

1497-
template <ttg::ExecutionSpace Space>
14981484
static parsec_hook_return_t device_static_evaluate(parsec_task_t* parsec_task) {
14991485

15001486
task_t *task = (task_t*)parsec_task;
@@ -1509,7 +1495,7 @@ namespace ttg_parsec {
15091495
gpu_task->task_type = 0; // user task
15101496
gpu_task->last_data_check_epoch = 0; // used internally
15111497
gpu_task->pushout = 0;
1512-
gpu_task->submit = &TT::device_static_submit<Space>;
1498+
gpu_task->submit = &TT::device_static_submit;
15131499

15141500
// one way to force the task device
15151501
// currently this will probably break all of PaRSEC if this hint
@@ -1527,7 +1513,7 @@ namespace ttg_parsec {
15271513
task->dev_ptr->task_class = *task->parsec_task.task_class;
15281514

15291515
// first invocation of the coroutine to get the coroutine handle
1530-
static_op<Space>(parsec_task);
1516+
static_op(parsec_task);
15311517

15321518
/* when we come back here, the flows in gpu_task are set (see register_device_memory) */
15331519

@@ -1577,7 +1563,6 @@ namespace ttg_parsec {
15771563

15781564
}
15791565

1580-
template <ttg::ExecutionSpace Space>
15811566
static parsec_hook_return_t device_static_op(parsec_task_t* parsec_task) {
15821567
static_assert(derived_has_device_op());
15831568

@@ -1649,7 +1634,6 @@ namespace ttg_parsec {
16491634
}
16501635
#endif // TTG_HAVE_DEVICE
16511636

1652-
template <ttg::ExecutionSpace Space>
16531637
static parsec_hook_return_t static_op(parsec_task_t *parsec_task) {
16541638

16551639
task_t *task = (task_t*)parsec_task;
@@ -1675,14 +1659,14 @@ namespace ttg_parsec {
16751659

16761660
if constexpr (!ttg::meta::is_void_v<keyT> && !ttg::meta::is_empty_tuple_v<input_values_tuple_type>) {
16771661
auto input = make_tuple_of_ref_from_array(task, std::make_index_sequence<numinvals>{});
1678-
TTG_PROCESS_TT_OP_RETURN(suspended_task_address, task->coroutine_id, baseobj->template op<Space>(task->key, std::move(input), obj->output_terminals));
1662+
TTG_PROCESS_TT_OP_RETURN(suspended_task_address, task->coroutine_id, baseobj->op(task->key, std::move(input), obj->output_terminals));
16791663
} else if constexpr (!ttg::meta::is_void_v<keyT> && ttg::meta::is_empty_tuple_v<input_values_tuple_type>) {
1680-
TTG_PROCESS_TT_OP_RETURN(suspended_task_address, task->coroutine_id, baseobj->template op<Space>(task->key, obj->output_terminals));
1664+
TTG_PROCESS_TT_OP_RETURN(suspended_task_address, task->coroutine_id, baseobj->op(task->key, obj->output_terminals));
16811665
} else if constexpr (ttg::meta::is_void_v<keyT> && !ttg::meta::is_empty_tuple_v<input_values_tuple_type>) {
16821666
auto input = make_tuple_of_ref_from_array(task, std::make_index_sequence<numinvals>{});
1683-
TTG_PROCESS_TT_OP_RETURN(suspended_task_address, task->coroutine_id, baseobj->template op<Space>(std::move(input), obj->output_terminals));
1667+
TTG_PROCESS_TT_OP_RETURN(suspended_task_address, task->coroutine_id, baseobj->op(std::move(input), obj->output_terminals));
16841668
} else if constexpr (ttg::meta::is_void_v<keyT> && ttg::meta::is_empty_tuple_v<input_values_tuple_type>) {
1685-
TTG_PROCESS_TT_OP_RETURN(suspended_task_address, task->coroutine_id, baseobj->template op<Space>(obj->output_terminals));
1669+
TTG_PROCESS_TT_OP_RETURN(suspended_task_address, task->coroutine_id, baseobj->op(obj->output_terminals));
16861670
} else {
16871671
ttg::abort();
16881672
}
@@ -1758,7 +1742,6 @@ namespace ttg_parsec {
17581742
return PARSEC_HOOK_RETURN_DONE;
17591743
}
17601744

1761-
template <ttg::ExecutionSpace Space>
17621745
static parsec_hook_return_t static_op_noarg(parsec_task_t *parsec_task) {
17631746
task_t *task = static_cast<task_t*>(parsec_task);
17641747

@@ -1774,9 +1757,9 @@ namespace ttg_parsec {
17741757
assert(detail::parsec_ttg_caller == NULL);
17751758
detail::parsec_ttg_caller = task;
17761759
if constexpr (!ttg::meta::is_void_v<keyT>) {
1777-
TTG_PROCESS_TT_OP_RETURN(suspended_task_address, task->coroutine_id, baseobj->template op<Space>(task->key, obj->output_terminals));
1760+
TTG_PROCESS_TT_OP_RETURN(suspended_task_address, task->coroutine_id, baseobj->op(task->key, obj->output_terminals));
17781761
} else if constexpr (ttg::meta::is_void_v<keyT>) {
1779-
TTG_PROCESS_TT_OP_RETURN(suspended_task_address, task->coroutine_id, baseobj->template op<Space>(obj->output_terminals));
1762+
TTG_PROCESS_TT_OP_RETURN(suspended_task_address, task->coroutine_id, baseobj->op(obj->output_terminals));
17801763
} else // unreachable
17811764
ttg:: abort();
17821765
detail::parsec_ttg_caller = NULL;
@@ -4330,7 +4313,7 @@ namespace ttg_parsec {
43304313
void make_executable() override {
43314314
world.impl().register_tt_profiling(this);
43324315
register_static_op_function();
4333-
ttg::TTBase::make_executable();
4316+
::ttg::TTBase::make_executable();
43344317
}
43354318

43364319
/// keymap accessor
@@ -4376,6 +4359,7 @@ namespace ttg_parsec {
43764359
return ttg::device::Device(dm(key), ttg::ExecutionSpace::L0);
43774360
} else {
43784361
throw std::runtime_error("Unknown device type!");
4362+
return ttg::device::Device{};
43794363
}
43804364
};
43814365
}

0 commit comments

Comments
 (0)