flux2-dev-qb2 / code /models /tt_dit /utils /cpp /planar_concat.cpp
stisiTT's picture
Add files using upload-large-folder tool
9aa90e0 verified
Raw History Blame Contribute Delete
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.
#include "planar_concat.hpp"
#include "transpose_avx2.hpp"
#include <atomic>
#include <condition_variable>
#include <cstring>
#include <functional>
#include <mutex>
#include <queue>
#include <thread>
#include <vector>
#include <emmintrin.h> // SSE2 (_mm_stream_si128, _mm_sfence)
#include <immintrin.h> // 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<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