PXL
stage_builder.hpp
Go to the documentation of this file.
1 // SPDX-License-Identifier: Apache-2.0
2 // Copyright 2024 XCENA Inc.
3 
4 #pragma once
5 
6 #include <atomic>
7 #include <cstdint>
8 #include <future>
9 #include <memory>
10 #include <optional>
11 #include <string>
12 #include <type_traits>
13 #include <vector>
14 #include <xpti/xpti.hpp>
15 
16 #include "pxl/direction.hpp"
17 #include "pxl/launch_result.hpp"
18 #include "pxl/map.hpp"
19 #include "pxl/module.hpp"
20 #include "pxl/ndarray.hpp"
21 #include "pxl/stream.hpp"
22 #include "pxl/task_count.hpp"
23 #include "pxl/type.hpp"
24 
25 namespace pxl
26 {
27 
28 namespace impl
29 {
30 class DeviceHandle;
31 
51 struct KnobConfig
52 {
53  std::optional<uint32_t> batchSize;
54  std::optional<uint32_t> clusterBitmap;
55  std::optional<pxl::LocalityMode> locality;
56 
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;
63 
64  // stream knob — 3-state semantics:
65  // unset (nullopt): leave Map's default stream untouched
66  // set, non-null: route through user-supplied Stream
67  // set, null: explicitly reset Map back to its default stream
68  // (Without this distinction, stream(nullptr) was a silent no-op and there
69  // was no way to opt back out of a previously set user stream.)
70  std::optional<pxl::Stream*> stream; // nullopt = leave default; non-null = user stream; null = reset to default
71  std::optional<xpti::ProfileConfig> profileConfig;
72 };
73 
74 struct StageEntry
75 {
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;
81  KnobConfig knobs;
82 };
83 
84 struct StageDirectionMetadata
85 {
86  ArgDirections directions;
87 
88  // The shared_ptr control block owns the directions; there is no pointee.
89  void operator()(void*) const noexcept
90  {
91  }
92 };
93 
94 inline void attachStageDirections(StageEntry& entry, ArgDirections directions)
95 {
96  if (directions.args.empty())
97  {
98  return;
99  }
100  entry.constStorage.emplace_back(nullptr, StageDirectionMetadata{std::move(directions)});
101 }
102 
103 inline const ArgDirections& stageDirections(const StageEntry& entry)
104 {
105  static const ArgDirections LegacyDirections;
106  for (auto it = entry.constStorage.rbegin(); it != entry.constStorage.rend(); ++it)
107  {
108  if (const auto* metadata = std::get_deleter<StageDirectionMetadata>(*it))
109  {
110  return metadata->directions;
111  }
112  }
113  return LegacyDirections;
114 }
115 
116 struct StageGroup
117 {
118  std::vector<StageEntry> entries;
119 };
120 
121 } // namespace impl
122 
123 class StageGroupBuilder;
124 
149 // [[nodiscard]] guards against building a stage chain and forgetting `.run()`.
150 class [[nodiscard]] StageBuilder
151 {
152 public:
153  StageBuilder(std::shared_ptr<impl::DeviceHandle> device);
155 
157  StageBuilder& operator=(StageBuilder&&) noexcept;
158  StageBuilder(const StageBuilder&) = delete;
159  StageBuilder& operator=(const StageBuilder&) = delete;
160 
161  template <auto Func, typename... Args,
162  detail::EnableCompatibleKernelArguments<Func, Args...> = 0>
163  StageBuilder& stage(uint32_t taskCount, Args&&... args)
164  {
165  const char* kernelName = nullptr;
166 #if defined(__PXCC_ANALYSIS__)
167  kernelName = "__pxcc_placeholder__";
168 #else
169  static_assert(sizeof(pxl::KernelTraits<Func>) > 0,
170  "Kernel function not found in KernelTraits — was it compiled with pxcc?");
171  kernelName = pxl::KernelTraits<Func>::name;
172 #endif
173 
174  impl::StageEntry entry;
175  entry.kernelName = kernelName;
176  entry.taskCount = taskCount;
177  entry.argList.reserve(sizeof...(Args));
178  if constexpr (sizeof...(Args) > 0)
179  {
180  buildArgList(entry.argList, entry.constStorage, entry.syncTargets, std::forward<Args>(args)...);
181  }
182  impl::attachStageDirections(entry, detail::MakeKernelDirections<Func>());
183 
184  impl::StageGroup group;
185  group.entries.push_back(std::move(entry));
186  stages_.push_back(std::move(group));
187  return *this;
188  }
189 
196  template <auto Func, typename... Args,
197  detail::EnableCompatibleKernelArguments<Func, Args...> = 0>
199 
200  template <auto Func, typename... Args,
201  detail::EnableCompatibleKernelArguments<Func, Args...> = 0>
202  detail::TaskCountRequired<StageBuilder> stage(Args&&... args) & = delete;
203 
205  StageBuilder& reserve(uint32_t n);
206 
209  StageBuilder& tasks(uint32_t taskCount);
210 
216  StageBuilder& batchSize(uint32_t n);
217 
219  StageBuilder& clusterBitmap(uint32_t bitmap);
220 
223 
226 
228  StageBuilder& onMessage(Map::MessageCallback cb, void* arg = nullptr);
229 
231  StageBuilder& onError(Map::ErrorCallback cb, void* arg = nullptr);
232 
235 
238  StageBuilder& enableProfile(const xpti::ProfileConfig& config);
239 
241  std::future<LaunchResult> runAsync();
242 
243 private:
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)
247  {
248  addArg(argList, storage, syncTargets, std::forward<T>(arg));
249  if constexpr (sizeof...(rest) > 0)
250  {
251  buildArgList(argList, storage, syncTargets, std::forward<Rest>(rest)...);
252  }
253  }
254 
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)
259  {
260  argList.emplace_back(ArgInfo_t{ArgType::DeviceMemory, static_cast<const void*>(ptr), 0});
261  }
262 
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)
268  {
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);
272  }
273 
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,
277  const pxl::NDArray<T>& ndarrayArg)
278  {
279  // Copy the NDArray descriptor into heap storage so we don't stash the
280  // caller's stack address into argList. The caller's NDArray may be a
281  // function-local rvalue that dies before runAsync's worker reads it;
282  // without this copy the worker would dereference a dead stack frame.
283  // The underlying data buffer (the device-memory pointer behind data())
284  // must remain valid until execution completes.
285  auto storedNdarray = std::make_shared<pxl::NDArray<T>>(ndarrayArg);
286  argList.emplace_back(ArgInfo_t{ArgType::NDArray, storedNdarray.get(), sizeof(pxl::NDArray<T>)});
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);
291  }
292 
293  std::shared_ptr<impl::DeviceHandle> device_;
294  std::vector<impl::StageGroup> stages_;
295  uint32_t reserveCount_ = 0;
296  // Atomic: same rationale as LaunchBuilder::started_ — guard against
297  // accidental cross-thread reuse, return the documented "already called"
298  // error rather than UB.
299  std::atomic<bool> started_{false};
300 
301  StageGroupBuilder takeStageGroup() &&;
302 
303  friend class StageGroupBuilder;
304  template <typename>
306 };
307 
324 class [[nodiscard]] StageGroupBuilder
325 {
326 public:
327  template <auto Func, typename... Args,
328  detail::EnableCompatibleKernelArguments<Func, Args...> = 0>
329  StageGroupBuilder& add(uint32_t taskCount, Args&&... args)
330  {
331  const char* kernelName = nullptr;
332 #if defined(__PXCC_ANALYSIS__)
333  kernelName = "__pxcc_placeholder__";
334 #else
335  static_assert(sizeof(pxl::KernelTraits<Func>) > 0,
336  "Kernel function not found in KernelTraits — was it compiled with pxcc?");
337  kernelName = pxl::KernelTraits<Func>::name;
338 #endif
339 
340  impl::StageEntry entry;
341  entry.kernelName = kernelName;
342  entry.taskCount = taskCount;
343  entry.argList.reserve(sizeof...(Args));
344  if constexpr (sizeof...(Args) > 0)
345  {
346  parent_.buildArgList(entry.argList, entry.constStorage, entry.syncTargets,
347  std::forward<Args>(args)...);
348  }
349  impl::attachStageDirections(entry, detail::MakeKernelDirections<Func>());
350  currentGroup_.entries.push_back(std::move(entry));
351  return *this;
352  }
353 
359  template <auto Func, typename... Args,
360  detail::EnableCompatibleKernelArguments<Func, Args...> = 0>
362 
363  template <auto Func, typename... Args,
364  detail::EnableCompatibleKernelArguments<Func, Args...> = 0>
366 
368  StageGroupBuilder& tasks(uint32_t taskCount);
369 
370  template <auto Func, typename... Args,
371  detail::EnableCompatibleKernelArguments<Func, Args...> = 0>
372  StageBuilder& stage(uint32_t taskCount, Args&&... args)
373  {
374  flush();
375  return parent_.stage<Func>(taskCount, std::forward<Args>(args)...);
376  }
377 
380  template <auto Func, typename... Args,
381  detail::EnableCompatibleKernelArguments<Func, Args...> = 0>
383 
384  template <auto Func, typename... Args,
385  detail::EnableCompatibleKernelArguments<Func, Args...> = 0>
386  detail::TaskCountRequired<StageBuilder> stage(Args&&... args) & = delete;
387 
389  {
390  flush();
391  return parent_.stage();
392  }
393 
395  {
396  flush();
397  return parent_.run();
398  }
399 
400  std::future<LaunchResult> runAsync()
401  {
402  flush();
403  return parent_.runAsync();
404  }
405 
407  {
408  parent_.reserve(n);
409  return *this;
410  }
411 
418 
420  StageGroupBuilder& clusterBitmap(uint32_t bitmap);
421 
424 
427 
430 
432  StageGroupBuilder& onError(Map::ErrorCallback cb, void* arg = nullptr);
433 
436 
439  StageGroupBuilder& enableProfile(const xpti::ProfileConfig& config);
440 
443 
444 private:
446  // Owning constructor: takes a heap-allocated StageBuilder when the group
447  // is created from a temporary parent (e.g. Launcher::stage() no-arg).
448  // Without this, parent_ would dangle once the temporary expires.
449  StageGroupBuilder(std::unique_ptr<StageBuilder> owned);
450  void flush();
451  StageGroupBuilder takeStageGroup() &&;
452 
453  // Heap-owned parent for the owning construction path; nullptr when the
454  // group is bound to an externally-owned StageBuilder.
455  std::unique_ptr<StageBuilder> owned_;
456  // Always-valid reference to the active parent (owned_.get() when owning,
457  // or the externally-supplied StageBuilder otherwise).
458  StageBuilder& parent_;
459  impl::StageGroup currentGroup_;
460  bool flushed_ = false;
461 
462  friend class StageBuilder;
463  friend class Launcher;
464  template <typename>
465  friend class detail::TaskCountConfigured;
466 };
467 
468 inline StageGroupBuilder StageBuilder::takeStageGroup() &&
469 {
470  return StageGroupBuilder(std::make_unique<StageBuilder>(std::move(*this)));
471 }
472 
473 inline StageGroupBuilder StageGroupBuilder::takeStageGroup() &&
474 {
475  flush();
476  if (owned_)
477  {
478  return StageGroupBuilder(std::move(owned_));
479  }
480  return StageGroupBuilder(std::make_unique<StageBuilder>(std::move(parent_)));
481 }
482 
483 template <auto Func, typename... Args,
486 {
487  stage<Func>(uint32_t{0}, std::forward<Args>(args)...);
488  return detail::TaskCountRequired<StageBuilder>(std::move(*this));
489 }
490 
491 template <auto Func, typename... Args,
494 {
495  add<Func>(uint32_t{0}, std::forward<Args>(args)...);
496  return detail::TaskCountRequired<StageGroupBuilder>(std::move(*this));
497 }
498 
499 template <auto Func, typename... Args,
502 {
503  auto& parent = stage<Func>(uint32_t{0}, std::forward<Args>(args)...);
504  return detail::TaskCountRequired<StageBuilder>(std::move(parent));
505 }
506 
507 } // namespace pxl
Execution context for kernel launches on XCENA devices.
Definition: launcher.hpp:267
std::function< void(void *message, void *arg)> MessageCallback
Callback function type for task message.
Definition: map.hpp:69
std::function< void(void *arg)> ErrorCallback
Callback function type for task error.
Definition: map.hpp:78
std::function< void(void *arg)> CompletionCallback
Callback function type for task completion.
Definition: map.hpp:59
NDArray is a class that can be used in both host and device code. NDArray is automatically divided in...
Definition: ndarray.hpp:25
T * data()
Definition: ndarray.hpp:84
std::int64_t size() const
Definition: ndarray.hpp:114
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.
LaunchResult run()
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.
Definition: stream.hpp:20
Runnable builder that preserves the value category of settings.
Definition: task_count.hpp:24
Holds a launch request that still needs its task count.
Definition: task_count.hpp:103
std::enable_if_t< KernelArgumentsCompatibleV< Func, Actual... >, int > EnableCompatibleKernelArguments
Definition: direction.hpp:184
Definition: config.hpp:11
LocalityMode
Definition: type.hpp:92
Traits class that maps a kernel function pointer to its name string.
Definition: module.hpp:30