diff --git a/src/infiniop/ops/upsample_bilinear/cpu/upsample_bilinear_cpu.cc b/src/infiniop/ops/upsample_bilinear/cpu/upsample_bilinear_cpu.cc index 5f0ffe0ca..084c817ce 100644 --- a/src/infiniop/ops/upsample_bilinear/cpu/upsample_bilinear_cpu.cc +++ b/src/infiniop/ops/upsample_bilinear/cpu/upsample_bilinear_cpu.cc @@ -2,6 +2,7 @@ #include "../../../devices/cpu/common_cpu.h" #include #include +#include #include #include @@ -104,6 +105,14 @@ void calculate_cpu_impl( size_t out_h = info.h_out(); size_t out_w = info.w_out(); bool align_corners = info.align_corners(); + const ptrdiff_t input_stride_n = info.input_stride(0); + const ptrdiff_t input_stride_c = info.input_stride(1); + const ptrdiff_t input_stride_h = info.input_stride(2); + const ptrdiff_t input_stride_w = info.input_stride(3); + const ptrdiff_t output_stride_n = info.output_stride(0); + const ptrdiff_t output_stride_c = info.output_stride(1); + const ptrdiff_t output_stride_h = info.output_stride(2); + const ptrdiff_t output_stride_w = info.output_stride(3); auto out_ptr = reinterpret_cast(output); auto in_ptr = reinterpret_cast(input); @@ -116,24 +125,24 @@ void calculate_cpu_impl( #pragma omp parallel for schedule(static) for (ptrdiff_t nc = 0; nc < (ptrdiff_t)n_c; ++nc) { - // 当前 channel 的输入输出起始指针 - const T *src_base = in_ptr + nc * in_h * in_w; - T *dst_base = out_ptr + nc * out_h * out_w; + const size_t n = static_cast(nc) / C; + const size_t c = static_cast(nc) % C; + const ptrdiff_t input_base = n * input_stride_n + c * input_stride_c; + const ptrdiff_t output_base = n * output_stride_n + c * output_stride_c; for (size_t h = 0; h < out_h; ++h) { const auto &hp = h_params[h]; - // 缓存行指针,避免内层循环重复计算乘法 - const T *src_row0 = src_base + hp.idx0 * in_w; - const T *src_row1 = src_base + hp.idx1 * in_w; + const ptrdiff_t row0 = input_base + hp.idx0 * input_stride_h; + const ptrdiff_t row1 = input_base + hp.idx1 * input_stride_h; for (size_t w = 0; w < out_w; ++w) { const auto &wp = w_params[w]; // 获取四个采样点的值 - float val00 = utils::cast(src_row0[wp.idx0]); - float val01 = utils::cast(src_row0[wp.idx1]); - float val10 = utils::cast(src_row1[wp.idx0]); - float val11 = utils::cast(src_row1[wp.idx1]); + float val00 = utils::cast(in_ptr[row0 + wp.idx0 * input_stride_w]); + float val01 = utils::cast(in_ptr[row0 + wp.idx1 * input_stride_w]); + float val10 = utils::cast(in_ptr[row1 + wp.idx0 * input_stride_w]); + float val11 = utils::cast(in_ptr[row1 + wp.idx1 * input_stride_w]); // 双线性插值计算 // interpolation = (val00 * w0 + val01 * w1) * h_w0 + (val10 * w0 + val11 * w1) * h_w1 @@ -141,7 +150,7 @@ void calculate_cpu_impl( float val_h1 = val10 * wp.w0 + val11 * wp.w1; float result = val_h0 * hp.w0 + val_h1 * hp.w1; - dst_base[h * out_w + w] = utils::cast(result); + out_ptr[output_base + h * output_stride_h + w * output_stride_w] = utils::cast(result); } } } diff --git a/src/infiniop/ops/upsample_bilinear/cuda/kernel.cuh b/src/infiniop/ops/upsample_bilinear/cuda/kernel.cuh index 8ab3358b9..e9731d77e 100644 --- a/src/infiniop/ops/upsample_bilinear/cuda/kernel.cuh +++ b/src/infiniop/ops/upsample_bilinear/cuda/kernel.cuh @@ -2,6 +2,7 @@ #define __UPSAMPLE_BILINEAR_CUDA_CUH__ #include +#include #include namespace op::upsample_bilinear::cuda { @@ -39,6 +40,14 @@ __global__ void upsample_bilinear_kernel( size_t W_in, size_t H_out, size_t W_out, + ptrdiff_t input_stride_n, + ptrdiff_t input_stride_c, + ptrdiff_t input_stride_h, + ptrdiff_t input_stride_w, + ptrdiff_t output_stride_n, + ptrdiff_t output_stride_c, + ptrdiff_t output_stride_h, + ptrdiff_t output_stride_w, float scale_h, // 预计算的缩放比例 float scale_w, // 预计算的缩放比例 bool align_corners) { @@ -82,20 +91,19 @@ __global__ void upsample_bilinear_kernel( w0 = clamp(w0, 0, static_cast(W_in) - 1); w1 = clamp(w1, 0, static_cast(W_in) - 1); - // 6. 读取数据 - // 计算当前 Batch 和 Channel 的 Input 基地址 - const T *img_base = input + (n_idx * C + c_idx) * H_in * W_in; - - float val00 = static_cast(img_base[h0 * W_in + w0]); - float val01 = static_cast(img_base[h0 * W_in + w1]); - float val10 = static_cast(img_base[h1 * W_in + w0]); - float val11 = static_cast(img_base[h1 * W_in + w1]); + const ptrdiff_t input_base = n_idx * input_stride_n + c_idx * input_stride_c; + float val00 = static_cast(input[input_base + h0 * input_stride_h + w0 * input_stride_w]); + float val01 = static_cast(input[input_base + h0 * input_stride_h + w1 * input_stride_w]); + float val10 = static_cast(input[input_base + h1 * input_stride_h + w0 * input_stride_w]); + float val11 = static_cast(input[input_base + h1 * input_stride_h + w1 * input_stride_w]); // 7. 双线性插值计算 // result = (val00 * w0 + val01 * w1) * h0 + (val10 * w0 + val11 * w1) * h1 float val = h0_lambda * (w0_lambda * val00 + w1_lambda * val01) + h1_lambda * (w0_lambda * val10 + w1_lambda * val11); - output[i] = static_cast(val); + const ptrdiff_t output_offset = n_idx * output_stride_n + c_idx * output_stride_c + + h_out_idx * output_stride_h + w_out_idx * output_stride_w; + output[output_offset] = static_cast(val); } } diff --git a/src/infiniop/ops/upsample_bilinear/info.h b/src/infiniop/ops/upsample_bilinear/info.h index e1d3f9eb9..7e6854b61 100644 --- a/src/infiniop/ops/upsample_bilinear/info.h +++ b/src/infiniop/ops/upsample_bilinear/info.h @@ -3,6 +3,8 @@ #include "../../../utils.h" #include "../../tensor.h" +#include +#include #include namespace op::upsample_bilinear { @@ -23,6 +25,8 @@ class UpsampleBilinearInfo { size_t _w_in; // Input Width size_t _h_out; // Output Height size_t _w_out; // Output Width + std::array _input_strides; + std::array _output_strides; int dtype() const { return _dtype; } bool align_corners() const { return _align_corners; } @@ -32,16 +36,22 @@ class UpsampleBilinearInfo { size_t w_in() const { return _w_in; } size_t h_out() const { return _h_out; } size_t w_out() const { return _w_out; } + ptrdiff_t input_stride(size_t dim) const { return _input_strides[dim]; } + ptrdiff_t output_stride(size_t dim) const { return _output_strides[dim]; } // 构造函数 UpsampleBilinearInfo(int dtype, bool align_corners, size_t n, size_t c, size_t h_in, size_t w_in, - size_t h_out, size_t w_out) + size_t h_out, size_t w_out, + std::array input_strides, + std::array output_strides) : _dtype(dtype), _align_corners(align_corners), _n(n), _c(c), _h_in(h_in), _w_in(w_in), - _h_out(h_out), _w_out(w_out) {} + _h_out(h_out), _w_out(w_out), + _input_strides(input_strides), + _output_strides(output_strides) {} static utils::Result create( infiniopTensorDescriptor_t out_desc, @@ -49,10 +59,9 @@ class UpsampleBilinearInfo { int align_corners) { // C 接口通常传入 int 替代 bool // 1. 检查维度数量 - // 至少需要 2 维 (H, W) - // 修复: 使用 size_t 避免与 ndim() 返回值比较时的 signed/unsigned 警告 + // Normalize [H, W], [C, H, W], and [N, C, H, W] to NCHW. size_t ndim = input_desc->ndim(); - if (ndim < 2) { + if (ndim < 2 || ndim > 4) { return INFINI_STATUS_BAD_TENSOR_SHAPE; } if (out_desc->ndim() != ndim) { @@ -67,32 +76,21 @@ class UpsampleBilinearInfo { // 3. 检查 Batch/Channel 维度一致性 // 除了最后两维 (H, W),前面的维度必须完全匹配 - size_t n = 1; - size_t c = 1; - - // 解析 N 和 C 用于 Info 缓存 - // 逻辑: - // ndim = 4: [N, C, H, W] -> n=dims[0], c=dims[1] - // ndim = 3: [C, H, W] -> n=1, c=dims[0] - // ndim = 2: [H, W] -> n=1, c=1 - // 其他情况将所有非 spatial 维度累乘到 c 中 (视为 flattened channels) - - for (size_t i = 0; i < ndim - 2; ++i) { // 循环变量 i 也建议改为 size_t + for (size_t i = 0; i < ndim - 2; ++i) { if (input_desc->shape()[i] != out_desc->shape()[i]) { return INFINI_STATUS_BAD_TENSOR_SHAPE; } + } - // 简单 heuristic 来填充 n 和 c - if (ndim == 4 && i == 0) { - n = input_desc->shape()[i]; - } else if (ndim == 4 && i == 1) { - c = input_desc->shape()[i]; - } else if (ndim == 3 && i == 0) { - c = input_desc->shape()[i]; - } else { - // 对于 >4 维的情况,简单地归约为 c - c *= input_desc->shape()[i]; - } + size_t n = ndim == 4 ? input_desc->shape()[0] : 1; + size_t c = ndim == 4 ? input_desc->shape()[1] + : (ndim == 3 ? input_desc->shape()[0] : 1); + std::array input_strides{0, 0, 0, 0}; + std::array output_strides{0, 0, 0, 0}; + const size_t stride_offset = 4 - ndim; + for (size_t i = 0; i < ndim; ++i) { + input_strides[stride_offset + i] = input_desc->strides()[i]; + output_strides[stride_offset + i] = out_desc->strides()[i]; } // 4. 获取空间维度 @@ -114,7 +112,9 @@ class UpsampleBilinearInfo { h_in, w_in, h_out, - w_out}); + w_out, + input_strides, + output_strides}); } }; diff --git a/src/infiniop/ops/upsample_bilinear/metax/upsample_bilinear_metax.maca b/src/infiniop/ops/upsample_bilinear/metax/upsample_bilinear_metax.maca index ccafac3e2..0e182ad5b 100644 --- a/src/infiniop/ops/upsample_bilinear/metax/upsample_bilinear_metax.maca +++ b/src/infiniop/ops/upsample_bilinear/metax/upsample_bilinear_metax.maca @@ -5,6 +5,7 @@ #include "upsample_bilinear_metax.h" #include #include +#include #include namespace op::upsample_bilinear::metax { @@ -60,6 +61,14 @@ __global__ void upsample_bilinear_kernel( size_t W_in, size_t H_out, size_t W_out, + ptrdiff_t input_stride_n, + ptrdiff_t input_stride_c, + ptrdiff_t input_stride_h, + ptrdiff_t input_stride_w, + ptrdiff_t output_stride_n, + ptrdiff_t output_stride_c, + ptrdiff_t output_stride_h, + ptrdiff_t output_stride_w, float scale_h, // 预计算的缩放比例 float scale_w, // 预计算的缩放比例 bool align_corners) { @@ -101,18 +110,18 @@ __global__ void upsample_bilinear_kernel( w0 = clamp(w0, 0, static_cast(W_in) - 1); w1 = clamp(w1, 0, static_cast(W_in) - 1); - // 6. 读取数据并转换为 float - const T *img_base = input + (n_idx * C + c_idx) * H_in * W_in; - - float val00 = to_float(img_base[h0 * W_in + w0]); - float val01 = to_float(img_base[h0 * W_in + w1]); - float val10 = to_float(img_base[h1 * W_in + w0]); - float val11 = to_float(img_base[h1 * W_in + w1]); + const ptrdiff_t input_base = n_idx * input_stride_n + c_idx * input_stride_c; + float val00 = to_float(input[input_base + h0 * input_stride_h + w0 * input_stride_w]); + float val01 = to_float(input[input_base + h0 * input_stride_h + w1 * input_stride_w]); + float val10 = to_float(input[input_base + h1 * input_stride_h + w0 * input_stride_w]); + float val11 = to_float(input[input_base + h1 * input_stride_h + w1 * input_stride_w]); // 7. 双线性插值计算 float val = h0_lambda * (w0_lambda * val00 + w1_lambda * val01) + h1_lambda * (w0_lambda * val10 + w1_lambda * val11); - output[i] = static_cast(val); + const ptrdiff_t output_offset = n_idx * output_stride_n + c_idx * output_stride_c + + h_out_idx * output_stride_h + w_out_idx * output_stride_w; + output[output_offset] = static_cast(val); } } @@ -165,6 +174,10 @@ void launch_kernel( out_ptr, in_ptr, N, C, H_in, W_in, H_out, W_out, + info.input_stride(0), info.input_stride(1), + info.input_stride(2), info.input_stride(3), + info.output_stride(0), info.output_stride(1), + info.output_stride(2), info.output_stride(3), scale_h, scale_w, align_corners); } diff --git a/src/infiniop/ops/upsample_bilinear/moore/upsample_bilinear_moore.mu b/src/infiniop/ops/upsample_bilinear/moore/upsample_bilinear_moore.mu index 546765c66..3dd9ad1ad 100644 --- a/src/infiniop/ops/upsample_bilinear/moore/upsample_bilinear_moore.mu +++ b/src/infiniop/ops/upsample_bilinear/moore/upsample_bilinear_moore.mu @@ -56,6 +56,10 @@ void launch_kernel( out_ptr, in_ptr, N, C, H_in, W_in, H_out, W_out, + info.input_stride(0), info.input_stride(1), + info.input_stride(2), info.input_stride(3), + info.output_stride(0), info.output_stride(1), + info.output_stride(2), info.output_stride(3), scale_h, scale_w, align_corners); } diff --git a/src/infiniop/ops/upsample_bilinear/moore/upsample_bilinear_moore_kernel.h b/src/infiniop/ops/upsample_bilinear/moore/upsample_bilinear_moore_kernel.h index e2ef3b02f..b1105e413 100644 --- a/src/infiniop/ops/upsample_bilinear/moore/upsample_bilinear_moore_kernel.h +++ b/src/infiniop/ops/upsample_bilinear/moore/upsample_bilinear_moore_kernel.h @@ -2,6 +2,7 @@ #define __UPSAMPLE_BILINEAR_MOORE_H__ #include +#include #include #include #include @@ -33,6 +34,14 @@ __global__ void upsample_bilinear_kernel( size_t W_in, size_t H_out, size_t W_out, + ptrdiff_t input_stride_n, + ptrdiff_t input_stride_c, + ptrdiff_t input_stride_h, + ptrdiff_t input_stride_w, + ptrdiff_t output_stride_n, + ptrdiff_t output_stride_c, + ptrdiff_t output_stride_h, + ptrdiff_t output_stride_w, float scale_h, float scale_w, bool align_corners) { @@ -67,16 +76,17 @@ __global__ void upsample_bilinear_kernel( w0 = clamp(w0, 0, static_cast(W_in) - 1); w1 = clamp(w1, 0, static_cast(W_in) - 1); - const T *img_base = input + (n_idx * C + c_idx) * H_in * W_in; - - float val00 = static_cast(img_base[h0 * W_in + w0]); - float val01 = static_cast(img_base[h0 * W_in + w1]); - float val10 = static_cast(img_base[h1 * W_in + w0]); - float val11 = static_cast(img_base[h1 * W_in + w1]); + const ptrdiff_t input_base = n_idx * input_stride_n + c_idx * input_stride_c; + float val00 = static_cast(input[input_base + h0 * input_stride_h + w0 * input_stride_w]); + float val01 = static_cast(input[input_base + h0 * input_stride_h + w1 * input_stride_w]); + float val10 = static_cast(input[input_base + h1 * input_stride_h + w0 * input_stride_w]); + float val11 = static_cast(input[input_base + h1 * input_stride_h + w1 * input_stride_w]); float val = h0_lambda * (w0_lambda * val00 + w1_lambda * val01) + h1_lambda * (w0_lambda * val10 + w1_lambda * val11); - output[i] = static_cast(val); + const ptrdiff_t output_offset = n_idx * output_stride_n + c_idx * output_stride_c + + h_out_idx * output_stride_h + w_out_idx * output_stride_w; + output[output_offset] = static_cast(val); } } diff --git a/src/infiniop/ops/upsample_bilinear/nvidia/upsample_bilinear_nvidia.cu b/src/infiniop/ops/upsample_bilinear/nvidia/upsample_bilinear_nvidia.cu index d7ee074f4..305101e97 100644 --- a/src/infiniop/ops/upsample_bilinear/nvidia/upsample_bilinear_nvidia.cu +++ b/src/infiniop/ops/upsample_bilinear/nvidia/upsample_bilinear_nvidia.cu @@ -70,6 +70,10 @@ void launch_kernel( out_ptr, in_ptr, N, C, H_in, W_in, H_out, W_out, + info.input_stride(0), info.input_stride(1), + info.input_stride(2), info.input_stride(3), + info.output_stride(0), info.output_stride(1), + info.output_stride(2), info.output_stride(3), scale_h, scale_w, align_corners); } diff --git a/test/infinicore/ops/upsample_bilinear.py b/test/infinicore/ops/upsample_bilinear.py index b6871c8a7..319aff946 100644 --- a/test/infinicore/ops/upsample_bilinear.py +++ b/test/infinicore/ops/upsample_bilinear.py @@ -23,6 +23,8 @@ ((2, 3, 6, 6), (12, 12), None, None), ((4, 3, 7, 7), 2.0, False, None), ((3, 3, 5, 5), (10, 10), True, None), + # Channel-first view of a contiguous NHWC tensor, as used by Qwen3-VL. + ((1, 8, 4, 4), (6, 8), True, (128, 1, 32, 8)), ] _TOLERANCE_MAP = {