aboutsummaryrefslogtreecommitdiff
path: root/main.cpp
diff options
context:
space:
mode:
Diffstat (limited to 'main.cpp')
-rw-r--r--main.cpp179
1 files changed, 179 insertions, 0 deletions
diff --git a/main.cpp b/main.cpp
new file mode 100644
index 0000000..bd03139
--- /dev/null
+++ b/main.cpp
@@ -0,0 +1,179 @@
+#include <hip/hip_runtime.h>
+#include <source_location>
+#include <stdexcept>
+#include <string>
+#include <vector>
+#include <cmath>
+#include <algorithm>
+
+
+inline void hip_check(
+ hipError_t err,
+ std::source_location loc = std::source_location::current())
+{
+ if (err != hipSuccess) {
+ throw std::runtime_error(
+ std::string(hipGetErrorString(err)) +
+ " at " + loc.file_name() + ":" + std::to_string(loc.line()));
+ }
+}
+
+
+struct PyramidLevel {
+ int width = 0;
+ int height = 0;
+ float4* color_smax = nullptr; /* rgb s_max*/
+ float4* tensor = nullptr; /* x=E, y=F, z=G, w= */
+};
+
+struct KuwaharaPyramid {
+ std::vector<PyramidLevel> levels; /* 0 = finest, back() = coarsest */
+};
+
+struct KuwaharaParams {
+ float radius;
+ float alpha;
+ int num_sectors; /* 4 or 8 */
+ float q;
+ float tau_w;
+ float p_s, p_d, tau_v;
+};
+
+/* allocations */
+
+PyramidLevel make_level(int width, int height) {
+ PyramidLevel lvl{width, height};
+ hip_check(hipMalloc(&lvl.color_smax, sizeof(float4) * width * height));
+ hip_check(hipMalloc(&lvl.tensor, sizeof(float4) * width * height));
+ return lvl;
+}
+
+void free_level(PyramidLevel& lvl) {
+ hipFree(lvl.color_smax);
+ hipFree(lvl.tensor);
+ lvl.color_smax = nullptr;
+ lvl.tensor = nullptr;
+}
+
+KuwaharaPyramid make_pyramid(int base_width, int base_height, int num_levels) {
+ KuwaharaPyramid pyr;
+ pyr.levels.reserve(num_levels);
+ int w = base_width, h = base_height;
+ for (int i = 0; i < num_levels; ++i) {
+ pyr.levels.push_back(make_level(w, h));
+ w = std::max(1, w / 2);
+ h = std::max(1, h / 2);
+ }
+ return pyr;
+}
+
+void destroy_pyramid(KuwaharaPyramid& pyr) {
+ for (auto& lvl : pyr.levels) {
+ free_level(lvl);
+ }
+ pyr.levels.clear();
+}
+
+/* Lanczos3 downsample */
+
+__device__ __forceinline__ float lanczos3(float x) {
+ const float a = 3.0f;
+ if (x == 0.0f) return 1.0f;
+ if (x <= -a || x >= a) {
+ return 0.0f;
+ }
+ const float px = M_PI * x;
+ return (a * sinf(px) * sinf(px / a)) / (px * px);
+}
+
+__device__ float4 lanczos3_sample_axis(
+ const float4* src, int src_len, int src_stride, int src_base, int dst_idx, float scale)
+{
+ float center = (dst_idx + 0.5f) * scale - 0.5f;
+ int icenter = static_cast<int>(floorf(center));
+
+ float filter_scale = fmaxf(scale, 1.0f);
+ int radius = static_cast<int>(ceilf(3.0f * filter_scale));
+
+ float4 sum = make_float4(0.0f, 0.0f, 0.0f, 0.0f);
+ float wsum = 0.0f;
+
+ for (int i = icenter - radius; i <= icenter + radius; ++i) {
+ int clamped = min(max(i, 0), src_len - 1);
+ float w = lanczos3((i - center) / filter_scale);
+ const float4& s = src[src_base + clamped * src_stride];
+ sum.x += s.x * w;
+ sum.y += s.y * w;
+ sum.z += s.z * w;
+ sum.w += s.w * w;
+ wsum += w;
+ }
+
+ if (wsum > 1e-6f) {
+ sum.x /= wsum;
+ sum.y /= wsum;
+ sum.z /= wsum;
+ sum.w /= wsum;
+ }
+ return sum;
+}
+
+__global__ void downsample_horizontal(
+ const float4* src, int src_w, int h,
+ float4* dst, int dst_w, float scale_x)
+{
+ int x = blockIdx.x * blockDim.x + threadIdx.x;
+ int y = blockIdx.y * blockDim.y + threadIdx.y;
+ if (x >= dst_w || y >= h) {
+ return;
+ }
+
+ int src_row_base = y * src_w;
+ dst[y * dst_w + x] =
+ lanczos3_sample_axis(src, src_w, 1, src_row_base, x, scale_x);
+}
+
+__global__ void downsample_vertical(
+ const float4* src, int w, int src_h,
+ float4* dst, int dst_h, float scale_y)
+{
+ int x = blockIdx.x * blockDim.x + threadIdx.x;
+ int y = blockIdx.y * blockDim.y + threadIdx.y;
+ if (x >= w || y >= dst_h) return;
+
+ dst[y * w + x] =
+ lanczos3_sample_axis(src, src_h, w, x, y, scale_y);
+}
+
+void downsample_level(const PyramidLevel& src, PyramidLevel& dst) {
+ float scale_x = static_cast<float>(src.width) / dst.width;
+ float scale_y = static_cast<float>(src.height) / dst.height;
+
+ float4* tmp;
+ hip_check(hipMalloc(&tmp, sizeof(float4) * dst.width * src.height));
+
+ dim3 block(16, 16);
+
+ dim3 grid_h((dst.width + block.x - 1) / block.x,
+ (src.height + block.y - 1) / block.y);
+ downsample_horizontal<<<grid_h, block>>>(
+ src.color_smax, src.width, src.height, tmp, dst.width, scale_x);
+ hip_check(hipGetLastError());
+
+ dim3 grid_v((dst.width + block.x - 1) / block.x,
+ (dst.height + block.y - 1) / block.y);
+ downsample_vertical<<<grid_v, block>>>(
+ tmp, dst.width, src.height, dst.color_smax, dst.height, scale_y);
+ hip_check(hipGetLastError());
+
+ hipFree(tmp);
+}
+
+void build_pyramid(KuwaharaPyramid& pyr) {
+ for (size_t i = 1; i < pyr.levels.size(); ++i) {
+ downsample_level(pyr.levels[i - 1], pyr.levels[i]);
+ }
+}
+
+
+int main() { }