// Fork-safe data parallelism for the runtime encode/decode kernels. // // These kernels run inside forked processes (for example PyTorch DataLoader // workers). GNU OpenMP keeps a process-wide thread pool that does not survive // fork(): once the parent has entered a parallel region -- ours, or any other // library's sharing libgomp -- the child hangs or crashes when it enters one. // This helper instead starts its threads per call and joins them before // returning, so no thread state outlives a call and fork() is always safe. // // Iterations are split into contiguous static blocks, like OpenMP's // ``schedule(static)``, and every iteration writes only its own outputs, so // results do not depend on the thread count. // // Starting a thread costs tens of microseconds, and one iteration costs from // well under a microsecond to about a hundred depending on the kernel and the // profile's primitive counts. So each call times its first iteration on the // calling thread and sizes the split of the remaining work ``W`` from it: the // calling thread starts threads one after another at ``s`` seconds each, so a // split over ``w`` threads takes about ``W / w + s * w``, least at // ``w = sqrt(W / s)``. #pragma once #include #include #include #include #include #include #include #include #include #ifdef __linux__ #include #endif namespace ac2 { // Process-wide cap on threads per call; 0 means every CPU the process may use. inline std::atomic& thread_budget() { static std::atomic value{0}; return value; } inline int available_cpus() { #ifdef __linux__ cpu_set_t allowed; CPU_ZERO(&allowed); if (sched_getaffinity(0, sizeof(allowed), &allowed) == 0) { return std::max(1, CPU_COUNT(&allowed)); } #endif return std::max(1, static_cast(std::thread::hardware_concurrency())); } // Measured cost of starting and joining one thread (seconds). constexpr double kThreadStartSeconds = 30e-6; // Threads for ``count`` iterations of ``seconds_each``: at most ``requested`` // when positive, at most the process budget, and at most ``sqrt(W / s)``. inline int64_t plan_threads(int64_t count, int requested, double seconds_each) { const int budget = thread_budget().load(); int64_t threads = budget > 0 ? budget : available_cpus(); if (requested > 0) threads = std::min(threads, requested); const double by_work = std::sqrt(count * seconds_each / kThreadStartSeconds); if (by_work < static_cast(threads)) threads = static_cast(by_work); return std::max(1, std::min(threads, count)); } // Call ``body(index)`` for every index in [0, count). Iteration 0 runs first // on the calling thread and sizes the split of the rest; the calling thread // also runs the first block of the rest. The first exception thrown by any // block is rethrown after all threads have joined. template void parallel_for(int64_t count, int threads, Body&& body) { if (count <= 0) return; const auto start = std::chrono::steady_clock::now(); body(0); const double seconds_each = std::chrono::duration(std::chrono::steady_clock::now() - start).count(); const int64_t rest = count - 1; if (rest == 0) return; const int64_t workers = plan_threads(rest, threads, seconds_each); auto run = [&](int64_t worker) { const int64_t begin = 1 + rest * worker / workers; const int64_t end = 1 + rest * (worker + 1) / workers; for (int64_t index = begin; index < end; ++index) body(index); }; if (workers == 1) { run(0); return; } std::exception_ptr failure; std::mutex failure_lock; auto guarded = [&](int64_t worker) { try { run(worker); } catch (...) { std::lock_guard hold(failure_lock); if (!failure) failure = std::current_exception(); } }; std::vector pool; pool.reserve(static_cast(workers - 1)); try { for (int64_t worker = 1; worker < workers; ++worker) pool.emplace_back(guarded, worker); } catch (...) { for (auto& thread : pool) thread.join(); throw; } guarded(0); for (auto& thread : pool) thread.join(); if (failure) std::rethrow_exception(failure); } } // namespace ac2