aboutsummaryrefslogtreecommitdiff
path: root/src/kernels
diff options
context:
space:
mode:
authorYousef Jan <yousefjan24000@gmail.com>2026-08-25 15:47:09 +0300
committerYousef Jan <yousefjan24000@gmail.com>2026-08-25 15:51:41 +0300
commit15bb2f12fd1b3a47ae0a790cfbd71f8837a3badc (patch)
tree5aedd8fce4a9cb411ec9e8148fc76eaaa57c7589 /src/kernels
parentce90395aa32ad37d5ebcb628b2b22d7bf7d9850b (diff)
Modularize and cleanHEADmain
Diffstat (limited to 'src/kernels')
-rw-r--r--src/kernels/pyramid.hip122
1 files changed, 122 insertions, 0 deletions
diff --git a/src/kernels/pyramid.hip b/src/kernels/pyramid.hip
new file mode 100644
index 0000000..80d9838
--- /dev/null
+++ b/src/kernels/pyramid.hip
@@ -0,0 +1,122 @@
+#include "pipeline.h"
+#include <algorithm>
+
+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;
+}
+
+Pyramid make_pyramid(int base_width, int base_height, int num_levels) {
+ Pyramid 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(Pyramid& pyr) {
+ for (auto& lvl : pyr.levels) free_level(lvl);
+ pyr.levels.clear();
+}
+
+/* Lanczos3 */
+
+__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 fscale = fmaxf(scale, 1.0f);
+ int radius = static_cast<int>(ceilf(3.0f * fscale));
+
+ 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 c = min(max(i, 0), src_len - 1);
+ float w = lanczos3((i - center) / fscale);
+ float4 s = src[src_base + c * 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) {
+ float inv = 1.0f / wsum;
+ sum.x *= inv; sum.y *= inv; sum.z *= inv; sum.w *= inv;
+ }
+ 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;
+
+ dst[y * dst_w + x] =
+ lanczos3_sample_axis(src, src_w, 1, y * src_w, 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);
+}
+
+static 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 + 15) / 16, (src.height + 15) / 16);
+ 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 + 15) / 16, (dst.height + 15) / 16);
+ 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(Pyramid& pyr) {
+ for (size_t i = 1; i < pyr.levels.size(); ++i)
+ downsample_level(pyr.levels[i - 1], pyr.levels[i]);
+}