//===----------------------------------------------------------------------===// // // Part of CUDASTF in CUDA C++ Core Libraries, // under the Apache License v2.0 with LLVM Exceptions. // See https://llvm.org/LICENSE.txt for license information. // SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception // SPDX-FileCopyrightText: Copyright (c) 2022-2024 NVIDIA CORPORATION & AFFILIATES. // //===----------------------------------------------------------------------===// /** * @file * @brief Generate a library call from nested CUDA graphs generated using algorithms */ #include using namespace cuda::experimental::stf; // Some fake library doing MATH void libMATH(graph_ctx ctx, logical_data> x, logical_data> y) { // We only want to have kernels with 4 CTAs to stress the system auto spec = par<4>(par<128>()); ctx.launch(spec, exec_place::current_device(), x.read(), y.write()).set_symbol("MATH1")->* [] __device__(auto t, auto x, auto y) { for (auto i : t.apply_partition(shape(x))) { y(i) = cos(cos(x(i))); } }; ctx.launch(spec, exec_place::current_device(), x.write(), y.read()).set_symbol("MATH2")->* [] __device__(auto t, auto x, auto y) { for (auto i : t.apply_partition(shape(x))) { x(i) = sin(sin(y(i))); }; }; } template void libMATH_AS_GRAPH(context_t& ctx, logical_data> x, logical_data> y) { static algorithm alg; alg.run_as_task(libMATH, ctx, x.rw(), y.write()); } // Some fake lib doing a SWAP template void libSWAP(context_t& ctx, logical_data> x, logical_data> y) { // We only want to have kernels with 4 CTAs to stress the system auto spec = par<4>(par<128>()); ctx.launch(spec, exec_place::current_device(), x.rw(), y.rw()).set_symbol("SWAP")->* [] __device__(auto t, auto x, auto y) { for (auto i : t.apply_partition(shape(x))) { auto tmp = x(i); x(i) = y(i); y(i) = tmp; } }; } template logical_data> libCOPY(context_t& ctx, logical_data> x) { logical_data> res = ctx.logical_data(x.shape()); // We only want to have kernels with 4 CTAs to stress the system auto spec = par<4>(par<128>()); ctx.launch(spec, exec_place::current_device(), x.read(), res.write()).set_symbol("SWAP")->* [] __device__(auto t, auto x, auto res) { for (auto i : t.apply_partition(shape(x))) { res(i) = x(i); } }; return res; } int main() { nvtx_range r("run"); stream_ctx ctx; const size_t N = 256 * 1024; const size_t K = 8; logical_data> lX[K]; logical_data> lY[K]; for (size_t i = 0; i < K; i++) { lX[i] = ctx.logical_data(N); lY[i] = ctx.logical_data(N); ctx.parallel_for(lX[i].shape(), lX[i].write(), lY[i].write()).set_symbol("INIT")->* [] __device__(size_t i, auto x, auto y) { x(i) = 2.0 * i + 12.0; y(i) = -3.0 * i + 17.0; }; } for (size_t i = 0; i < K; i++) { auto tmp = libCOPY(ctx, lX[i]); libSWAP(ctx, tmp, lY[i]); libMATH_AS_GRAPH(ctx, lX[i], lY[i]); libSWAP(ctx, lX[i], lY[i]); } ctx.finalize(); }