Download code/models/tt_dit/utils/cpp/planar_concat.cpp from stisiTT/flux2-dev-qb2: direct link, hf CLI and curl.
- Browser
- Download file 13.6 kB
-
https://huggingface.co/stisiTT/flux2-dev-qb2/resolve/main/code/models/tt_dit/utils/cpp/planar_concat.cpp
- Command line
-
hf download hf://stisiTT/flux2-dev-qb2/code/models/tt_dit/utils/cpp/planar_concat.cpp
-
curl -L -o planar_concat.cpp https://huggingface.co/stisiTT/flux2-dev-qb2/resolve/main/code/models/tt_dit/utils/cpp/planar_concat.cpp
13.6 kB
| // 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. | |
| 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<std::mutex> lock(mu_); | |
| stop_ = true; | |
| } | |
| cv_.notify_all(); | |
| for (auto& w : workers_) { | |
| w.join(); | |
| } | |
| } | |
| // Submit `n_tasks` tasks | |
| template <typename Fn> | |
| void run(int n_tasks, Fn fn) { | |
| if (n_tasks <= 0) { | |
| return; | |
| } | |
| std::atomic<int> remaining(n_tasks); | |
| std::atomic<int> next_idx(0); | |
| std::mutex done_mu; | |
| std::condition_variable done_cv; | |
| { | |
| std::unique_lock<std::mutex> 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<std::mutex> lk(done_mu); | |
| done_cv.notify_one(); | |
| } | |
| }); | |
| } | |
| } | |
| cv_.notify_all(); | |
| std::unique_lock<std::mutex> 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<void()> task; | |
| { | |
| std::unique_lock<std::mutex> 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<std::thread> workers_; | |
| std::queue<std::function<void()>> 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<int>(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<uintptr_t>(dst); | |
| if ((addr & 31u) == 0 && (n & 31u) == 0) { | |
| for (size_t i = 0; i < n; i += 32) { | |
| __m256i v = _mm256_loadu_si256(reinterpret_cast<const __m256i*>(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<const __m128i*>(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<std::ptrdiff_t>(row_stride) + h * static_cast<std::ptrdiff_t>(plane_W), | |
| src + (t * static_cast<std::ptrdiff_t>(src_h_per) + h) * static_cast<std::ptrdiff_t>(src_w_per), | |
| static_cast<size_t>(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<std::ptrdiff_t>(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<std::ptrdiff_t>(h) * src_w_per * T; | |
| uint8_t* dst_h = dst + static_cast<std::ptrdiff_t>(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<std::ptrdiff_t>(w_tile) * 32 * T; | |
| uint8_t* dst_w = dst_h + static_cast<std::ptrdiff_t>(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<std::ptrdiff_t>(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<std::ptrdiff_t>(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<std::ptrdiff_t>(n_w_full_tiles) * 32 * T; | |
| uint8_t* dst_w = dst_h + static_cast<std::ptrdiff_t>(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<std::ptrdiff_t>(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<std::ptrdiff_t>(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<std::ptrdiff_t>(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<std::ptrdiff_t>(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<ShardView>& 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<int>(shards.size()); | |
| pool.run(n, [&](int idx) { | |
| const ShardView& sv = shards[idx]; | |
| uint8_t* dst_base = out + static_cast<std::ptrdiff_t>(plane_offset) + | |
| static_cast<std::ptrdiff_t>(sv.r) * h_per * plane_W + | |
| static_cast<std::ptrdiff_t>(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<ShardView>& y_shards, | |
| int y_h_per, | |
| int y_w_per, | |
| const std::vector<ShardView>& cb_shards, | |
| int uv_h_per, | |
| int uv_w_per, | |
| const std::vector<ShardView>& 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<Task> tasks; | |
| tasks.reserve(y_shards.size() + cb_shards.size() + cr_shards.size()); | |
| auto add_component = | |
| [&](const std::vector<ShardView>& 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<int>(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<std::ptrdiff_t>(tk.plane_offset) + | |
| static_cast<std::ptrdiff_t>(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 | |