11 #include <type_traits>
13 #include <xpti/xpti.hpp>
30 class LaunchBuilderImpl;
31 class LaunchBuilderTestAccess;
234 LaunchBuilder(std::unique_ptr<impl::LaunchBuilderImpl> impl);
236 std::unique_ptr<impl::LaunchBuilderImpl> impl_;
241 std::atomic<bool> started_{
false};
244 friend class impl::LaunchBuilderTestAccess;
334 template <auto Func, typename... Args,
338 const char* kernelName =
nullptr;
340 #if defined(__PXCC_ANALYSIS__)
341 kernelName =
"__pxcc_placeholder__";
344 "Kernel function not found in KernelTraits — was it compiled with pxcc?");
348 std::vector<ArgInfo_t> argList;
349 std::vector<std::shared_ptr<void>> storage;
350 std::vector<std::pair<void*, size_t>> syncTargets;
351 auto directions = detail::MakeKernelDirections<Func>();
352 argList.reserve(
sizeof...(Args));
353 if constexpr (
sizeof...(Args) > 0)
355 buildArgList(argList, storage, syncTargets, std::forward<Args>(args)...);
358 if (impl_ ==
nullptr)
360 initFromArgs(argList, syncTargets);
364 return makeEmptyBuilder();
367 return execute(kernelName, taskCount,
368 std::move(argList), std::move(storage), std::move(syncTargets),
369 std::move(directions));
382 template <
auto Func,
typename... Args,
387 execute<Func>(uint32_t{0}, std::forward<Args>(args)...));
406 template <
auto Func,
typename... Args,
410 if (impl_ ==
nullptr)
412 std::vector<ArgInfo_t> tempArgList;
413 std::vector<std::shared_ptr<void>> tempStorage;
414 std::vector<std::pair<void*, size_t>> tempSyncTargets;
415 tempArgList.
reserve(
sizeof...(Args));
416 if constexpr (
sizeof...(Args) > 0)
418 buildArgList(tempArgList, tempStorage, tempSyncTargets, std::forward<Args>(args)...);
420 initFromArgs(tempArgList, tempSyncTargets);
425 builder.
stage<Func>(taskCount, std::forward<Args>(args)...);
436 template <
auto Func,
typename... Args,
441 stage<Func>(uint32_t{0}, std::forward<Args>(args)...));
480 std::vector<ArgInfo_t> argList,
481 std::vector<std::shared_ptr<void>> storage,
482 std::vector<std::pair<void*, size_t>> syncTargets);
485 std::vector<ArgInfo_t> argList,
486 std::vector<std::shared_ptr<void>> storage,
487 std::vector<std::pair<void*, size_t>> syncTargets,
490 void initFromArgs(
const std::vector<ArgInfo_t>& argList,
491 const std::vector<std::pair<void*, size_t>>& syncTargets);
494 template <
typename T,
typename... Rest>
495 void buildArgList(std::vector<ArgInfo_t>& argList, std::vector<std::shared_ptr<void>>& storage,
496 std::vector<std::pair<void*, size_t>>& syncTargets, T&& arg, Rest&&... rest)
498 addArg(argList, storage, syncTargets, std::forward<T>(arg));
499 if constexpr (
sizeof...(rest) > 0)
501 buildArgList(argList, storage, syncTargets, std::forward<Rest>(rest)...);
505 template <
typename T>
506 std::enable_if_t<std::is_pointer_v<std::decay_t<T>>>
507 addArg(std::vector<ArgInfo_t>& argList, std::vector<std::shared_ptr<void>>&,
508 std::vector<std::pair<void*, size_t>>&, T&& ptr)
513 template <
typename T>
514 std::enable_if_t<std::is_arithmetic_v<std::decay_t<T>> && !std::is_pointer_v<std::decay_t<T>> &&
515 !std::is_same_v<std::decay_t<T>,
bool>>
516 addArg(std::vector<ArgInfo_t>& argList, std::vector<std::shared_ptr<void>>& storage,
517 std::vector<std::pair<void*, size_t>>&, T&& value)
519 auto stored = std::make_shared<std::decay_t<T>>(std::forward<T>(value));
520 argList.emplace_back(ArgInfo_t{
ArgType::Constant, stored.get(),
sizeof(std::decay_t<T>)});
521 storage.push_back(stored);
524 template <
typename T>
525 void addArg(std::vector<ArgInfo_t>& argList, std::vector<std::shared_ptr<void>>& storage,
526 std::vector<std::pair<void*, size_t>>& syncTargets,
541 auto storedNdarray = std::make_shared<pxl::NDArray<T>>(ndarrayArg);
543 storage.push_back(storedNdarray);
545 const auto byteSize = ndarrayArg.
size();
546 syncTargets.emplace_back(
const_cast<void*
>(
static_cast<const void*
>(ndarrayArg.
data())),
547 byteSize > 0 ?
static_cast<size_t>(byteSize) : 0);
550 std::shared_ptr<impl::DeviceHandle> impl_;
Configurable kernel execution request.
LaunchBuilder & enableProfile(const xpti::ProfileConfig &config)
Enables profiling for this kernel.
LaunchBuilder(LaunchBuilder &&other) noexcept
LaunchBuilder & clusterBitmap(uint32_t bitmap)
Sets the cluster bitmap (0 = use all clusters).
LaunchBuilder & tasks(uint32_t taskCount)
Sets how many tasks invoke the kernel.
LaunchBuilder & operator=(const LaunchBuilder &)=delete
LaunchBuilder & batchSize(uint32_t n)
Sets per-task batch size for kernel execution. If not called, the runtime uses an implementation-defi...
LaunchBuilder & onError(Map::ErrorCallback cb, void *arg=nullptr)
Registers an error callback for execution failure.
LaunchBuilder & operator=(LaunchBuilder &&other) noexcept
LaunchResult run()
Executes the kernel and waits for completion.
LaunchBuilder & stream(Stream *s)
Routes the kernel through a user-managed stream.
LaunchBuilder & onComplete(Map::CompletionCallback cb, void *arg=nullptr)
Registers a completion callback for successful execution.
LaunchBuilder(const LaunchBuilder &)=delete
LaunchBuilder & reserve(uint32_t n)
Sets the number of Sub resources to reserve (Mid-Level).
LaunchBuilder & locality(LocalityMode mode)
Sets the locality mode for this kernel. If not called, the runtime uses an implementation-defined def...
LaunchBuilder & onMessage(Map::MessageCallback cb, void *arg=nullptr)
Registers a message callback for device-originated messages.
std::future< LaunchResult > runAsync()
Executes the kernel asynchronously.
Execution context for kernel launches on XCENA devices.
uint32_t deviceId() const
Returns the device ID.
StageGroupBuilder stage()
Begins a parallel stage group.
Launcher(uint32_t deviceId)
Constructs a Launcher for an explicit device ID.
StageBuilder stage(uint32_t taskCount, Args &&... args)
Begins a stage chain with a sequential kernel stage.
Launcher(Launcher &&) noexcept
LaunchBuilder execute(uint32_t taskCount, Args &&... args)
Launches a kernel for deferred execution.
Launcher & operator=(const Launcher &)=delete
Launcher(const Launcher &)=delete
detail::TaskCountRequired< LaunchBuilder > execute(Args &&... args)
Launches a kernel with only the kernel's own arguments.
bool isValid() const
Returns whether the launcher was initialized successfully.
detail::TaskCountRequired< StageBuilder > stage(Args &&... args)
Begins a stage chain with only the kernel's own arguments.
Launcher()
Constructs a Launcher with deferred device detection.
std::function< void(void *message, void *arg)> MessageCallback
Callback function type for task message.
std::function< void(void *arg)> ErrorCallback
Callback function type for task error.
std::function< void(void *arg)> CompletionCallback
Callback function type for task completion.
NDArray is a class that can be used in both host and device code. NDArray is automatically divided in...
std::int64_t size() const
Builds a chain of kernel stages for sequential/parallel execution.
StageBuilder & stage(uint32_t taskCount, Args &&... args)
StageBuilder & reserve(uint32_t n)
Collects parallel kernels within a single stage group.
An asynchronous work queue for device operations.
Holds a launch request that still needs its task count.
std::enable_if_t< KernelArgumentsCompatibleV< Func, Actual... >, int > EnableCompatibleKernelArguments
Traits class that maps a kernel function pointer to its name string.