PXL
map.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 #include <cstddef>
6 #include <cstdint>
7 #include <functional>
8 #include <memory>
9 #include <type_traits>
10 #include <vector>
11 #include <xpti/xpti.hpp>
12 
13 #include "pxl/ndarray.hpp"
14 #include "pxl/type.hpp"
15 
16 namespace xpti
17 {
18 class Profile;
19 }
20 
21 namespace pxl
22 {
23 class Function;
24 class Job;
25 class Stream;
26 
27 namespace impl
28 {
29 class MapFactory;
30 class MapImpl;
31 } // namespace impl
32 
47 class Map
48 {
49 public:
50  friend class impl::MapFactory;
51 
59  using CompletionCallback = std::function<void(void* arg)>;
60 
69  using MessageCallback = std::function<void(void* message, void* arg)>;
70 
78  using ErrorCallback = std::function<void(void* arg)>;
79 
97  void setCompletionCallback(const CompletionCallback& callback, void* arg);
98 
114  void setMessageCallback(const MessageCallback& callback, void* arg);
115 
133  void setErrorCallback(const ErrorCallback& callback, void* arg);
134 
140  void setBatchSize(const uint32_t& batchSize);
141 
174  uint32_t getBatchSize() const;
175 
188  uint32_t getTaskCount() const;
189 
195  void setClusterBitmap(const uint32_t& clusterBitmap);
196 
202  void setLocalityMode(const LocalityMode& mode);
203 
219  void setStream();
220 
236  void setStream(Stream* stream);
237 
243  uint32_t streamId() const;
244 
246  {
248  }
249 
285  template <typename... ARGS>
286  Result execute(ARGS&&... args)
287  {
288  std::vector<ArgInfo_t> argInfos;
289  if (buildArgInfo(argInfos, std::forward<ARGS>(args)...) == false)
290  {
292  }
293 
294  return execute(argInfos);
295  }
296 
329  Result execute(std::vector<ArgInfo_t>& args);
330 
344  template <typename... ARG>
345  Result setInput(ARG&&... args)
346  {
347  std::vector<ArgInfo_t> argInfos;
348  if (buildArgInfo(argInfos, std::forward<ARG>(args)...) == false)
349  {
351  }
352  return setInput(argInfos);
353  }
354 
368  template <typename... ARG>
369  Result setOutput(ARG&&... args)
370  {
371  std::vector<ArgInfo_t> argInfos;
372  if (buildArgInfo(argInfos, std::forward<ARG>(args)...) == false)
373  {
375  }
376  return setOutput(argInfos);
377  }
378 
394  Result setInput(std::vector<ArgInfo_t>& args);
395 
411  Result setOutput(std::vector<ArgInfo_t>& args);
412 
475 
518 
547 
570 
590  std::vector<KernelError_t> getKernelError();
591 
608  std::string getStats() const;
609 
619  uint64_t id() const;
620 
624  uint32_t deviceId() const;
625 
652  Result enableProfiling(const xpti::ProfileConfig& config);
653 
673  template <xpti::QueryKey Key>
674  xpti::SessionData<Key> getProfilingData(const xpti::QueryFilter& filter = {}) const
675  {
676  // Use this map's id() as sessionId
677  return xpti::getSessionData<Key>(id(), filter);
678  }
679 
698  Result exportProfiling(const std::string& path, xpti::ExportFormat format) const;
699 
705  Result exportProfilingCsv(const std::string& path = "") const;
706 
707 protected:
711  virtual ~Map() = default;
712 
713 private:
714  Map(std::unique_ptr<impl::MapImpl> impl);
715 
716  template <typename ARG, typename... ARGS>
717  typename std::enable_if<std::is_trivial<typename std::remove_reference<ARG>::type>::value, bool>::type
718  buildArgInfo(std::vector<ArgInfo_t>& argInfos, ARG&& arg, ARGS&&... args)
719  {
720  argInfos.emplace_back(ArgInfo_t{ArgType::Constant, &arg, sizeof(ARG)});
721  if constexpr (sizeof...(args) > 0)
722  {
723  return buildArgInfo(argInfos, std::forward<ARGS>(args)...);
724  }
725  return true;
726  }
727 
728  template <typename ARG, typename... ARGS>
729  typename std::enable_if<!std::is_trivial<typename std::remove_reference<ARG>::type>::value, bool>::type
730  buildArgInfo([[maybe_unused]] std::vector<ArgInfo_t>& argInfos, [[maybe_unused]] ARG&& arg, [[maybe_unused]] ARGS&&... args)
731  {
732  return false;
733  }
734 
735  template <typename ARG, typename... ARGS>
736  bool buildArgInfo(std::vector<ArgInfo_t>& argInfos, NDArray<ARG>& ndarrayArg, ARGS&&... args)
737  {
738  argInfos.emplace_back(ArgInfo_t{ArgType::NDArray, &ndarrayArg, sizeof(NDArray<ARG>)});
739  if constexpr (sizeof...(args) > 0)
740  {
741  return buildArgInfo(argInfos, std::forward<ARGS>(args)...);
742  }
743  return true;
744  }
745 
746  template <typename ARG, typename... ARGS>
747  bool buildArgInfo(std::vector<ArgInfo_t>& argInfos, NDArray<ARG>&& ndarrayArg, ARGS&&... args) = delete;
748 
749  template <typename ARG, typename... ARGS>
750  bool buildArgInfo(std::vector<ArgInfo_t>& argInfos, ARG* arg, ARGS&&... args)
751  {
752  argInfos.emplace_back(ArgInfo_t{ArgType::DeviceMemory, arg, sizeof(void*)});
753  if constexpr (sizeof...(args) > 0)
754  {
755  return buildArgInfo(argInfos, std::forward<ARGS>(args)...);
756  }
757  return true;
758  }
759 
760  std::unique_ptr<impl::MapImpl> impl_;
761 };
762 
763 } // namespace pxl
The Map class represents a map operation in the XCENA execution framework.
Definition: map.hpp:48
void setStream()
Set the default stream for the map operation.
void setErrorCallback(const ErrorCallback &callback, void *arg)
Set the error callback function and user data.
Result execute(ARGS &&... args)
Execute the map operation. This method executes the map operation with the provided input arguments.
Definition: map.hpp:286
std::function< void(void *message, void *arg)> MessageCallback
Callback function type for task message.
Definition: map.hpp:69
void setBatchSize(const uint32_t &batchSize)
Sets batchSize for the task.
friend class impl::MapFactory
Definition: map.hpp:50
uint32_t deviceId() const
Returns the device id this Map runs on.
Result execute()
Definition: map.hpp:245
void setMessageCallback(const MessageCallback &callback, void *arg)
Set the message callback function and user data.
Result setOutput(ARG &&... args)
Set the output for the map operation. This method sets the output for the map operation.
Definition: map.hpp:369
Result setOutput(std::vector< ArgInfo_t > &args)
Set the output for the map operation. This method sets the output for the map operation.
uint64_t id() const
Gets ID of map object.
Result exportProfilingCsv(const std::string &path="") const
Export profiling data to CSV files (convenience wrapper).
Result cancel()
Request cancellation of the map operation and wait until it stops.
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
ExecuteStatus getExecuteStatus() const
Gets the execution status of the map operation.
Result enableProfiling(const xpti::ProfileConfig &config)
Enable profiling for this map with the given configuration.
Progress_t getProgress() const
Get dispatch/completion progress for the current or last run.
virtual ~Map()=default
Default virtual destructor.
Result exportProfiling(const std::string &path, xpti::ExportFormat format) const
Export profiling data in the specified format.
uint32_t getTaskCount() const
Returns the task count this map was built with.
Result setInput(std::vector< ArgInfo_t > &args)
Set the input for the map operation. This method sets the input for the map operation.
void setClusterBitmap(const uint32_t &clusterBitmap)
Set the cluster bitmap for the map operation.
void setCompletionCallback(const CompletionCallback &callback, void *arg)
Set the success callback function and user data.
Result execute(std::vector< ArgInfo_t > &args)
Execute the map operation with a vector of argument information. This method executes the map operati...
std::string getStats() const
Gets the execution statistics of the map operation.
xpti::SessionData< Key > getProfilingData(const xpti::QueryFilter &filter={}) const
Get profiling data with type-safe automatic type deduction.
Definition: map.hpp:674
uint32_t getBatchSize() const
Returns the batch size the next execution will use.
void setStream(Stream *stream)
Set the stream for the map operation.
std::vector< KernelError_t > getKernelError()
Gets the error report of the map operation.
void setLocalityMode(const LocalityMode &mode)
Set the locality mode for the map operation.
Result synchronize()
Synchronize the map operation. This method synchronizes the map operation, ensuring that all tasks ar...
uint32_t streamId() const
Returns the stream ID that the map operation is associated with.
Result setInput(ARG &&... args)
Set the input for the map operation. This method sets the input for the map operation.
Definition: map.hpp:345
An asynchronous work queue for device operations.
Definition: stream.hpp:20
Definition: config.hpp:11
Result
Definition: type.hpp:55
xpti::Profile Profile
Definition: profile.hpp:35
LocalityMode
Definition: type.hpp:92
ExecuteStatus
Definition: type.hpp:13
Definition: map.hpp:17