12 #include <type_traits>
14 #include <xpti/xpti.hpp>
53 std::optional<uint32_t> batchSize;
54 std::optional<uint32_t> clusterBitmap;
55 std::optional<pxl::LocalityMode> locality;
57 std::optional<pxl::Map::CompletionCallback> onComplete;
58 void* onCompleteArg =
nullptr;
59 std::optional<pxl::Map::MessageCallback> onMessage;
60 void* onMessageArg =
nullptr;
61 std::optional<pxl::Map::ErrorCallback> onError;
62 void* onErrorArg =
nullptr;
70 std::optional<pxl::Stream*> stream;
71 std::optional<xpti::ProfileConfig> profileConfig;
76 std::string kernelName;
77 uint32_t taskCount = 0;
78 std::vector<ArgInfo_t> argList;
79 std::vector<std::shared_ptr<void>> constStorage;
80 std::vector<std::pair<void*, size_t>> syncTargets;
84 struct StageDirectionMetadata
86 ArgDirections directions;
89 void operator()(
void*)
const noexcept
94 inline void attachStageDirections(StageEntry& entry, ArgDirections directions)
96 if (directions.args.empty())
100 entry.constStorage.emplace_back(
nullptr, StageDirectionMetadata{std::move(directions)});
103 inline const ArgDirections& stageDirections(
const StageEntry& entry)
105 static const ArgDirections LegacyDirections;
106 for (
auto it = entry.constStorage.rbegin(); it != entry.constStorage.rend(); ++it)
108 if (
const auto* metadata = std::get_deleter<StageDirectionMetadata>(*it))
110 return metadata->directions;
113 return LegacyDirections;
118 std::vector<StageEntry> entries;
123 class StageGroupBuilder;
161 template <auto Func, typename... Args,
165 const char* kernelName =
nullptr;
166 #if defined(__PXCC_ANALYSIS__)
167 kernelName =
"__pxcc_placeholder__";
170 "Kernel function not found in KernelTraits — was it compiled with pxcc?");
174 impl::StageEntry entry;
175 entry.kernelName = kernelName;
176 entry.taskCount = taskCount;
177 entry.argList.reserve(
sizeof...(Args));
178 if constexpr (
sizeof...(Args) > 0)
180 buildArgList(entry.argList, entry.constStorage, entry.syncTargets, std::forward<Args>(args)...);
182 impl::attachStageDirections(entry, detail::MakeKernelDirections<Func>());
184 impl::StageGroup group;
185 group.entries.push_back(std::move(entry));
186 stages_.push_back(std::move(group));
196 template <
auto Func,
typename... Args,
200 template <
auto Func,
typename... Args,
244 template <
typename T,
typename... Rest>
245 void buildArgList(std::vector<ArgInfo_t>& argList, std::vector<std::shared_ptr<void>>& storage,
246 std::vector<std::pair<void*, size_t>>& syncTargets, T&& arg, Rest&&... rest)
248 addArg(argList, storage, syncTargets, std::forward<T>(arg));
249 if constexpr (
sizeof...(rest) > 0)
251 buildArgList(argList, storage, syncTargets, std::forward<Rest>(rest)...);
255 template <
typename T>
256 std::enable_if_t<std::is_pointer_v<std::decay_t<T>>>
257 addArg(std::vector<ArgInfo_t>& argList, std::vector<std::shared_ptr<void>>&,
258 std::vector<std::pair<void*, size_t>>&, T&& ptr)
263 template <
typename T>
264 std::enable_if_t<std::is_arithmetic_v<std::decay_t<T>> && !std::is_pointer_v<std::decay_t<T>> &&
265 !std::is_same_v<std::decay_t<T>,
bool>>
266 addArg(std::vector<ArgInfo_t>& argList, std::vector<std::shared_ptr<void>>& storage,
267 std::vector<std::pair<void*, size_t>>&, T&& value)
269 auto stored = std::make_shared<std::decay_t<T>>(std::forward<T>(value));
270 argList.emplace_back(ArgInfo_t{
ArgType::Constant, stored.get(),
sizeof(std::decay_t<T>)});
271 storage.push_back(stored);
274 template <
typename T>
275 void addArg(std::vector<ArgInfo_t>& argList, std::vector<std::shared_ptr<void>>& storage,
276 std::vector<std::pair<void*, size_t>>& syncTargets,
285 auto storedNdarray = std::make_shared<pxl::NDArray<T>>(ndarrayArg);
287 storage.push_back(storedNdarray);
288 const auto byteSize = ndarrayArg.
size();
289 syncTargets.emplace_back(
const_cast<void*
>(
static_cast<const void*
>(ndarrayArg.
data())),
290 byteSize > 0 ?
static_cast<size_t>(byteSize) : 0);
293 std::shared_ptr<impl::DeviceHandle> device_;
294 std::vector<impl::StageGroup> stages_;
295 uint32_t reserveCount_ = 0;
299 std::atomic<bool> started_{
false};
301 StageGroupBuilder takeStageGroup() &&;
327 template <
auto Func,
typename... Args,
331 const char* kernelName =
nullptr;
332 #if defined(__PXCC_ANALYSIS__)
333 kernelName =
"__pxcc_placeholder__";
336 "Kernel function not found in KernelTraits — was it compiled with pxcc?");
340 impl::StageEntry entry;
341 entry.kernelName = kernelName;
342 entry.taskCount = taskCount;
343 entry.argList.reserve(
sizeof...(Args));
344 if constexpr (
sizeof...(Args) > 0)
346 parent_.buildArgList(entry.argList, entry.constStorage, entry.syncTargets,
347 std::forward<Args>(args)...);
349 impl::attachStageDirections(entry, detail::MakeKernelDirections<Func>());
350 currentGroup_.entries.push_back(std::move(entry));
359 template <
auto Func,
typename... Args,
363 template <
auto Func,
typename... Args,
370 template <
auto Func,
typename... Args,
375 return parent_.stage<Func>(taskCount, std::forward<Args>(args)...);
380 template <
auto Func,
typename... Args,
384 template <
auto Func,
typename... Args,
391 return parent_.stage();
397 return parent_.run();
403 return parent_.runAsync();
459 impl::StageGroup currentGroup_;
460 bool flushed_ = false;
465 friend class detail::TaskCountConfigured;
480 return StageGroupBuilder(std::make_unique<StageBuilder>(std::move(parent_)));
483 template <
auto Func,
typename... Args,
487 stage<Func>(uint32_t{0}, std::forward<Args>(args)...);
491 template <
auto Func,
typename... Args,
495 add<Func>(uint32_t{0}, std::forward<Args>(args)...);
499 template <
auto Func,
typename... Args,
503 auto& parent = stage<Func>(uint32_t{0}, std::forward<Args>(args)...);
Execution context for kernel launches on XCENA devices.
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(StageBuilder &&) noexcept
StageBuilder & onMessage(Map::MessageCallback cb, void *arg=nullptr)
Registers a message callback on the most-recently-added stage.
StageBuilder & onError(Map::ErrorCallback cb, void *arg=nullptr)
Registers an error callback on the most-recently-added stage.
StageBuilder & batchSize(uint32_t n)
Sets per-task batch size for the most-recently-added stage. Affects only the last entry added via sta...
detail::TaskCountRequired< StageBuilder > stage(Args &&... args) &=delete
StageBuilder & reserve(uint32_t n)
StageBuilder(std::shared_ptr< impl::DeviceHandle > device)
StageBuilder & locality(LocalityMode mode)
Sets locality mode for the most-recently-added stage.
detail::TaskCountRequired< StageBuilder > stage(Args &&... args) &&
StageBuilder & stream(Stream *s)
Routes the most-recently-added stage through a user-managed stream.
StageBuilder & clusterBitmap(uint32_t bitmap)
Sets cluster bitmap for the most-recently-added stage.
std::future< LaunchResult > runAsync()
StageBuilder & onComplete(Map::CompletionCallback cb, void *arg=nullptr)
Registers a completion callback on the most-recently-added stage.
StageBuilder & enableProfile(const xpti::ProfileConfig &config)
Enables profiling for the most-recently-added stage. enableProfiling failures are logged but not prop...
StageBuilder & tasks(uint32_t taskCount)
StageGroupBuilder stage()
Collects parallel kernels within a single stage group.
StageGroupBuilder & add(uint32_t taskCount, Args &&... args)
StageGroupBuilder & onComplete(Map::CompletionCallback cb, void *arg=nullptr)
Registers a completion callback on the most-recently-added kernel.
std::future< LaunchResult > runAsync()
StageGroupBuilder & reserve(uint32_t n)
StageGroupBuilder(StageGroupBuilder &&) noexcept
StageGroupBuilder & clusterBitmap(uint32_t bitmap)
Sets cluster bitmap for the most-recently-added kernel in the active group.
StageGroupBuilder & onError(Map::ErrorCallback cb, void *arg=nullptr)
Registers an error callback on the most-recently-added kernel.
StageGroupBuilder stage()
detail::TaskCountRequired< StageBuilder > stage(Args &&... args) &=delete
detail::TaskCountRequired< StageGroupBuilder > add(Args &&... args) &=delete
StageGroupBuilder & tasks(uint32_t taskCount)
Sets the task count of the most recently added kernel in this group.
StageGroupBuilder & stream(Stream *s)
Routes the most-recently-added kernel through a user-managed stream.
StageBuilder & stage(uint32_t taskCount, Args &&... args)
StageGroupBuilder & locality(LocalityMode mode)
Sets locality mode for the most-recently-added kernel in the active group.
StageGroupBuilder & onMessage(Map::MessageCallback cb, void *arg=nullptr)
Registers a message callback on the most-recently-added kernel.
detail::TaskCountRequired< StageGroupBuilder > add(Args &&... args) &&
StageGroupBuilder & batchSize(uint32_t n)
Sets per-task batch size for the most-recently-added kernel within the active parallel group....
detail::TaskCountRequired< StageBuilder > stage(Args &&... args) &&
StageGroupBuilder & enableProfile(const xpti::ProfileConfig &config)
Enables profiling for the most-recently-added kernel. enableProfiling failures are logged but not pro...
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.