MCPcopy Create free account
hub / github.com/NVIDIA/DALI / MakeExec2Config

Function MakeExec2Config

dali/pipeline/executor/executor_factory.cc:30–94  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

28namespace {
29
30auto MakeExec2Config(int batch_size, int num_thread, int device_id,
31 size_t bytes_per_sample_hint, ExecutorFlags flags,
32 QueueSizes prefetch_queue_depth) {
33 exec2::Executor2::Config cfg{};
34 cfg.async_output = false;
35 cfg.set_affinity = Test(flags, ExecutorFlags::SetAffinity);
36 cfg.thread_pool_threads = num_thread;
37 // TODO(michalz): Expose the thread configuration in the Pipeline (?)
38 // Alternatively, use cooperative parallelism with the CPU thread pool (?)
39
40 static int exec2_max_threads = []() {
41 const char *env = getenv("DALI_EXEC2_MAX_THREADS");
42 constexpr int kDefaultMaxThreads = 4;
43 if (env) {
44 int value = atoi(env);
45 if (value <= 0)
46 value = kDefaultMaxThreads;
47 return value;
48 } else {
49 return kDefaultMaxThreads;
50 }
51 }();
52 static std::optional<int> exec2_num_threads = []()->std::optional<int> {
53 const char *env = getenv("DALI_EXEC2_NUM_THREADS");
54 if (env) {
55 int value = atoi(env);
56 if (value >= 0)
57 return value;
58 }
59 return std::nullopt;;
60 }();
61
62 cfg.operator_threads = exec2_num_threads.value_or(std::min(num_thread, exec2_max_threads));
63 if (device_id != CPU_ONLY_DEVICE_ID)
64 cfg.device = device_id;
65 cfg.max_batch_size = batch_size;
66 cfg.cpu_queue_depth = prefetch_queue_depth.cpu_size;
67 cfg.gpu_queue_depth = prefetch_queue_depth.gpu_size;
68 cfg.queue_policy = exec2::QueueDepthPolicy::Legacy;
69 switch (flags & ExecutorFlags::StreamPolicyMask) {
70 case ExecutorFlags::StreamPolicyPerOperator:
71 cfg.stream_policy = exec2::StreamPolicy::PerOperator;
72 break;
73 case ExecutorFlags::StreamPolicySingle:
74 cfg.stream_policy = exec2::StreamPolicy::Single;
75 break;
76 case ExecutorFlags::StreamPolicyPerBackend:
77 default:
78 cfg.stream_policy = exec2::StreamPolicy::PerBackend;
79 break;
80 }
81 switch (flags & ExecutorFlags::ConcurrencyMask) {
82 case ExecutorFlags::ConcurrencyNone:
83 cfg.concurrency = exec2::OperatorConcurrency::None;
84 break;
85 case ExecutorFlags::ConcurrencyFull:
86 cfg.concurrency = exec2::OperatorConcurrency::Full;
87 break;

Callers 1

GetExecutorImplFunction · 0.85

Calls 2

TestFunction · 0.85
minFunction · 0.50

Tested by

no test coverage detected