// SPDX-FileCopyrightText: (c) 2026 Tenstorrent USA, Inc. // SPDX-License-Identifier: Apache-2.0 // // Outer scatter loop + persistent std::thread pool for the C++ planar // concat path. Inner SIMD work lives in transpose_avx2.cpp. #include "planar_concat.hpp" #include "transpose_avx2.hpp" #include #include #include #include #include #include #include #include #include // SSE2 (_mm_stream_si128, _mm_sfence) #include // AVX2 (_mm256_stream_si256) namespace tt_dit_planar { // --------------------------------------------------------------------------- Static thread pool namespace { class ThreadPool { public: explicit ThreadPool(int n_threads) : n_threads_(n_threads), stop_(false) { workers_.reserve(n_threads); for (int i = 0; i < n_threads; ++i) { workers_.emplace_back([this] { run_worker(); }); } } ~ThreadPool() { { std::unique_lock lock(mu_); stop_ = true; } cv_.notify_all(); for (auto& w : workers_) { w.join(); } } // Submit `n_tasks` tasks template void run(int n_tasks, Fn fn) { if (n_tasks <= 0) { return; } std::atomic remaining(n_tasks); std::atomic next_idx(0); std::mutex done_mu; std::condition_variable done_cv; { std::unique_lock lock(mu_); for (int i = 0; i < n_tasks; ++i) { tasks_.emplace([&, this] { int idx; while ((idx = next_idx.fetch_add(1, std::memory_order_relaxed)) < n_tasks) { fn(idx); } if (remaining.fetch_sub(1, std::memory_order_acq_rel) == 1) { std::lock_guard lk(done_mu); done_cv.notify_one(); } }); } } cv_.notify_all(); std::unique_lock lk(done_mu); done_cv.wait(lk, [&] { return remaining.load() == 0; }); } int n_threads() const { return n_threads_; } private: void run_worker() { for (;;) { std::function task; { std::unique_lock lock(mu_); cv_.wait(lock, [this] { return stop_ || !tasks_.empty(); }); if (stop_ && tasks_.empty()) { return; } task = std::move(tasks_.front()); tasks_.pop(); } task(); } } int n_threads_; std::vector workers_; std::queue> tasks_; std::mutex mu_; std::condition_variable cv_; bool stop_; }; // Process-wide singleton static int g_requested_threads = 0; static std::once_flag g_init_flag; static ThreadPool* g_pool = nullptr; ThreadPool& get_pool() { std::call_once(g_init_flag, [] { int n = g_requested_threads > 0 ? g_requested_threads : static_cast(std::thread::hardware_concurrency()); if (n <= 0) { n = 1; } if (n > 8) { n = 8; } g_pool = new ThreadPool(n); }); return *g_pool; } } // namespace void set_thread_pool_size(int n_threads) { if (n_threads > 0) { g_requested_threads = n_threads; } } // --------------------------------------------------------------------------- Inner per-shard scatter kernels namespace { // Non-temporal copy of `n` bytes: writes via SSE/AVX streaming stores so the destination cache lines aren't pulled static inline void stream_copy_n(uint8_t* __restrict dst, const uint8_t* __restrict src, size_t n) { const uintptr_t addr = reinterpret_cast(dst); if ((addr & 31u) == 0 && (n & 31u) == 0) { for (size_t i = 0; i < n; i += 32) { __m256i v = _mm256_loadu_si256(reinterpret_cast(src + i)); _mm256_stream_si256(reinterpret_cast<__m256i*>(dst + i), v); } return; } if ((addr & 15u) == 0 && (n & 15u) == 0) { for (size_t i = 0; i < n; i += 16) { __m128i v = _mm_loadu_si128(reinterpret_cast(src + i)); _mm_stream_si128(reinterpret_cast<__m128i*>(dst + i), v); } return; } std::memcpy(dst, src, n); } // CTHW: src layout is (T, h_per, w_per), W innermost (stride 1) void scatter_one_cthw( const uint8_t* src, uint8_t* dst, int T, int src_h_per, int src_w_per, int valid_h, int valid_w, int plane_W, int row_stride) { for (int t = 0; t < T; ++t) { for (int h = 0; h < valid_h; ++h) { stream_copy_n( dst + t * static_cast(row_stride) + h * static_cast(plane_W), src + (t * static_cast(src_h_per) + h) * static_cast(src_w_per), static_cast(valid_w)); } } // Drain WC buffers so subsequent readers (other threads, the Python caller) see the data _mm_sfence(); } // CHWT: src layout is (h_per, w_per, T), T innermost (stride 1) void scatter_one_chwt( const uint8_t* src, uint8_t* dst, int T, int src_h_per, int src_w_per, int valid_h, int valid_w, int plane_W, int row_stride) { (void)src_h_per; // CHWT source h-stride is src_w_per*T; h_per only bounds the loop (valid_h). const std::ptrdiff_t src_w_stride = static_cast(T); // bytes between adjacent w in source const int n_w_full_tiles = valid_w / 32; const int w_tail = valid_w - n_w_full_tiles * 32; // padded tail cols are never written const int n_t_full_tiles = T / 32; const int t_tail = T - n_t_full_tiles * 32; for (int h = 0; h < valid_h; ++h) { const uint8_t* src_h = src + static_cast(h) * src_w_per * T; uint8_t* dst_h = dst + static_cast(h) * plane_W; // Full 32×32 (W × T) tiles. for (int w_tile = 0; w_tile < n_w_full_tiles; ++w_tile) { const uint8_t* src_w = src_h + static_cast(w_tile) * 32 * T; uint8_t* dst_w = dst_h + static_cast(w_tile) * 32; for (int t_tile = 0; t_tile < n_t_full_tiles; ++t_tile) { const uint8_t* src_tile = src_w + t_tile * 32; uint8_t* dst_tile = dst_w + static_cast(t_tile) * 32 * row_stride; transpose_32x32_u8(src_tile, src_w_stride, dst_tile, row_stride); } if (t_tail > 0) { const uint8_t* src_tile = src_w + n_t_full_tiles * 32; uint8_t* dst_tile = dst_w + static_cast(n_t_full_tiles) * 32 * row_stride; transpose_32xN_u8(src_tile, src_w_stride, dst_tile, row_stride, t_tail); } } // W-tail (when w_per % 32 != 0) if (w_tail > 0) { const uint8_t* src_w = src_h + static_cast(n_w_full_tiles) * 32 * T; uint8_t* dst_w = dst_h + static_cast(n_w_full_tiles) * 32; if (w_tail == 16) { for (int t_tile = 0; t_tile < n_t_full_tiles; ++t_tile) { transpose_32x16_u8( src_w + t_tile * 32, src_w_stride, dst_w + static_cast(t_tile) * 32 * row_stride, row_stride); } if (t_tail > 0) { // T-tail × 16 W-cols: small enough to keep the bounce buffer alignas(32) uint8_t tmp_src[32 * 32]; alignas(32) uint8_t tmp_dst[32 * 32]; std::memset(tmp_src, 0, sizeof(tmp_src)); for (int i = 0; i < 16; ++i) { std::memcpy(tmp_src + i * 32, src_w + i * src_w_stride + n_t_full_tiles * 32, t_tail); } transpose_32x32_u8(tmp_src, 32, tmp_dst, 32); uint8_t* dst_tile = dst_w + static_cast(n_t_full_tiles) * 32 * row_stride; for (int i = 0; i < t_tail; ++i) { std::memcpy(dst_tile + i * row_stride, tmp_dst + i * 32, 16); } } } else { alignas(32) uint8_t tmp_src[32 * 32]; alignas(32) uint8_t tmp_dst[32 * 32]; for (int t_tile = 0; t_tile < n_t_full_tiles; ++t_tile) { std::memset(tmp_src, 0, sizeof(tmp_src)); for (int i = 0; i < w_tail; ++i) { std::memcpy(tmp_src + i * 32, src_w + i * src_w_stride + t_tile * 32, 32); } transpose_32x32_u8(tmp_src, 32, tmp_dst, 32); uint8_t* dst_tile = dst_w + static_cast(t_tile) * 32 * row_stride; for (int i = 0; i < 32; ++i) { std::memcpy(dst_tile + i * row_stride, tmp_dst + i * 32, w_tail); } } if (t_tail > 0) { std::memset(tmp_src, 0, sizeof(tmp_src)); for (int i = 0; i < w_tail; ++i) { std::memcpy(tmp_src + i * 32, src_w + i * src_w_stride + n_t_full_tiles * 32, t_tail); } transpose_32x32_u8(tmp_src, 32, tmp_dst, 32); uint8_t* dst_tile = dst_w + static_cast(n_t_full_tiles) * 32 * row_stride; for (int i = 0; i < t_tail; ++i) { std::memcpy(dst_tile + i * row_stride, tmp_dst + i * 32, w_tail); } } } } } } } // namespace void scatter_component( const std::vector& shards, DimOrder dim_order, uint8_t* out, int T, int plane_offset, int plane_W, int row_stride, int h_per, int w_per) { auto& pool = get_pool(); const int n = static_cast(shards.size()); pool.run(n, [&](int idx) { const ShardView& sv = shards[idx]; uint8_t* dst_base = out + static_cast(plane_offset) + static_cast(sv.r) * h_per * plane_W + static_cast(sv.c) * w_per; if (dim_order == DimOrder::CTHW) { scatter_one_cthw(sv.data, dst_base, T, h_per, w_per, h_per, w_per, plane_W, row_stride); } else { scatter_one_chwt(sv.data, dst_base, T, h_per, w_per, h_per, w_per, plane_W, row_stride); } }); } void planar_concat( const std::vector& y_shards, int y_h_per, int y_w_per, const std::vector& cb_shards, int uv_h_per, int uv_w_per, const std::vector& cr_shards, DimOrder dim_order, int T, int H, int W, int out_H, int out_W, uint8_t* out) { (void)H; (void)W; // Output geometry is the logical (cropped) frame; sources keep the padded per-shard dims. const int out_Hu = out_H / 2; const int out_Wu = out_W / 2; const int out_hw = out_H * out_W; const int out_uv = out_Hu * out_Wu; const int row_stride = out_hw + 2 * out_uv; // Build a flat task list across all 3 components so the thread pool can load-balance Y struct Task { const ShardView* shard; int plane_offset; int plane_W; int h_per; int w_per; int bound_h; int bound_w; }; std::vector tasks; tasks.reserve(y_shards.size() + cb_shards.size() + cr_shards.size()); auto add_component = [&](const std::vector& s, int plane_off, int plane_w_arg, int h_p, int w_p, int bnd_h, int bnd_w) { for (const auto& sv : s) { tasks.push_back({&sv, plane_off, plane_w_arg, h_p, w_p, bnd_h, bnd_w}); } }; add_component(y_shards, 0, out_W, y_h_per, y_w_per, out_H, out_W); add_component(cb_shards, out_hw, out_Wu, uv_h_per, uv_w_per, out_Hu, out_Wu); add_component(cr_shards, out_hw + out_uv, out_Wu, uv_h_per, uv_w_per, out_Hu, out_Wu); auto& pool = get_pool(); pool.run(static_cast(tasks.size()), [&](int idx) { const Task& tk = tasks[idx]; const ShardView& sv = *tk.shard; const int r0 = sv.r * tk.h_per; const int c0 = sv.c * tk.w_per; const int valid_h = (tk.h_per < tk.bound_h - r0) ? tk.h_per : (tk.bound_h - r0); const int valid_w = (tk.w_per < tk.bound_w - c0) ? tk.w_per : (tk.bound_w - c0); if (valid_h <= 0 || valid_w <= 0) { return; // shard lies entirely in the padded tail } uint8_t* dst_base = out + static_cast(tk.plane_offset) + static_cast(r0) * tk.plane_W + c0; if (dim_order == DimOrder::CTHW) { scatter_one_cthw(sv.data, dst_base, T, tk.h_per, tk.w_per, valid_h, valid_w, tk.plane_W, row_stride); } else { scatter_one_chwt(sv.data, dst_base, T, tk.h_per, tk.w_per, valid_h, valid_w, tk.plane_W, row_stride); } }); } } // namespace tt_dit_planar