From 2287c5123653274b9a0833569cafe5746b8c009a Mon Sep 17 00:00:00 2001 From: Thomas Benson Date: Sun, 20 Sep 2026 10:50:42 -0400 Subject: [PATCH] Optimize channelize_poly CUDA kernels and dispatch Add fused FIR+DFT kernels that keep filtered samples on chip, avoiding the intermediate tensor and separate cuFFT launch: small-channel leaves for critically sampled M=2..6 (replacing FusedChan) and general leaves for M=8, 10, 16, 20, 32, 40, 64, and 80, plus oversampled M=3..6. Rework the FIR+cuFFT backends: tuning of tile sizes and launch configuration. Select backends through SelectPlan/ExecutePlan using device attributes (SM count, L2 size, FP64 throughput, memory bus width). Tests compare every launchable backend, including windowed launches, against the host implementation. channelize_poly_bench gains single-case options. Across a 4,832-shape sweep, the geometric-mean speedup over the previous implementation is 1.89x on L4 and 2.15x on GH200. Signed-off-by: Thomas Benson --- .../signalimage/filtering/channelize_poly.rst | 19 +- docs_input/executor_compatibility.rst | 2 +- examples/channelize_poly_bench.cu | 75 +- include/matx/kernels/channelize_poly.cuh | 1743 +++++++++++++---- include/matx/transforms/channelize_poly.h | 1418 +++++++++----- test/00_transform/ChannelizePoly.cu | 1315 ++++++++++++- test/00_transform/StreamingChannelize.cu | 2 +- test/test_vectors/generators/00_transforms.py | 31 +- 8 files changed, 3670 insertions(+), 935 deletions(-) diff --git a/docs_input/api/signalimage/filtering/channelize_poly.rst b/docs_input/api/signalimage/filtering/channelize_poly.rst index bb20a167a..225d04d82 100644 --- a/docs_input/api/signalimage/filtering/channelize_poly.rst +++ b/docs_input/api/signalimage/filtering/channelize_poly.rst @@ -9,6 +9,24 @@ Polyphase channelizer with a configurable number of channels .. doxygenfunction:: matx::channelize_poly(const InType &in, const FilterType &f, index_t num_channels, index_t decimation_factor) +CUDA performance +~~~~~~~~~~~~~~~~ + +For some CUDA inputs and channelizer configurations, MatX uses a fused kernel +that performs the polyphase filtering and FFT in one launch. Other +configurations use the general backend, which performs the filtering and FFT +separately. Kernel selection is automatic. + +Limitations +~~~~~~~~~~~ + +The general backend uses cuFFT, which supports half-precision (``matxFp16`` and +``matxBf16``) transforms only for power-of-two sizes. On CUDA, half-precision +outputs with any other channel count are supported only for critically sampled +channelizers (``decimation_factor == num_channels``) with 3, 5, or 6 channels, +which always use the fused kernel. Other such configurations raise a +``matxInvalidParameter`` error. + Examples ~~~~~~~~ @@ -23,4 +41,3 @@ Examples :start-after: example-begin channelize_poly-test-2 :end-before: example-end channelize_poly-test-2 :dedent: - diff --git a/docs_input/executor_compatibility.rst b/docs_input/executor_compatibility.rst index 467c0f7f6..45a42eaa7 100644 --- a/docs_input/executor_compatibility.rst +++ b/docs_input/executor_compatibility.rst @@ -78,7 +78,7 @@ existing operators do not implicitly become distributed operations. "cart2sph", "|yes|", "|yes|", "|yes|", "|no|", "Element-wise coordinate conversion expression." "ceil", "|yes|", "|yes|", "|yes|", "|no|", "Element-wise expression." "cgsolve", "|no|", "|yes|", "|no|", "|no|", "CUDA iterative solver path." - "channelize_poly", "|yes|", "|yes|", "|no|", "|no|", "Polyphase channelizer; host path directly computes the per-branch FIR and DFT stages." + "channelize_poly", "|yes|", "|yes|", "|no|", "|no|", "Polyphase channelizer; host path directly computes the per-branch FIR and DFT stages. CUDA has limitations for half-precision outputs; see :ref:`channelize_poly_func`." "chirp", "|yes|", "|yes|", "|yes|", "|no|", "Generator expression." "chol", "|yes|", "|yes|", "|yes|", "|partial|", "Host support requires the CPU solver backend. CUDAJITExecutor support uses cuSolverDx through MathDx for supported rank 2-4 square float, double, complex-float, and complex-double matrices. Experimental distributedCUDAExecutor support is limited to aligned batch sharding with fully local matrix dimensions." "clone", "|yes|", "|yes|", "|yes|", "|no|", "View expression." diff --git a/examples/channelize_poly_bench.cu b/examples/channelize_poly_bench.cu index 90250fe9e..a72ec782f 100644 --- a/examples/channelize_poly_bench.cu +++ b/examples/channelize_poly_bench.cu @@ -38,14 +38,13 @@ #include #include #include +#include #include using namespace matx; -// This example is used primarily for development purposes to benchmark the performance of the -// polyphase channelizer kernel(s). Typically, the parameters below (batch size, filter -// length, input signal length, and channel range) will be adjusted to a range of interest -// and the benchmark will be run with and without the proposed kernel changes. +// This example is used primarily for development purposes to benchmark the +// performance of the polyphase channelizer kernels. constexpr int NUM_WARMUP_ITERATIONS = 2; @@ -62,13 +61,16 @@ const char *TypeName() { } template -void ChannelizePolyBench(matx::index_t num_channels, matx::index_t decimation_factor) +void ChannelizePolyBench(matx::index_t num_channels, matx::index_t decimation_factor, + matx::index_t custom_batches, matx::index_t custom_filter_len_per_channel, + matx::index_t custom_input_len, int warmup_iterations, int iterations) { - struct { + struct TestCase { matx::index_t num_batches; matx::index_t filter_len_per_channel; matx::index_t input_len; - } test_cases[] = { + }; + std::vector test_cases = { { 1, 17, 256 }, { 1, 17, 3000 }, { 1, 17, 31000 }, @@ -78,6 +80,9 @@ void ChannelizePolyBench(matx::index_t num_channels, matx::index_t decimation_fa { 1, 17, 8192*1024 }, { 42, 17, 8192*1024 } }; + if (custom_input_len > 0) { + test_cases = {{custom_batches, custom_filter_len_per_channel, custom_input_len}}; + } cudaStream_t stream; cudaStreamCreate(&stream); @@ -85,15 +90,15 @@ void ChannelizePolyBench(matx::index_t num_channels, matx::index_t decimation_fa cudaEventCreate(&start); cudaEventCreate(&stop); - cudaExecutor exec{}; + cudaExecutor exec{stream}; - for (size_t i = 0; i < sizeof(test_cases)/sizeof(test_cases[0]); i++) { + for (size_t i = 0; i < test_cases.size(); i++) { const matx::index_t num_batches = test_cases[i].num_batches; const matx::index_t filter_len = test_cases[i].filter_len_per_channel * num_channels; const matx::index_t input_len = test_cases[i].input_len; const matx::index_t output_len_per_channel = (input_len + decimation_factor - 1) / decimation_factor; - if (input_len < num_channels * 100) { + if (custom_input_len <= 0 && input_len < num_channels * 100) { continue; } @@ -103,7 +108,7 @@ void ChannelizePolyBench(matx::index_t num_channels, matx::index_t decimation_fa (input = static_cast(1)).run(exec); (filter = static_cast(1)).run(exec); - for (int k = 0; k < NUM_WARMUP_ITERATIONS; k++) { + for (int k = 0; k < warmup_iterations; k++) { (output = channelize_poly(input, filter, num_channels, decimation_factor)).run(exec); } @@ -111,7 +116,7 @@ void ChannelizePolyBench(matx::index_t num_channels, matx::index_t decimation_fa float elapsed_ms = 0.0f; cudaEventRecord(start, stream); - for (int k = 0; k < NUM_ITERATIONS; k++) { + for (int k = 0; k < iterations; k++) { (output = channelize_poly(input, filter, num_channels, decimation_factor)).run(exec); } cudaEventRecord(stop, stream); @@ -119,7 +124,7 @@ void ChannelizePolyBench(matx::index_t num_channels, matx::index_t decimation_fa MATX_CUDA_CHECK_LAST_ERROR(); cudaEventElapsedTime(&elapsed_ms, start, stop); - const double avg_elapsed_us = (static_cast(elapsed_ms)/NUM_ITERATIONS)*1.0e3; + const double avg_elapsed_us = (static_cast(elapsed_ms)/iterations)*1.0e3; printf("Batches: %5" MATX_INDEX_T_FMT " Channels: %5" MATX_INDEX_T_FMT " Decimation: %5" MATX_INDEX_T_FMT " FilterLen: %5" MATX_INDEX_T_FMT " InputLen: %7" MATX_INDEX_T_FMT " Elapsed Usecs: %12.1f MPts/sec: %12.3f\n", num_batches, num_channels, decimation_factor, filter_len, input_len, avg_elapsed_us, @@ -143,6 +148,11 @@ struct BenchConfig { Domain filter_domain = Domain::Real; matx::index_t M = 10; // number of channels matx::index_t D = -1; // decimation factor (-1 means D = M) + matx::index_t batches = 1; + matx::index_t filter_len_per_channel = 17; + matx::index_t input_len = -1; + int warmup_iterations = NUM_WARMUP_ITERATIONS; + int iterations = NUM_ITERATIONS; }; void PrintUsage(const char *prog) { @@ -151,7 +161,13 @@ void PrintUsage(const char *prog) { printf(" --filter-type Filter type: float, double, cf, cd (default: float)\n"); printf(" -M Number of channels (default: 10)\n"); printf(" -D Decimation factor, 0 < D <= M (default: M)\n"); + printf(" --batches Batch count for a single custom case (default: 1)\n"); + printf(" --filter-per-channel Filter taps per channel for a custom case (default: 17)\n"); + printf(" --input-len Run one custom case with input length N > 0\n"); + printf(" --warmups Warmup iterations (default: %d)\n", NUM_WARMUP_ITERATIONS); + printf(" --iterations Timed iterations (default: %d)\n", NUM_ITERATIONS); printf("\n"); + printf("--batches and --filter-per-channel require --input-len.\n"); printf("Type shorthands: float, double, cf (complex), cd (complex)\n"); } @@ -179,7 +195,9 @@ void DispatchBench(const BenchConfig &cfg) { TypeName(), TypeName(), TypeName()); printf("M: %" MATX_INDEX_T_FMT " D: %" MATX_INDEX_T_FMT "\n\n", cfg.M, cfg.D); - ChannelizePolyBench(cfg.M, cfg.D); + ChannelizePolyBench( + cfg.M, cfg.D, cfg.batches, cfg.filter_len_per_channel, cfg.input_len, + cfg.warmup_iterations, cfg.iterations); } void RunBench(const BenchConfig &cfg) { @@ -213,6 +231,7 @@ int main(int argc, char **argv) MATX_ENTER_HANDLER(); BenchConfig cfg; + bool requires_input_len = false; for (int i = 1; i < argc; i++) { if (strcmp(argv[i], "--help") == 0 || strcmp(argv[i], "-h") == 0) { @@ -232,6 +251,22 @@ int main(int argc, char **argv) cfg.M = static_cast(atol(argv[++i])); } else if (strcmp(argv[i], "-D") == 0 && i + 1 < argc) { cfg.D = static_cast(atol(argv[++i])); + } else if (strcmp(argv[i], "--batches") == 0 && i + 1 < argc) { + cfg.batches = static_cast(atol(argv[++i])); + requires_input_len = true; + } else if (strcmp(argv[i], "--filter-per-channel") == 0 && i + 1 < argc) { + cfg.filter_len_per_channel = static_cast(atol(argv[++i])); + requires_input_len = true; + } else if (strcmp(argv[i], "--input-len") == 0 && i + 1 < argc) { + cfg.input_len = static_cast(atol(argv[++i])); + if (cfg.input_len <= 0) { + fprintf(stderr, "Error: --input-len must be positive\n"); + return 1; + } + } else if (strcmp(argv[i], "--warmups") == 0 && i + 1 < argc) { + cfg.warmup_iterations = atoi(argv[++i]); + } else if (strcmp(argv[i], "--iterations") == 0 && i + 1 < argc) { + cfg.iterations = atoi(argv[++i]); } else { fprintf(stderr, "Unknown option: %s\n", argv[i]); PrintUsage(argv[0]); @@ -255,6 +290,18 @@ int main(int argc, char **argv) return 1; } + if (requires_input_len && cfg.input_len < 0) { + fprintf(stderr, "Error: --batches and --filter-per-channel require --input-len\n"); + return 1; + } + + if (cfg.batches <= 0 || cfg.filter_len_per_channel <= 0 || + cfg.warmup_iterations < 0 || cfg.iterations <= 0) { + fprintf(stderr, + "Error: custom dimensions and iteration counts must be positive (warmups may be zero)\n"); + return 1; + } + RunBench(cfg); matx::ClearCachesAndAllocations(); diff --git a/include/matx/kernels/channelize_poly.cuh b/include/matx/kernels/channelize_poly.cuh index 02ecdada3..f477ea494 100644 --- a/include/matx/kernels/channelize_poly.cuh +++ b/include/matx/kernels/channelize_poly.cuh @@ -45,50 +45,291 @@ #include "matx/kernels/tensor_accessor.h" #include #include +#include namespace matx { +namespace detail { + +// Packed even/odd values used only by the specialized M=2 leaves. +template +struct alignas(sizeof(T) < 16 ? 2 * sizeof(T) : 16) ChannelizePolyM2Pair { + T even; + T odd; +}; // detail constants that require both host and device visibility at compile // time. Scoped to matx::detail::cpoly so the transform's helpers can sit in // the same namespace without colliding with other transforms' internals. -namespace detail { namespace cpoly { // Number of output elements generated per thread constexpr index_t ElemsPerThread = 1; + // Largest dynamic shared memory allocation a kernel can launch with without + // opting in through cudaFuncSetAttribute. + constexpr size_t MaxDefaultDynamicSmemBytes = 48 * 1024; + + // Number of independent time lanes in a maximally decimated tiled CTA. Each lane computes + // several output rows so its channel's filter tap can be reused across those rows. + constexpr int SmemTiledMaxDecYThreads = 4; + // Maximum number of filter rotations per channel for the SmemTiled kernel. // This is used to determine if the filter can be stored in shared memory. // The number of rotations can exceed this value, but the filter will be // read from global memory rather than cached in shared memory. constexpr int SmemTiledMaxRotations = 32; -} // namespace cpoly -} // namespace detail -#ifdef __CUDACC__ + __MATX_HOST__ __MATX_DEVICE__ constexpr index_t SmemTiledInputHeight( + index_t taps, index_t channels, index_t decimation, int nout) + { + // Near-critical oversampling can need one older branch row while + // the next output group has already loaded a newer row into the ring. + const bool extra_row = decimation < channels && + static_cast(nout - 1) * decimation > + static_cast(nout - 2) * channels; + return taps + nout - 1 + extra_row; + } -namespace detail { + constexpr int FusedRadixThreads = 256; -template -__MATX_DEVICE__ __MATX_INLINE__ auto channelize_cast_filter(FilterT v) -{ - if constexpr (is_complex_v) { - // Complex filter: keep full complex multiply - return static_cast(v); - } else if constexpr (is_complex_v) { - // Real filter + complex accumulator: promote to scalar only - using accum_scalar_t = typename inner_op_type_t::type; - return static_cast(v); - } else { - return static_cast(v); + // Zero-padded load of input(base + offset). Requires offset >= 0. Unsigned + // arithmetic wraps modulo 2^N without overflow; for any representable base, + // a negative base + offset maps to at least 2^(N-1) > input_len, and a sum + // cannot wrap past 2^N, so one compare checks both bounds exactly. + template + __MATX_HOST__ __MATX_DEVICE__ __MATX_INLINE__ auto LoadRelativeInput( + const Input &input, IdxT base, IdxT input_len, int32_t offset) + { + using input_t = cuda::std::remove_cvref_t; + using unsigned_index_t = cuda::std::make_unsigned_t; + const auto index = static_cast(base) + + static_cast(offset); + return index < static_cast(input_len) + ? input(static_cast(index)) : input_t{}; + } + + // Small channel counts are most efficient when a thread owns a complete output row. Two rows + // per thread provide useful FIR ILP for narrow FP32 leaves. One row bounds register and + // shared-memory use for M>=4 and for FP64 leaves above M=2. + template + __MATX_HOST__ __MATX_DEVICE__ constexpr int FusedSmallOutputsPerThread() + { + static_assert(NUM_CHAN >= 2 && NUM_CHAN <= 6); + if constexpr (NUM_CHAN >= 4 || (cuda::std::is_same_v && NUM_CHAN >= 3)) { + return 1; + } else { + return 2; + } } -} -template -__MATX_DEVICE__ __MATX_INLINE__ auto channelize_cast_input(InputT v) + // Staging a complete overlapping tile is not amortized by very short filters. Keep that case + // inside the same kernel instantiation, but read the few FIR values directly before applying + // the fixed butterfly. Real input can use the direct leaf a little longer because its + // global-memory footprint is half that of complex input. + template + __MATX_HOST__ __MATX_DEVICE__ constexpr int FusedSmallDirectMaxTaps() + { + if constexpr (cuda::std::is_same_v && NUM_CHAN >= 5) { + return 16; + } else { + return 4; + } + } + + __MATX_HOST__ __MATX_DEVICE__ constexpr int FusedRadixOddFactor(int n) + { + while ((n & 1) == 0) n /= 2; + return n; + } + + __MATX_HOST__ __MATX_DEVICE__ constexpr int FusedRadixRowStride(int n) + { + // Pack oversampled M=3..6 rows into subwarps. Critical small-M leaves + // use their separate row-owned mapping. + int stride = n >= 3 && n <= 6 ? 1 : 32; + while (stride < n) stride *= 2; + return stride; + } + + // Reverse the lowest log2(N) bits for a power-of-two FFT size N. + template + __MATX_HOST__ __MATX_DEVICE__ __MATX_INLINE__ int32_t BitReverse(int32_t index) + { + static_assert(N > 0 && (N & (N - 1)) == 0); + int32_t reversed = 0; + uint32_t remaining = static_cast(index); + // BREV/CLZ could replace this loop, but measured slightly slower. + MATX_LOOP_UNROLL + for (int bit = 1; bit < N; bit *= 2) { + reversed = 2 * reversed + static_cast(remaining & 1); + remaining /= 2; + } + return reversed; + } + + template + struct FusedRadixConfig { + static_assert(NUM_CHAN >= 2); + static constexpr int Radix1 = FusedRadixOddFactor(NUM_CHAN); + static constexpr int Radix2 = NUM_CHAN / Radix1; + static constexpr bool PackedPow2 = Radix1 == 1 && NUM_CHAN >= 8 && NUM_CHAN < 32; + // Fill subwarps while retaining the existing tile height and footprint. + static constexpr int PackedRows = cuda::std::is_same_v ? 16 : 32; + // Smaller row-owned CTAs improve coverage. Real M=5/6 keeps 128 threads + // to amortize FIR staging without sacrificing short-signal coverage. + static constexpr int SmallThreads = + (NUM_CHAN == 2 && cuda::std::is_same_v) || + (NUM_CHAN >= 3 && NUM_CHAN <= 4) || + (is_complex_v && (NUM_CHAN >= 3 || sizeof(InputType) > 8)) + ? 64 : NUM_CHAN >= 5 ? 128 : FusedRadixThreads; + static constexpr int Threads = + MaximallyDecimated && NUM_CHAN <= 6 ? SmallThreads : + PackedPow2 ? cuda::std::min(FusedRadixThreads, PackedRows * NUM_CHAN) : + FusedRadixThreads; + static constexpr int RowStride = PackedPow2 ? NUM_CHAN : FusedRadixRowStride(NUM_CHAN); + // Packed FIR rows retain their low-overhead tiny DFT. + static constexpr bool WarpFft = (RowStride >= 32 || PackedPow2) && Radix2 <= 32; + // Large register butterflies need enough resident warps to hide latency. + static constexpr int MinBlocksPerSm = + WarpFft && 2 * Radix1 * sizeof(AccumType) >= 64 ? 1024 / Threads : 0; + static constexpr int FirGroups = Threads / RowStride; + // One output per thread bounds the packed small-M tile's footprint. + // Wide double FIR tiles retain smaller groups to leave room for taps. + static constexpr int FirOutputsPerThread = + NUM_CHAN >= 3 && NUM_CHAN <= 6 ? 1 : + PackedPow2 ? PackedRows * NUM_CHAN / Threads : + cuda::std::is_same_v + ? (!WarpFft ? 1 : Radix1 == 1 ? 2 : cuda::std::min(FirGroups, 4)) : 4; + static constexpr int NRows = NUM_CHAN == 2 + ? 1 : FirGroups * FirOutputsPerThread; + static constexpr int SmallOutputsPerThread = [] { + if constexpr (NUM_CHAN <= 6) { + return FusedSmallOutputsPerThread(); + } else { + return 0; + } + }(); + }; + + // Shared by launch sizing and device access. Element counts use native + // input/filter types, including when M=2 accesses them as packed pairs. + template + struct FusedRadixSmemLayout { + index_t filter_elements; + index_t input_elements; + size_t input_offset; + size_t work_offset = 0; + size_t stage1_offset = 0; + size_t twiddle_cross_offset = 0; + size_t twiddle_radix2_offset = 0; + size_t bytes; + + // Preserve the caller's index width: device tiles use 32-bit counts, + // while host eligibility checks may size filters too large to launch. + template + __MATX_HOST__ __MATX_DEVICE__ constexpr FusedRadixSmemLayout( + IdxT taps_per_channel, index_t decimation_factor) + { + using config = FusedRadixConfig; + using complex_t = typename scalar_to_complex::ctype; + constexpr size_t input_alignment = + NUM_CHAN == 2 && MaximallyDecimated + ? alignof(ChannelizePolyM2Pair) + : alignof(InputType); + filter_elements = taps_per_channel * NUM_CHAN; + input_offset = MATX_ROUND_UP( + static_cast(taps_per_channel) * NUM_CHAN * + sizeof(FilterType), + input_alignment); + if constexpr (NUM_CHAN <= 6 && MaximallyDecimated) { + constexpr int block_rows = config::Threads * config::SmallOutputsPerThread; + input_elements = (taps_per_channel + block_rows - 1) * NUM_CHAN; + } else if constexpr (NUM_CHAN == 2) { + input_elements = 2 * taps_per_channel + + config::Threads * config::SmallOutputsPerThread; + } else if constexpr (MaximallyDecimated) { + input_elements = (taps_per_channel + config::NRows - 1) * NUM_CHAN; + } else { + input_elements = taps_per_channel * NUM_CHAN + + (config::NRows - 1) * static_cast(decimation_factor) + 1; + } + bytes = input_offset + static_cast(input_elements) * sizeof(InputType); + if constexpr (NUM_CHAN > 2 && !(NUM_CHAN <= 6 && MaximallyDecimated)) { + constexpr size_t work_bytes = config::NRows * NUM_CHAN * sizeof(complex_t); + work_offset = MATX_ROUND_UP(bytes, alignof(complex_t)); + bytes = work_offset + work_bytes; + if constexpr (config::Radix1 > 1) { + stage1_offset = bytes; + twiddle_cross_offset = stage1_offset + (config::WarpFft ? 0 : work_bytes); + bytes = twiddle_cross_offset + NUM_CHAN * sizeof(complex_t); + } + twiddle_radix2_offset = bytes; + bytes = twiddle_radix2_offset + + (config::Radix2 + config::Radix2 / 2) * sizeof(complex_t); + } + } + }; + + // Whether the cached small-channel leaf's FIR tile fits the default + // shared-memory limit. + template + __MATX_HOST__ __MATX_DEVICE__ constexpr bool FusedSmallCachedFits(index_t taps_per_channel) + { + const FusedRadixSmemLayout layout(taps_per_channel, NUM_CHAN); + return layout.bytes <= MaxDefaultDynamicSmemBytes; + } + + // Keep host allocation and device leaf selection in sync. The direct leaf is + // faster for short filters and is the fallback when the cached tile does not fit. + template + __MATX_HOST__ __MATX_DEVICE__ constexpr bool FusedSmallUseDirect( + index_t taps_per_channel, index_t output_rows) + { + using input_t = typename InType::value_type; + if constexpr (is_tensor_view_v) { + bool use_direct = taps_per_channel <= + FusedSmallDirectMaxTaps(); + if constexpr (!cuda::std::is_same_v && NUM_CHAN >= 5) { + // For very long outputs, caching saves enough global reads + // to outweigh the direct leaf's lower setup cost. + use_direct &= output_rows <= (index_t{1} << 20); + } + if (use_direct) return true; + } + return !FusedSmallCachedFits(taps_per_channel); + } + + // Build and batch-bind the FIR accessors in their common setup order. + template + __MATX_HOST__ __MATX_DEVICE__ __MATX_INLINE__ auto MakeAccessors( + const OutType &output, const InType &input, const FilterType &filter, index_t batch) + { + TensorAccessor input_acc(input); + TensorAccessor output_acc(output); + TensorAccessor filter_acc(filter); + const auto in_batch_idx = BlockToIdx(input, batch, 1); + const auto out_batch_idx = BlockToIdx(output, batch, 2); + auto input_b = bind_first_n(input_acc, in_batch_idx); + auto output_b = bind_first_n(output_acc, out_batch_idx); + return cuda::std::make_tuple(input_b, output_b, filter_acc); + } +} // namespace cpoly + +template +__MATX_HOST__ __MATX_DEVICE__ __MATX_INLINE__ auto channelize_cast_operand(ValueT v) { - if constexpr (is_complex_v) { - return static_cast(v); + if constexpr (cuda::std::is_same_v) { + return v; + } else if constexpr (is_complex_v) { + // Component-wise conversion also supports mixed complex-half types. + using scalar_t = typename inner_op_type_t::type; + return AccumT{static_cast(v.real()), static_cast(v.imag())}; } else if constexpr (is_complex_v) { + // Preserve a real operand so channelize_cmac can use its cheaper + // real-by-complex specialization. using accum_scalar_t = typename inner_op_type_t::type; return static_cast(v); } else { @@ -100,7 +341,7 @@ __MATX_DEVICE__ __MATX_INLINE__ auto channelize_cast_input(InputT v) // ~8 mixed FMUL/FADD/FSUB. Falls back to the default operator* + operator+= // for real or mixed-precision types. template -__MATX_DEVICE__ __MATX_INLINE__ void channelize_cmac( +__MATX_HOST__ __MATX_DEVICE__ __MATX_INLINE__ void channelize_cmac( AccumT &accum, FilterValT hv, InputValT iv) { if constexpr (is_complex_v && is_complex_v && is_complex_v) { @@ -131,6 +372,8 @@ __MATX_DEVICE__ __MATX_INLINE__ void channelize_cmac( } // namespace detail +#ifdef __CUDACC__ + // out_elem_offset shifts the global per-channel output element (time) index // used for the input footprint and polyphase phase, while the write row stays // local, so this kernel can emit an arbitrary window @@ -139,7 +382,9 @@ __MATX_DEVICE__ __MATX_INLINE__ void channelize_cmac( // can be non-zero for a streaming channelizer call. template __launch_bounds__(THREADS) -__global__ void ChannelizePoly1D(OutType output, InType input, FilterType filter, index_t decimation_factor, uint32_t smem_filter_bytes, index_t out_elem_offset) +__global__ void ChannelizePoly1D( + OutType output, InType input, FilterType filter, index_t decimation_factor, + uint32_t smem_filter_bytes, index_t out_elem_offset, int elem_block_offset) { using output_t = typename OutType::value_type; using input_t = typename InType::value_type; @@ -160,8 +405,11 @@ __global__ void ChannelizePoly1D(OutType output, InType input, FilterType filter const index_t filter_full_len = filter.Size(0); const index_t filter_phase_len = (filter_full_len + num_channels - 1) / num_channels; - const int elem_block = blockIdx.x; - const int channel = blockIdx.y; + // Channels vary fastest across CTAs so the CTAs that read the same input + // span run together and share it through L2. Time blocks use grid.y; the + // launch splits grids taller than the grid.y limit using elem_block_offset. + const int channel = static_cast(blockIdx.x); + const int elem_block = static_cast(blockIdx.y) + elem_block_offset; const int tid = threadIdx.x; constexpr index_t ELEMS_PER_BLOCK = detail::cpoly::ElemsPerThread * THREADS; @@ -169,25 +417,9 @@ __global__ void ChannelizePoly1D(OutType output, InType input, FilterType filter const index_t last_out_elem = cuda::std::min( output_len_per_channel - 1, first_out_elem + ELEMS_PER_BLOCK - 1); - // Wrap input/output/filter in TensorAccessor and bind the per-block batch - // coords once. After binding, per-access calls supply only the inner - // indices: (sample_idx) for input, (t, channel) for output. On the fast - // path this collapses to base_ptr[stride*inner + ...] arithmetic with no - // per-access stride reload; on the slow path it forwards to operator(). - // - // Note on output layout: MatX arranges output as [batch..., elem, channel] - // where channel is the LAST dim (see the size asserts in - // channelize_poly_impl). We therefore bind only the batch dims (first - // OutRank-2) and pass both elem (t) and channel at access time. - detail::TensorAccessor input_acc(input); - detail::TensorAccessor output_acc(output); - detail::TensorAccessor filter_acc(filter); - - const auto in_batch_idx = BlockToIdx(input, blockIdx.z, 1); // last slot unused - const auto out_batch_idx = BlockToIdx(output, blockIdx.z, 2); // last two slots unused - - auto input_b = detail::bind_first_n(input_acc, in_batch_idx); - auto output_b = detail::bind_first_n(output_acc, out_batch_idx); + // Bind batch dimensions; output accesses retain (time, channel) indices. + auto [input_b, output_b, filter_acc] = + detail::cpoly::MakeAccessors(output, input, filter, blockIdx.z); if constexpr (MaximallyDecimated) { // Maximally decimated (D == M) path: filter phase is fixed per channel @@ -230,8 +462,8 @@ __global__ void ChannelizePoly1D(OutType output, InType input, FilterType filter for (int i = 0; i < niter; i++) { const input_t in_val = input_b(sample_idx); detail::channelize_cmac(accum, - detail::channelize_cast_filter(*h), - detail::channelize_cast_input(in_val)); + detail::channelize_cast_operand(*h), + detail::channelize_cast_operand(in_val)); sample_idx -= num_channels; h++; } @@ -261,8 +493,8 @@ __global__ void ChannelizePoly1D(OutType output, InType input, FilterType filter const input_t in_val = input_b(sample_idx); const filter_t h_val = filter_acc(h_ind); detail::channelize_cmac(accum, - detail::channelize_cast_filter(h_val), - detail::channelize_cast_input(in_val)); + detail::channelize_cast_operand(h_val), + detail::channelize_cast_operand(in_val)); h_ind += num_channels; sample_idx -= num_channels; } @@ -308,8 +540,8 @@ __global__ void ChannelizePoly1D(OutType output, InType input, FilterType filter const input_t in_val = input_b(sample_idx); const filter_t h_val = filter_acc(h_ind); detail::channelize_cmac(accum, - detail::channelize_cast_filter(h_val), - detail::channelize_cast_input(in_val)); + detail::channelize_cast_operand(h_val), + detail::channelize_cast_operand(in_val)); h_ind += num_channels; sample_idx -= num_channels; } @@ -329,12 +561,12 @@ __global__ void ChannelizePoly1D(OutType output, InType input, FilterType filter // FilterInSmem: when true, filter taps are cached in shared memory; // when false, filter taps are read from global/L2. // -// Block: dim3(CTILE, NOUT) +// Block: dim3(CTILE, MaximallyDecimated ? SmemTiledMaxDecYThreads : NOUT) // Grid: dim3(time_blocks, channel_tiles, batches) // // Shared memory layout (FilterInSmem = true): // smem_filter: [P][CTILE] for D==M, or [CTILE][K][P] for D -__launch_bounds__(CTILE * NOUT) +__launch_bounds__(CTILE * (MaximallyDecimated ? detail::cpoly::SmemTiledMaxDecYThreads : NOUT)) // See ChannelizePoly1D for the out_elem_offset output-element window semantics. // Every local output element is shifted to its global index before it drives the // input footprint / phase / circular-buffer block index (via max_bidx and the @@ -380,8 +612,10 @@ __global__ void ChannelizePoly1D_SmemTiled( const int32_t cx = static_cast(threadIdx.x); const int32_t ty = static_cast(threadIdx.y); const int32_t tid = ty * CTILE + cx; - assert(blockDim.x * blockDim.y == CTILE * NOUT && blockDim.z == 1); - const int32_t nthreads = CTILE * NOUT; + assert(blockDim.x * blockDim.y == CTILE * + (MaximallyDecimated ? detail::cpoly::SmemTiledMaxDecYThreads : NOUT) && blockDim.z == 1); + const int32_t nthreads = CTILE * + (MaximallyDecimated ? detail::cpoly::SmemTiledMaxDecYThreads : NOUT); const int32_t tile_base = static_cast(blockIdx.y) * CTILE; const int32_t c = tile_base + cx; const bool active = (c < M); @@ -390,7 +624,8 @@ __global__ void ChannelizePoly1D_SmemTiled( const int32_t L = MaximallyDecimated ? 0 : (M - static_cast(decimation_factor)); const int32_t filter_stride = K * P; // per-channel filter block size - const int32_t height = P + NOUT - 1; + const int32_t height = MaximallyDecimated ? P + NOUT - 1 : + static_cast(detail::cpoly::SmemTiledInputHeight(P, M, decimation_factor, NOUT)); filter_t *smem_filter_base = nullptr; input_t *smem_input = nullptr; @@ -411,18 +646,8 @@ __global__ void ChannelizePoly1D_SmemTiled( smem_input = reinterpret_cast(smem_raw); } - // TensorAccessors bind per-block batch coords once. After binding, - // input_b(sample_idx) reads one input sample for this batch, and - // output_b(t, ch) writes one output sample. Fast path folds the strides - // into pointer arithmetic; slow path forwards to operator(). - detail::TensorAccessor input_acc(input); - detail::TensorAccessor output_acc(output); - detail::TensorAccessor filter_acc(filter); - - const auto in_batch_idx = BlockToIdx(input, blockIdx.z, 1); - const auto out_batch_idx = BlockToIdx(output, blockIdx.z, 2); - auto input_b = detail::bind_first_n(input_acc, in_batch_idx); - auto output_b = detail::bind_first_n(output_acc, out_batch_idx); + auto [input_b, output_b, filter_acc] = + detail::cpoly::MakeAccessors(output, input, filter, blockIdx.z); if constexpr (FilterInSmem) { // Load filter into smem @@ -456,11 +681,6 @@ __global__ void ChannelizePoly1D_SmemTiled( for (int32_t k = 0; k < K; k++) { rotations[k] = static_cast((static_cast(k) * decimation_factor) % M); } - // TODO-PERF: Each channel rotates through K filter phases. Currently, we store all K phases - // per channel in shared memory. It is possible for K * CTILE to exceed the total number - // of channels, in which case we would be better off storing the full filter in shared memory. - // Furthermore, if there are fewer than K output points per channel generated in this CTA, - // then we could store only the required phases. for (int32_t i = tid; i < CTILE * filter_stride; i += nthreads) { const int32_t local_channel = i / filter_stride; const int32_t kp = i % filter_stride; @@ -541,120 +761,83 @@ __global__ void ChannelizePoly1D_SmemTiled( __syncthreads(); if constexpr (MaximallyDecimated) { - // bidx = t, causal_count = t+1, last_arrived >= s always true. - // Track buf_row incrementally across iterations to avoid modulo. - - // Seed buf_row for ty=0 at the global block index of start_elem - int32_t buf_row_base = static_cast((start_elem + out_elem_offset) % height); - // Per-iteration advance: NOUT output steps = NOUT buf_row advance. This is a defensive modulo - // so that we can keep buf_row_base in [0, height) with only a conditional subtraction. - const int32_t nout_wrap = NOUT % height; - + // The dispatcher launches this path with NOUT=16. Critical 32x4 plans use the + // oversampled path below with K=1. + constexpr int YTHREADS = detail::cpoly::SmemTiledMaxDecYThreads; + static_assert(NOUT % YTHREADS == 0); + constexpr int OUTPUTS_PER_THREAD = NOUT / YTHREADS; const IdxT last_start = start_elem + ((last_elem - start_elem) / NOUT) * NOUT; - for (IdxT next_start = start_elem; next_start <= last_start; next_start += NOUT) { - const IdxT t = next_start + ty; - if (t <= last_elem && active) { - accum_t accum{}; - - const IdxT tg = t + out_elem_offset; // global output element index - - // buf_row for this thread's output step - int32_t my_buf_row = buf_row_base + ty; - if (my_buf_row >= height) my_buf_row -= height; - - const IdxT newest_raw = static_cast(s) + tg * M; - int32_t h_skip = 0; - int32_t niter = static_cast( - cuda::std::min(static_cast(P), tg + 1)); - if (newest_raw >= input_len) { - h_skip = 1; - niter = static_cast( - cuda::std::min(static_cast(P - 1), tg)); - if (--my_buf_row < 0) my_buf_row += height; - } + for (IdxT next_start = start_elem; + next_start <= last_start; next_start += NOUT) { + accum_t accum[OUTPUTS_PER_THREAD]{}; + int32_t sample_row[OUTPUTS_PER_THREAD]; + bool valid[OUTPUTS_PER_THREAD]; + #pragma unroll + for (int q = 0; q < OUTPUTS_PER_THREAD; q++) { + const IdxT t = next_start + ty + q * YTHREADS; + valid[q] = t <= last_elem; + const IdxT tg = t + out_elem_offset; + sample_row[q] = static_cast(tg % height); + } - const int32_t prologue = cuda::std::min(my_buf_row + 1, niter); - const int32_t epilogue = niter - prologue; - // Single running counter instead of separate `p` and `h_ind`. - // Each layout's access pattern reduces to (init, stride): - // Full: smem[(p+h_skip)*M + c] -> init h_skip*M+c, stride M - // Rotated: smem[(p+h_skip)*CTILE + cx] -> init h_skip*CTILE+cx, stride CTILE - // Global: filter[c + (p+h_skip)*M] -> init c+h_skip*M, stride M - int32_t filter_idx; - if constexpr (FilterInSmem && !FilterFullLayout) { - filter_idx = h_skip * CTILE + cx; - } else { - filter_idx = c + h_skip * M; - } - for (int32_t i = 0; i < prologue; i++) { - filter_t hv; - if constexpr (FilterInSmem) { - hv = smem_filter_base[filter_idx]; - } else { - hv = filter_acc(static_cast(filter_idx)); - } - const input_t iv = smem_input[my_buf_row * CTILE + cx]; - detail::channelize_cmac(accum, - detail::channelize_cast_filter(hv), - detail::channelize_cast_input(iv)); - my_buf_row--; - if constexpr (FilterInSmem && !FilterFullLayout) { - filter_idx += CTILE; - } else { - filter_idx += M; - } - } - my_buf_row = height - 1; - for (int32_t i = 0; i < epilogue; i++) { + if (active) { + int32_t filter_idx = (FilterInSmem && !FilterFullLayout) ? cx : c; + for (int32_t p = 0; p < P; p++) { filter_t hv; if constexpr (FilterInSmem) { hv = smem_filter_base[filter_idx]; } else { - hv = filter_acc(static_cast(filter_idx)); + hv = (filter_idx < filter_full_len) + ? filter_acc(static_cast(filter_idx)) + : static_cast(0); } - const input_t iv = smem_input[my_buf_row * CTILE + cx]; - detail::channelize_cmac(accum, - detail::channelize_cast_filter(hv), - detail::channelize_cast_input(iv)); - my_buf_row--; - if constexpr (FilterInSmem && !FilterFullLayout) { - filter_idx += CTILE; - } else { - filter_idx += M; + const auto hav = detail::channelize_cast_operand(hv); + #pragma unroll + for (int q = 0; q < OUTPUTS_PER_THREAD; q++) { + const input_t iv = smem_input[sample_row[q] * CTILE + cx]; + detail::channelize_cmac( + accum[q], hav, detail::channelize_cast_operand(iv)); + if (--sample_row[q] < 0) { + sample_row[q] += height; + } } + filter_idx += (FilterInSmem && !FilterFullLayout) ? CTILE : M; } + } - output_b(t, c) = static_cast(accum); + #pragma unroll + for (int q = 0; q < OUTPUTS_PER_THREAD; q++) { + const IdxT t = next_start + ty + q * YTHREADS; + if (active && valid[q]) { + output_b(t, c) = static_cast(accum[q]); + } } if (next_start < last_start) { - // Incremental buf_row_base advance (no modulo) - buf_row_base += nout_wrap; - if (buf_row_base >= height) buf_row_base -= height; - - // Ensure all threads have finished reading smem_input before overwriting + // Ensure all threads have finished reading smem_input before overwriting it. __syncthreads(); - // Load NOUT new rows. For D==M, exactly NOUT new samples arrive. - // Each ty-lane loads one row (its cx column). - { - const IdxT next_end = cuda::std::min(next_start + static_cast(2 * NOUT) - 1, last_elem); - const int32_t new_rows = static_cast(max_bidx(next_end) - loaded_up_to); - int32_t lr = static_cast((loaded_up_to + 1) % height); - if (ty < new_rows) { - int32_t my_lr = lr + ty; - if (my_lr >= height) my_lr -= height; - load_smem_elem(my_lr, cx, loaded_up_to + 1 + ty); + // Only load rows needed by this block. Speculative rows + // beyond its final output can overflow 32-bit raw indices. + const int32_t new_rows = static_cast( + cuda::std::min(static_cast(NOUT), + last_elem - next_start - NOUT + 1)); + #pragma unroll + for (int q = 0; q < OUTPUTS_PER_THREAD; q++) { + const int32_t row = ty + q * YTHREADS; + if (row < new_rows) { + const IdxT global_row = next_start + out_elem_offset + NOUT + row; + const int32_t smem_row = static_cast(global_row % height); + load_smem_elem(smem_row, cx, global_row); } - loaded_up_to += new_rows; } - // Ensure new rows are visible before next iteration's compute + // Ensure new rows are visible before the next compute. __syncthreads(); } } } else { - // D < M: oversampled path + // Oversampled path (D < M), also used by critical 32x4 plans. const IdxT last_start = start_elem + ((last_elem - start_elem) / NOUT) * NOUT; for (IdxT next_start = start_elem; next_start <= last_start; next_start += NOUT) { const IdxT t = next_start + ty; @@ -703,6 +886,7 @@ __global__ void ChannelizePoly1D_SmemTiled( } else { filter_idx = phase + h_skip * M; } + #pragma unroll 8 for (int32_t i = 0; i < prologue; i++) { filter_t hv; if constexpr (FilterInSmem) { @@ -712,8 +896,8 @@ __global__ void ChannelizePoly1D_SmemTiled( } const input_t iv = smem_input[buf_row * CTILE + cx]; detail::channelize_cmac(accum, - detail::channelize_cast_filter(hv), - detail::channelize_cast_input(iv)); + detail::channelize_cast_operand(hv), + detail::channelize_cast_operand(iv)); buf_row--; if constexpr (FilterInSmem && !FilterFullLayout) { filter_idx += 1; @@ -722,6 +906,7 @@ __global__ void ChannelizePoly1D_SmemTiled( } } buf_row = height - 1; + #pragma unroll 8 for (int32_t i = 0; i < epilogue; i++) { filter_t hv; if constexpr (FilterInSmem) { @@ -731,8 +916,8 @@ __global__ void ChannelizePoly1D_SmemTiled( } const input_t iv = smem_input[buf_row * CTILE + cx]; detail::channelize_cmac(accum, - detail::channelize_cast_filter(hv), - detail::channelize_cast_input(iv)); + detail::channelize_cast_operand(hv), + detail::channelize_cast_operand(iv)); buf_row--; if constexpr (FilterInSmem && !FilterFullLayout) { filter_idx += 1; @@ -820,14 +1005,8 @@ __global__ void ChannelizePoly1D_Smem(OutType output, InType input, FilterType f const int32_t ty = static_cast(threadIdx.y); const int32_t by = static_cast(blockDim.y); - // TensorAccessors with per-block batch binding (see ChannelizePoly1D_SmemTiled). - detail::TensorAccessor input_acc(input); - detail::TensorAccessor output_acc(output); - detail::TensorAccessor filter_acc(filter); - const auto in_batch_idx = BlockToIdx(input, blockIdx.z, 1); - const auto out_batch_idx = BlockToIdx(output, blockIdx.z, 2); - auto input_b = detail::bind_first_n(input_acc, in_batch_idx); - auto output_b = detail::bind_first_n(output_acc, out_batch_idx); + auto [input_b, output_b, filter_acc] = + detail::cpoly::MakeAccessors(output, input, filter, blockIdx.z); for (int32_t t = tid; t < filter_full_len; t += nthreads) { smem_h[t] = filter_acc(t); @@ -905,16 +1084,16 @@ __global__ void ChannelizePoly1D_Smem(OutType output, InType input, FilterType f // Apply the filter h in reverse order below to flip the filter for convolution for (int32_t k = 0; k < prologue_count; k++) { detail::channelize_cmac(accum, - detail::channelize_cast_filter(*h), - detail::channelize_cast_input(*sample)); + detail::channelize_cast_operand(*h), + detail::channelize_cast_operand(*sample)); sample += num_channels; h -= num_channels; } sample = smem_input + (num_channels - 1 - chan); for (int32_t k = 0; k < epilogue_count; k++) { detail::channelize_cmac(accum, - detail::channelize_cast_filter(*h), - detail::channelize_cast_input(*sample)); + detail::channelize_cast_operand(*h), + detail::channelize_cast_operand(*sample)); sample += num_channels; h -= num_channels; } @@ -926,188 +1105,1092 @@ __global__ void ChannelizePoly1D_Smem(OutType output, InType input, FilterType f } } -// See ChannelizePoly1D for the out_elem_offset output-element window semantics. -template -__launch_bounds__(THREADS) -__global__ void ChannelizePoly1D_FusedChan(OutType output, InType input, FilterType filter, index_t out_elem_offset) -{ - using output_t = typename OutType::value_type; +template +struct ChannelizePolyFusedSmallTraits { using input_t = typename InType::value_type; + using output_t = typename OutType::value_type; using filter_t = typename FilterType::value_type; - static_assert(! is_complex_v, - "channelize_poly: accumulator type must be real; it will be treated as complex when necessary"); - // If the output is complex, then then accumulator is complex. Otherwise, the accumulator is real. - using filtering_accum_t = cuda::std::conditional_t || is_complex_v, - typename detail::scalar_to_complex::ctype, AccumType>; using complex_accum_t = typename detail::scalar_to_complex::ctype; + using accum_t = cuda::std::conditional_t< + is_complex_v || is_complex_v, complex_accum_t, AccumType>; + static constexpr bool ComplexInput = is_complex_v; + static constexpr bool UseM2PairLeaf = + (cuda::std::is_same_v || + cuda::std::is_same_v) && + !is_complex_v && cuda::std::is_same_v && + (cuda::std::is_same_v || + cuda::std::is_same_v); + + static_assert(NUM_CHAN >= 2 && NUM_CHAN <= 6); + static_assert(is_complex_v); + static_assert(!is_complex_v); +}; + +template +__device__ __forceinline__ Complex ChannelizePolyFusedSmallComplex(Real real, Imag imag) +{ + using scalar_t = typename inner_op_type_t::type; + return Complex{static_cast(real), static_cast(imag)}; +} + +template +__device__ __forceinline__ void ChannelizePolyFusedM2LoadFilter( + FilterAccessor &filter, detail::ChannelizePolyM2Pair *smem_filter, + int32_t taps_per_channel, index_t filter_len, int32_t tid) +{ + for (int32_t p = tid; p < taps_per_channel; p += THREADS) { + const index_t h = static_cast(p) * 2; + smem_filter[p] = { + filter(h), (h + 1 < filter_len) ? filter(h + 1) : Scalar{0}}; + } +} + +template +__device__ __forceinline__ void ChannelizePolyFusedM2Store( + OutputAccessor &output, index_t output_len, index_t block_start, + int32_t tid, const Scalar (&a0r)[OUTPUTS_PER_THREAD], const Scalar (&a0i)[OUTPUTS_PER_THREAD], + const Scalar (&a1r)[OUTPUTS_PER_THREAD], const Scalar (&a1i)[OUTPUTS_PER_THREAD]) +{ + #pragma unroll + for (int q = 0; q < OUTPUTS_PER_THREAD; q++) { + const index_t t = block_start + tid + q * THREADS; + if (t < output_len) { + output(t, static_cast(0)) = + cuda::std::complex{a0r[q] + a1r[q], + a0i[q] + a1i[q]}; + output(t, static_cast(1)) = + cuda::std::complex{a0r[q] - a1r[q], + a0i[q] - a1i[q]}; + } + } +} + +template +__device__ __forceinline__ cuda::std::complex +ChannelizePolyTwiddle(int32_t numerator, int32_t denominator) +{ + const Scalar angle = static_cast(2) * + static_cast(M_PI) * static_cast(numerator) / + static_cast(denominator); + Scalar sinx, cosx; + if constexpr (cuda::std::is_same_v) { + sincos(angle, &sinx, &cosx); + } else { + sincosf(angle, &sinx, &cosx); + } + return {cosx, sinx}; +} + +// A quarter wave covers every radix-2 root in the current dispatch (up to 64). +// Store native-precision constants once instead of evaluating sincos per CTA. +template +static __device__ __constant__ const Scalar ChannelizePolyRadix2Cos[] = { + static_cast(1.0), // cos(0) + static_cast(0.99518472667219693), // cos(pi/32) + static_cast(0.98078528040323043), // cos(2*pi/32) + static_cast(0.95694033573220882), // cos(3*pi/32) + static_cast(0.92387953251128674), // cos(4*pi/32) + static_cast(0.88192126434835505), // cos(5*pi/32) + static_cast(0.83146961230254524), // cos(6*pi/32) + static_cast(0.77301045336273699), // cos(7*pi/32) + static_cast(0.70710678118654757), // cos(8*pi/32) + static_cast(0.63439328416364549), // cos(9*pi/32) + static_cast(0.55557023301960218), // cos(10*pi/32) + static_cast(0.47139673682599764), // cos(11*pi/32) + static_cast(0.38268343236508978), // cos(12*pi/32) + static_cast(0.29028467725446239), // cos(13*pi/32) + static_cast(0.19509032201612828), // cos(14*pi/32) + static_cast(0.098017140329560604), // cos(15*pi/32) + static_cast(0.0), // cos(pi/2) +}; + +template +__device__ __forceinline__ cuda::std::complex +ChannelizePolyRadix2Twiddle(int32_t j) +{ + if constexpr (RADIX_2 <= 64) { + const int32_t phase = j * (64 / RADIX_2); + const int32_t cos_index = phase <= 16 ? phase : 32 - phase; + const int32_t sin_index = phase <= 16 ? 16 - phase : phase - 16; + const Scalar cosx = ChannelizePolyRadix2Cos[cos_index]; + const Scalar sinx = ChannelizePolyRadix2Cos[sin_index]; + return {phase <= 16 ? cosx : -cosx, sinx}; + } else { + // Keep larger compile-time radices usable without expanding the table. + return ChannelizePolyTwiddle(j, RADIX_2); + } +} + +template +__device__ __forceinline__ void ChannelizePolyRadix3( + const Complex &x0, const Complex &x1, const Complex &x2, Complex &y0, Complex &y1, Complex &y2) +{ + using scalar_t = typename inner_op_type_t::type; + const scalar_t half = static_cast(0.5); // cos(pi/3) + const scalar_t sin60 = static_cast(0.8660254037844386); // sin(pi/3) + const Complex sum = x1 + x2; + const Complex diff = x1 - x2; + const Complex base = x0 - half * sum; + const Complex jdiff = ChannelizePolyFusedSmallComplex( + -sin60 * diff.imag(), sin60 * diff.real()); + y0 = x0 + sum; + y1 = base + jdiff; + y2 = base - jdiff; +} + +template +__device__ __forceinline__ void ChannelizePolyRadix5( + const Complex &x0, const Complex &x1, const Complex &x2, + const Complex &x3, const Complex &x4, Complex &y0, Complex &y1, + Complex &y2, Complex &y3, Complex &y4) +{ + using scalar_t = typename inner_op_type_t::type; + const scalar_t c1 = static_cast(0.30901699437494745); // cos(2*pi/5) + const scalar_t c2 = static_cast(-0.80901699437494745); // cos(4*pi/5) + const scalar_t s1 = static_cast(0.95105651629515357); // sin(2*pi/5) + const scalar_t s2 = static_cast(0.58778525229247314); // sin(4*pi/5) + const Complex sum14 = x1 + x4; + const Complex diff14 = x1 - x4; + const Complex sum23 = x2 + x3; + const Complex diff23 = x2 - x3; + const Complex base1 = x0 + c1 * sum14 + c2 * sum23; + const Complex base2 = x0 + c2 * sum14 + c1 * sum23; + const Complex odd1 = s1 * diff14 + s2 * diff23; + const Complex odd2 = s2 * diff14 - s1 * diff23; + const Complex jodd1 = ChannelizePolyFusedSmallComplex(-odd1.imag(), odd1.real()); + const Complex jodd2 = ChannelizePolyFusedSmallComplex(-odd2.imag(), odd2.real()); + y0 = x0 + sum14 + sum23; + y1 = base1 + jodd1; + y2 = base2 + jodd2; + y3 = base2 - jodd2; + y4 = base1 - jodd1; +} + +// Positive-sign, unnormalized DFT used by the channelizer. These fixed +// butterflies replace FusedChan's per-CTA sincos table and quadratic DFT for +// the small channel counts that it historically handled. +template +__device__ __forceinline__ void ChannelizePolySmallDFT( + const Complex (&x)[NUM_CHAN], Complex (&y)[NUM_CHAN]) +{ + using scalar_t = typename inner_op_type_t::type; + static_assert(NUM_CHAN >= 2 && NUM_CHAN <= 6); + if constexpr (NUM_CHAN == 2) { + y[0] = x[0] + x[1]; + y[1] = x[0] - x[1]; + } else if constexpr (NUM_CHAN == 3) { + ChannelizePolyRadix3(x[0], x[1], x[2], y[0], y[1], y[2]); + } else if constexpr (NUM_CHAN == 4) { + Complex sum[2], diff[2]; + #pragma unroll + for (int i = 0; i < 2; ++i) { + sum[i] = x[i] + x[i + 2]; + diff[i] = x[i] - x[i + 2]; + } + const Complex jdiff1 = ChannelizePolyFusedSmallComplex( + -diff[1].imag(), diff[1].real()); + y[0] = sum[0] + sum[1]; + y[1] = diff[0] + jdiff1; + y[2] = sum[0] - sum[1]; + y[3] = diff[0] - jdiff1; + } else if constexpr (NUM_CHAN == 5) { + ChannelizePolyRadix5(x[0], x[1], x[2], x[3], x[4], y[0], y[1], y[2], y[3], y[4]); + } else { + const scalar_t half = static_cast(0.5); // cos(pi/3) + const scalar_t sin60 = static_cast(0.8660254037844386); // sin(pi/3) + Complex even[3], odd[3]; + ChannelizePolyRadix3(x[0], x[2], x[4], even[0], even[1], even[2]); + ChannelizePolyRadix3(x[1], x[3], x[5], odd[0], odd[1], odd[2]); + Complex twiddled[3]; + twiddled[0] = odd[0]; + twiddled[1] = ChannelizePolyFusedSmallComplex(half, sin60) * odd[1]; + twiddled[2] = ChannelizePolyFusedSmallComplex(-half, sin60) * odd[2]; + #pragma unroll + for (int i = 0; i < 3; ++i) { + y[i] = even[i] + twiddled[i]; + y[i + 3] = even[i] - twiddled[i]; + } + } +} + +template +__device__ __forceinline__ void ChannelizePolyFusedSmallTransformStore( + OutputAccessor &output, index_t output_len, index_t block_start, int32_t tid, + const AccumType (&accum)[OUTPUTS_PER_THREAD][NUM_CHAN]) +{ + using scalar_t = typename inner_op_type_t::type; + using complex_t = typename detail::scalar_to_complex::ctype; + static_assert(cuda::std::is_same_v || + cuda::std::is_same_v); + #pragma unroll + for (int q = 0; q < OUTPUTS_PER_THREAD; q++) { + const index_t t = block_start + tid + q * THREADS; + if (t < output_len) { + complex_t transformed[NUM_CHAN]; + complex_t dft_input[NUM_CHAN]; + #pragma unroll + for (int c = 0; c < NUM_CHAN; c++) { + dft_input[c] = static_cast(accum[q][c]); + } + ChannelizePolySmallDFT(dft_input, transformed); + #pragma unroll + for (int c = 0; c < NUM_CHAN; c++) { + output(t, static_cast(c)) = transformed[c]; + } + } + } +} + +// Critically sampled row-owned leaf for the small-M portion of the unified +// fused radix backend. A CTA stages one overlapping FIR tile; each thread then +// computes every branch and the fixed K-by-power-of-two DFT for its rows. +template +__device__ __forceinline__ void ChannelizePolyFusedSmallCachedBody( + OutType &output, const InType &input, const FilterType &filter, + index_t out_elem_offset, uint8_t *smem_raw) +{ + constexpr int BLOCK_ROWS = THREADS * OUTPUTS_PER_THREAD; + using traits = ChannelizePolyFusedSmallTraits; + using input_t = typename traits::input_t; + using filter_t = typename traits::filter_t; + using accum_t = typename traits::accum_t; constexpr int InRank = InType::Rank(); constexpr int OutRank = OutType::Rank(); - constexpr int OutElemRank = OutRank-2; + constexpr int OutElemRank = OutRank - 2; + const index_t input_len = input.Size(InRank - 1); + const index_t output_len = output.Size(OutElemRank); + const index_t filter_len = filter.Size(0); + const int32_t P = static_cast((filter_len + NUM_CHAN - 1) / NUM_CHAN); + + const detail::cpoly::FusedRadixSmemLayout< + NUM_CHAN, true, input_t, filter_t, AccumType> layout(P, NUM_CHAN); + filter_t *smem_filter = reinterpret_cast(smem_raw); + input_t *smem_input = reinterpret_cast(smem_raw + layout.input_offset); + + auto [input_b, output_b, filter_acc] = + detail::cpoly::MakeAccessors(output, input, filter, blockIdx.z); + + const int32_t tid = static_cast(threadIdx.x); + const int32_t filter_elements = static_cast(layout.filter_elements); + for (int32_t i = tid; i < filter_elements; i += THREADS) { + smem_filter[i] = (i < filter_len) + ? filter_acc(static_cast(i)) : filter_t{}; + } - const index_t input_len = input.Size(InRank-1); - const index_t output_len_per_channel = output.Size(OutElemRank); - const index_t filter_full_len = filter.Size(0); - const index_t filter_phase_len = (filter_full_len + NUM_CHAN - 1) / NUM_CHAN; + const index_t block_start = static_cast(blockIdx.x) * BLOCK_ROWS; + const index_t input_base = (block_start + out_elem_offset - (P - 1)) * NUM_CHAN; + const int32_t input_elements = static_cast(layout.input_elements); + for (int32_t i = tid; i < input_elements; i += THREADS) { + // Linear staging preserves fully coalesced global loads. This is faster than transposing + // the tile while it is staged, despite the modest shared-memory conflicts during FIR reuse. + smem_input[i] = detail::cpoly::LoadRelativeInput(input_b, input_base, input_len, i); + } + __syncthreads(); - const int elem_block = blockIdx.x; - const int tid = threadIdx.x; + accum_t accum[OUTPUTS_PER_THREAD][NUM_CHAN]{}; + const index_t t = block_start + tid; + if (t < output_len) { + const int32_t output_row = P - 1 + tid; + for (int32_t p = 0; p < P; p++) { + #pragma unroll + for (int c = 0; c < NUM_CHAN; c++) { + // Reuse each filter coefficient across the output accumulators. + const filter_t hv = smem_filter[p * NUM_CHAN + c]; + #pragma unroll + for (int q = 0; q < OUTPUTS_PER_THREAD; q++) { + if (q == 0 || t + q * THREADS < output_len) { + const input_t xv = smem_input[ + (output_row + q * THREADS - p) * NUM_CHAN + (NUM_CHAN - 1 - c)]; + detail::channelize_cmac( + accum[q][c], detail::channelize_cast_operand(hv), + detail::channelize_cast_operand(xv)); + } + } + } + } + } - // TensorAccessors + batch bind. Batch coords come from BlockToIdx; inner - // indices (sample for input, (t, chan) for output) are supplied per access. - detail::TensorAccessor input_acc(input); - detail::TensorAccessor output_acc(output); - detail::TensorAccessor filter_acc(filter); - const auto in_batch_idx = BlockToIdx(input, blockIdx.z, 1); - const auto out_batch_idx = BlockToIdx(output, blockIdx.z, 2); - auto input_b = detail::bind_first_n(input_acc, in_batch_idx); - auto output_b = detail::bind_first_n(output_acc, out_batch_idx); + ChannelizePolyFusedSmallTransformStore< + THREADS, NUM_CHAN, OUTPUTS_PER_THREAD>( + output_b, output_len, block_start, tid, accum); +} - constexpr index_t ELEMS_PER_BLOCK = detail::cpoly::ElemsPerThread * THREADS; - const index_t first_out_elem = elem_block * detail::cpoly::ElemsPerThread * THREADS; - const index_t last_out_elem = cuda::std::min( - output_len_per_channel - 1, first_out_elem + ELEMS_PER_BLOCK - 1); +// Direct short-filter leaf of the same small-M backend. This deliberately shares the enclosing +// global-kernel instantiation with the cached leaf: P is a uniform runtime choice, so supporting it +// does not add another exported CUDA kernel or another host dispatch specialization. +template +__device__ __forceinline__ void ChannelizePolyFusedSmallDirectBody( + OutType &output, const InType &input, const FilterType &filter, index_t out_elem_offset) +{ + constexpr int BLOCK_ROWS = THREADS * OUTPUTS_PER_THREAD; + using traits = ChannelizePolyFusedSmallTraits; + using input_t = typename traits::input_t; + using filter_t = typename traits::filter_t; + using accum_t = typename traits::accum_t; - // Versions of CUDA prior to 11.8 do not allow static shared memory allocations of - // cuda::std::complex types due to it having no trivial constructor. This workaround - // prevents an 'initializer not allowed for __shared__ variable' error. - __shared__ __align__(16) uint8_t smem_eij_workaround[sizeof(complex_accum_t)*NUM_CHAN*NUM_CHAN]; - complex_accum_t (&smem_eij)[NUM_CHAN][NUM_CHAN] = reinterpret_cast(smem_eij_workaround); - // Pre-compute the DFT complex exponentials and store in shared memory - for (int t = tid; t < NUM_CHAN*NUM_CHAN; t += THREADS) { - const int i = t / NUM_CHAN; - const int j = t % NUM_CHAN; - if constexpr (cuda::std::is_same_v) { - const double arg = 2.0 * M_PI * j * i / NUM_CHAN; - double sinx, cosx; - sincos(arg, &sinx, &cosx); - complex_accum_t eij { static_cast(cosx), static_cast(sinx) }; - smem_eij[i][j] = eij; - } else { - const float arg = 2.0f * static_cast(M_PI) * j * i / NUM_CHAN; - float sinx, cosx; - sincosf(arg, &sinx, &cosx); - complex_accum_t eij { static_cast(cosx), static_cast(sinx) }; - smem_eij[i][j] = eij; + constexpr int InRank = InType::Rank(); + constexpr int OutRank = OutType::Rank(); + constexpr int OutElemRank = OutRank - 2; + const index_t input_len = input.Size(InRank - 1); + const index_t output_len = output.Size(OutElemRank); + const index_t filter_len = filter.Size(0); + const int32_t P = static_cast((filter_len + NUM_CHAN - 1) / NUM_CHAN); + + auto [input_b, output_b, filter_acc] = + detail::cpoly::MakeAccessors(output, input, filter, blockIdx.z); + + const int32_t tid = static_cast(threadIdx.x); + const index_t block_start = static_cast(blockIdx.x) * BLOCK_ROWS; + if constexpr (NUM_CHAN == 2 && cuda::std::is_same_v) { + // This leaf has at most four taps per channel; fixed loops avoid + // runtime FIR-loop overhead for double accumulation. + #pragma unroll + for (int q = 0; q < OUTPUTS_PER_THREAD; q++) { + accum_t accum[1][NUM_CHAN]{}; + const index_t t = block_start + tid + q * THREADS; + if (t < output_len) { + const index_t g = t + out_elem_offset; + const int32_t last = static_cast( + cuda::std::min(static_cast(P), g + 1)); + constexpr int max_taps = detail::cpoly::FusedSmallDirectMaxTaps< + NUM_CHAN, input_t, AccumType>(); + #pragma unroll + for (int32_t p = 0; p < max_taps; p++) { + if (p >= last) continue; + #pragma unroll + for (int c = 0; c < NUM_CHAN; c++) { + const index_t h = static_cast(p) * NUM_CHAN + c; + const index_t x = (g - p) * NUM_CHAN + NUM_CHAN - 1 - c; + if (h >= filter_len || x >= input_len) continue; + detail::channelize_cmac( + accum[0][c], detail::channelize_cast_operand(filter_acc(h)), + detail::channelize_cast_operand(input_b(x))); + } + } + } + ChannelizePolyFusedSmallTransformStore( + output_b, output_len, block_start + q * THREADS, tid, accum); } + return; + } + accum_t accum[OUTPUTS_PER_THREAD][NUM_CHAN]{}; + #pragma unroll + for (int q = 0; q < OUTPUTS_PER_THREAD; q++) { + const index_t t = block_start + tid + q * THREADS; + if (t < output_len) { + const index_t g = t + out_elem_offset; + constexpr bool peel_bounds = + detail::cpoly::FusedSmallDirectMaxTaps< + NUM_CHAN, input_t, AccumType>() > 4; + const int32_t last = peel_bounds ? static_cast( + cuda::std::min(static_cast(P), g + 1)) : P; + for (int32_t p = 0; p < last; p++) { + const index_t input_row = g - p; + if constexpr (!peel_bounds) { + if (input_row < 0) continue; + } + const index_t input_base = input_row * NUM_CHAN; + // Interior taps have complete input and filter rows. Preserve + // short-direct scheduling for the other types/channel counts. + if (peel_bounds && p > 0 && p + 1 < P) { + #pragma unroll + for (int c = 0; c < NUM_CHAN; c++) { + const filter_t hv = filter_acc(static_cast(p) * NUM_CHAN + c); + const input_t xv = input_b(input_base + NUM_CHAN - 1 - c); + detail::channelize_cmac( + accum[q][c], detail::channelize_cast_operand(hv), + detail::channelize_cast_operand(xv)); + } + continue; + } + #pragma unroll + for (int c = 0; c < NUM_CHAN; c++) { + const index_t h = static_cast(p) * NUM_CHAN + c; + const index_t x = input_base + NUM_CHAN - 1 - c; + if (h < filter_len && x < input_len) { + const filter_t hv = filter_acc(h); + const input_t xv = input_b(x); + detail::channelize_cmac( + accum[q][c], detail::channelize_cast_operand(hv), + detail::channelize_cast_operand(xv)); + } + } + } + } + } + + ChannelizePolyFusedSmallTransformStore< + THREADS, NUM_CHAN, OUTPUTS_PER_THREAD>( + output_b, output_len, block_start, tid, accum); +} + +// Maximally-decimated two-channel leaf of the fused radix/power-of-two +// family. A CTA stages a contiguous input window once, then reuses it across +// all FIR outputs before applying the exact two-point inverse DFT. +template +__device__ __forceinline__ void ChannelizePolyFusedM2D2Body( + OutType &output, const InType &input, const FilterType &filter, + index_t out_elem_offset, index_t elems_per_channel_per_cta, uint8_t *smem_raw) +{ + constexpr int BLOCK_ROWS = THREADS * OUTPUTS_PER_THREAD; + constexpr int NUM_CHAN = 2; + using traits = ChannelizePolyFusedSmallTraits; + using input_t = typename traits::input_t; + using scalar_t = AccumType; + using filter_pair_t = detail::ChannelizePolyM2Pair; + constexpr bool ComplexInput = traits::ComplexInput; + using shared_input_t = detail::ChannelizePolyM2Pair; + + constexpr int InRank = InType::Rank(); + constexpr int OutRank = OutType::Rank(); + constexpr int OutElemRank = OutRank - 2; + const index_t input_len = input.Size(InRank - 1); + const index_t output_len = output.Size(OutElemRank); + const index_t filter_len = filter.Size(0); + const int32_t P = static_cast((filter_len + 1) / NUM_CHAN); + + // Keep the maximum shared layout, but stage only the active output span. + const int32_t block_rows = ComplexInput && + cuda::std::is_same_v + ? static_cast(elems_per_channel_per_cta) : BLOCK_ROWS; + + if constexpr (cuda::std::is_same_v) { + if (detail::cpoly::FusedSmallUseDirect< + NUM_CHAN, OutType, InType, FilterType, AccumType>(P, output_len)) { + // Keep the usual small shared allocation and launch geometry even + // when bypassing staging, so this needs no separate host policy. + ChannelizePolyFusedSmallDirectBody< + THREADS, NUM_CHAN, OUTPUTS_PER_THREAD, IsUnitStride, + OutType, InType, FilterType, AccumType>( + output, input, filter, out_elem_offset); + return; + } + } + + const detail::cpoly::FusedRadixSmemLayout< + NUM_CHAN, true, input_t, scalar_t, AccumType> layout(P, NUM_CHAN); + filter_pair_t *smem_filter = reinterpret_cast(smem_raw); + shared_input_t *smem_input = reinterpret_cast(smem_raw + layout.input_offset); + + auto [input_b, output_b, filter_acc] = + detail::cpoly::MakeAccessors(output, input, filter, blockIdx.z); + + const int32_t tid = static_cast(threadIdx.x); + ChannelizePolyFusedM2LoadFilter(filter_acc, smem_filter, P, filter_len, tid); + + const index_t block_start = static_cast(blockIdx.x) * block_rows; + const index_t input_base = (block_start + out_elem_offset - (P - 1)) * NUM_CHAN; + const int32_t input_rows = + static_cast(layout.input_elements / NUM_CHAN) - (BLOCK_ROWS - block_rows); + for (int32_t row = tid; row < input_rows; row += THREADS) { + const int32_t offset = row * NUM_CHAN; + const input_t even = detail::cpoly::LoadRelativeInput( + input_b, input_base, input_len, offset); + const input_t odd = detail::cpoly::LoadRelativeInput( + input_b, input_base, input_len, offset + 1); + smem_input[row] = {even, odd}; } __syncthreads(); - filtering_accum_t accum[NUM_CHAN]; - for (index_t t = first_out_elem+tid; t <= last_out_elem; t += THREADS) { - for (int i = 0; i < NUM_CHAN; i++) { - accum[i] = static_cast(0); + scalar_t a0r[OUTPUTS_PER_THREAD]{}; + scalar_t a0i[OUTPUTS_PER_THREAD]{}; + scalar_t a1r[OUTPUTS_PER_THREAD]{}; + scalar_t a1i[OUTPUTS_PER_THREAD]{}; + for (int32_t p = 0; p < P; p++) { + const filter_pair_t hv = smem_filter[p]; + #pragma unroll + for (int q = 0; q < OUTPUTS_PER_THREAD; q++) { + if (q > 0 && q * THREADS >= block_rows) continue; + const int32_t row = (P - 1) + tid + q * THREADS - p; + const shared_input_t xv = smem_input[row]; + if constexpr (ComplexInput) { + a0r[q] = hv.even * xv.odd.real() + a0r[q]; + a0i[q] = hv.even * xv.odd.imag() + a0i[q]; + a1r[q] = hv.odd * xv.even.real() + a1r[q]; + a1i[q] = hv.odd * xv.even.imag() + a1i[q]; + } else { + a0r[q] = hv.even * xv.odd + a0r[q]; + a1r[q] = hv.odd * xv.even + a1r[q]; + } } - const index_t g = t + out_elem_offset; // global output element (time) index - index_t first_ind = cuda::std::max(static_cast(0), g - filter_phase_len + 1); - index_t sample_idx = g * NUM_CHAN + NUM_CHAN - 1; - index_t j_start = g; - index_t h_ind { 0 }; - index_t niter = j_start - first_ind + 1; - // For the last signal element, we need bounds-checking because we may need to zero-pad the signal. - if (niter > 0) { - for (int chan = 0; chan < NUM_CHAN; chan++) { - const filter_t h_val = (h_ind < filter_full_len) ? filter_acc(h_ind) : static_cast(0); - if (sample_idx < input_len) { - detail::channelize_cmac(accum[chan], - detail::channelize_cast_filter(h_val), - detail::channelize_cast_input(input_b(sample_idx))); + } + + const index_t store_end = ComplexInput && + cuda::std::is_same_v + ? cuda::std::min(output_len, block_start + block_rows) : output_len; + ChannelizePolyFusedM2Store( + output_b, store_end, block_start, tid, a0r, a0i, a1r, a1i); +} + +// Two-channel, decimation-one leaf of the fused radix/power-of-two family. +// Each output alternates the two filter phases while advancing by one raw +// input sample. Staging a contiguous raw-input window avoids launching the +// mostly-idle generic channel-tiled kernel for only two active channels. +template +__device__ __forceinline__ void ChannelizePolyFusedM2D1Body( + OutType &output, const InType &input, const FilterType &filter, index_t out_elem_offset, + uint8_t *smem_raw) +{ + constexpr int BLOCK_OUTPUTS = THREADS * OUTPUTS_PER_THREAD; + constexpr int NUM_CHAN = 2; + using traits = ChannelizePolyFusedSmallTraits; + using input_t = typename traits::input_t; + using scalar_t = AccumType; + using filter_pair_t = detail::ChannelizePolyM2Pair; + constexpr bool ComplexInput = traits::ComplexInput; + using shared_input_t = input_t; + + constexpr int InRank = InType::Rank(); + constexpr int OutRank = OutType::Rank(); + constexpr int OutElemRank = OutRank - 2; + const index_t input_len = input.Size(InRank - 1); + const index_t output_len = output.Size(OutElemRank); + const index_t filter_len = filter.Size(0); + const int32_t P = static_cast((filter_len + 1) / NUM_CHAN); + + const detail::cpoly::FusedRadixSmemLayout< + NUM_CHAN, false, input_t, scalar_t, AccumType> layout(P, 1); + filter_pair_t *smem_filter = reinterpret_cast(smem_raw); + shared_input_t *smem_input = reinterpret_cast(smem_raw + layout.input_offset); + + auto [input_b, output_b, filter_acc] = + detail::cpoly::MakeAccessors(output, input, filter, blockIdx.z); + + const int32_t tid = static_cast(threadIdx.x); + ChannelizePolyFusedM2LoadFilter(filter_acc, smem_filter, P, filter_len, tid); + + const index_t block_start = static_cast(blockIdx.x) * BLOCK_OUTPUTS; + const index_t input_base = block_start + out_elem_offset - 2 * P; + const int32_t staged_inputs = static_cast(layout.input_elements); + for (int32_t i = tid; i < staged_inputs; i += THREADS) { + smem_input[i] = detail::cpoly::LoadRelativeInput(input_b, input_base, input_len, i); + } + __syncthreads(); + + scalar_t a0r[OUTPUTS_PER_THREAD]{}; + scalar_t a0i[OUTPUTS_PER_THREAD]{}; + scalar_t a1r[OUTPUTS_PER_THREAD]{}; + scalar_t a1i[OUTPUTS_PER_THREAD]{}; + for (int32_t p = 0; p < P; p++) { + const filter_pair_t hv = smem_filter[p]; + #pragma unroll + for (int q = 0; q < OUTPUTS_PER_THREAD; q++) { + const int32_t local = tid + q * THREADS; + // XOR gives the sum's LSB (1 means odd) without risking overflow. + const bool odd_output = ((block_start ^ local ^ out_elem_offset) & 1) != 0; + const int32_t pair_base = local + 2 * P - 2 * p; + const shared_input_t even = smem_input[pair_base - (odd_output ? 1 : 0)]; + const shared_input_t odd = smem_input[pair_base - (odd_output ? 0 : 1)]; + const scalar_t h0 = odd_output ? hv.odd : hv.even; + const scalar_t h1 = odd_output ? hv.even : hv.odd; + if constexpr (ComplexInput) { + a0r[q] = h0 * even.real() + a0r[q]; + a0i[q] = h0 * even.imag() + a0i[q]; + a1r[q] = h1 * odd.real() + a1r[q]; + a1i[q] = h1 * odd.imag() + a1i[q]; + } else { + a0r[q] = h0 * even + a0r[q]; + a1r[q] = h1 * odd + a1r[q]; + } + } + } + + ChannelizePolyFusedM2Store( + output_b, output_len, block_start, tid, a0r, a0i, a1r, a1i); +} + +// General body for fused channel counts K*2^n, where K is a supported small +// odd radix. Keeping the filtered branch values on chip removes the global- +// memory round trip between the FIR kernel and cuFFT. +template +__device__ __forceinline__ void ChannelizePolyFusedRadixPow2Body( + OutType &output, const InType &input, const FilterType &filter, + IdxT elems_per_channel_per_cta, IdxT decimation_factor, IdxT out_elem_offset, uint8_t *smem_raw) +{ + using config = detail::cpoly::FusedRadixConfig; + constexpr int RADIX_1 = config::Radix1; + constexpr int RADIX_2 = config::Radix2; + constexpr int ROW_STRIDE = config::RowStride; + constexpr int FIR_OUTPUTS_PER_THREAD = config::FirOutputsPerThread; + static_assert(RADIX_1 == 1 || RADIX_1 == 3 || RADIX_1 == 5); + static_assert(ROW_STRIDE <= THREADS); + static_assert(NROWS == config::NRows); + // FFT loops execute complete warps, including padded final output rows. + static_assert(!config::WarpFft || (THREADS % 32 == 0 && (NROWS * RADIX_2) % 32 == 0)); + + using output_t = typename OutType::value_type; + using input_t = typename InType::value_type; + using filter_t = typename FilterType::value_type; + static_assert(!is_complex_v, + "channelize_poly: accumulator type must be real; " + "it will be treated as complex when necessary"); + using filtering_accum_t = cuda::std::conditional_t< + is_complex_v || is_complex_v, + typename detail::scalar_to_complex::ctype, AccumType>; + using complex_accum_t = typename detail::scalar_to_complex::ctype; + + constexpr int InRank = InType::Rank(); + constexpr int OutRank = OutType::Rank(); + constexpr int OutElemRank = OutRank - 2; + + const IdxT input_len = static_cast(input.Size(InRank - 1)); + const IdxT output_len = static_cast(output.Size(OutElemRank)); + const IdxT filter_len = static_cast(filter.Size(0)); + const int32_t P = static_cast((filter_len + NUM_CHAN - 1) / NUM_CHAN); + const detail::cpoly::FusedRadixSmemLayout< + NUM_CHAN, MaximallyDecimated, input_t, filter_t, AccumType> + layout(P, decimation_factor); + const int32_t input_elements = static_cast(layout.input_elements); + const int32_t height = input_elements / NUM_CHAN; + + filter_t *smem_filter = reinterpret_cast(smem_raw); + input_t *smem_input = reinterpret_cast(smem_raw + layout.input_offset); + complex_accum_t *smem_work = reinterpret_cast(smem_raw + layout.work_offset); + complex_accum_t *smem_stage1 = + reinterpret_cast(smem_raw + layout.stage1_offset); + complex_accum_t *twiddle_cross = reinterpret_cast( + smem_raw + layout.twiddle_cross_offset); + complex_accum_t *twiddle_radix2 = reinterpret_cast( + smem_raw + layout.twiddle_radix2_offset); + + auto [input_b, output_b, filter_acc] = + detail::cpoly::MakeAccessors(output, input, filter, blockIdx.z); + + const int32_t tid = static_cast(threadIdx.x); + const int32_t fir_group = tid / ROW_STRIDE; + const int32_t channel = tid % ROW_STRIDE; + const bool active = channel < NUM_CHAN; + + const int32_t filter_elements = static_cast(layout.filter_elements); + for (int32_t i = tid; i < filter_elements; i += THREADS) { + smem_filter[i] = (i < filter_len) + ? filter_acc(static_cast(i)) : static_cast(0); + } + if constexpr (RADIX_1 > 1) { + for (int32_t i = tid; i < RADIX_1 * RADIX_2; i += THREADS) { + const int32_t k = i / RADIX_2; + const int32_t n = i % RADIX_2; + twiddle_cross[i] = ChannelizePolyTwiddle(k * n, NUM_CHAN); + } + } + if constexpr (RADIX_2 > 1) { + if (tid < RADIX_2 / 2) { + const auto twiddle = ChannelizePolyRadix2Twiddle(tid); + // Smaller stages reuse a subset of the largest stage's roots. + #pragma unroll + for (int32_t len = 2; len <= RADIX_2; len *= 2) { + const int32_t stride = RADIX_2 / len; + if (tid % stride == 0) { + twiddle_radix2[len + tid / stride] = twiddle; } - h_ind++; - sample_idx--; } } - niter--; + } + + const IdxT start_elem = static_cast(blockIdx.x) * elems_per_channel_per_cta; + const IdxT last_elem = cuda::std::min( + output_len - 1, start_elem + elems_per_channel_per_cta - 1); + if constexpr (MaximallyDecimated) { + const IdxT first_global_row = start_elem + out_elem_offset - (P - 1); + const IdxT input_base = first_global_row * NUM_CHAN; + for (int32_t linear = tid; + linear < input_elements; linear += THREADS) { + int32_t destination = linear; + if constexpr (!config::WarpFft) { + const int32_t row_offset = linear / NUM_CHAN; + const int32_t col = linear % NUM_CHAN; + const IdxT global_row = first_global_row + row_offset; + int32_t smem_row = static_cast(global_row % height); + if (smem_row < 0) smem_row += height; + destination = smem_row * NUM_CHAN + col; + } + smem_input[destination] = detail::cpoly::LoadRelativeInput( + input_b, input_base, input_len, linear); + } + } + __syncthreads(); - // The central elements require no bounds checking on the filter or signal. - for (index_t i = 0; i < niter-1; i++) { - for (int chan = 0; chan < NUM_CHAN; chan++) { - const filter_t h_val = filter_acc(h_ind); - detail::channelize_cmac(accum[chan], - detail::channelize_cast_filter(h_val), - detail::channelize_cast_input(input_b(sample_idx))); - h_ind++; - sample_idx--; + const IdxT last_start = start_elem + ((last_elem - start_elem) / NROWS) * NROWS; + int32_t newest_row = P + (fir_group + 1) * FIR_OUTPUTS_PER_THREAD - 2; + int32_t reload_row = 0; + for (IdxT next_start = start_elem; + next_start <= last_start; next_start += NROWS) { + IdxT first_raw = 0; + if constexpr (!MaximallyDecimated) { + const IdxT input_base = (next_start + out_elem_offset) * decimation_factor; + const int32_t first_offset = static_cast(decimation_factor) - 1 - P * NUM_CHAN; + first_raw = input_base + first_offset; + for (int32_t i = tid; i < input_elements; i += THREADS) { + smem_input[i] = detail::cpoly::LoadRelativeInput(input_b, first_raw, input_len, i); } + __syncthreads(); } - // For the first signal element / last filter tap, we need to bounds check the filter. - if (niter > 0) { - for (int chan = 0; chan < NUM_CHAN; chan++) { - if (h_ind >= filter_full_len) { - break; + filtering_accum_t filtered[FIR_OUTPUTS_PER_THREAD]{}; + bool valid[FIR_OUTPUTS_PER_THREAD]; + #pragma unroll + for (int q = 0; q < FIR_OUTPUTS_PER_THREAD; q++) { + const IdxT t = next_start + fir_group * FIR_OUTPUTS_PER_THREAD + q; + valid[q] = active && t <= last_elem; + } + + if constexpr (MaximallyDecimated) { + if (active) { + int32_t sample_row = newest_row; + if constexpr (!config::WarpFft) { + const IdxT first_t = next_start + fir_group * FIR_OUTPUTS_PER_THREAD; + sample_row = static_cast( + (first_t + out_elem_offset + FIR_OUTPUTS_PER_THREAD - 1) % height); + } + for (int32_t r = 0; + r < P + FIR_OUTPUTS_PER_THREAD - 1; r++) { + const input_t iv = smem_input[sample_row * NUM_CHAN + (NUM_CHAN - 1 - channel)]; + const auto iav = detail::channelize_cast_operand(iv); + #pragma unroll + for (int q = 0; q < FIR_OUTPUTS_PER_THREAD; q++) { + const int32_t p = q + r - (FIR_OUTPUTS_PER_THREAD - 1); + if (valid[q] && p >= 0 && p < P) { + detail::channelize_cmac( + filtered[q], detail::channelize_cast_operand( + smem_filter[p * NUM_CHAN + channel]), + iav); + } + } + if (--sample_row < 0) sample_row += height; + } + } + } else if (active) { + #pragma unroll + for (int q = 0; q < FIR_OUTPUTS_PER_THREAD; q++) { + if (valid[q]) { + const IdxT t = next_start + + fir_group * FIR_OUTPUTS_PER_THREAD + q + out_elem_offset; + const IdxT last_arrived = t * decimation_factor + decimation_factor - 1; + const int32_t remapped = + (channel + NUM_CHAN - + static_cast(decimation_factor)) % NUM_CHAN; + const int32_t branch = NUM_CHAN - 1 - remapped; + if (last_arrived >= branch) { + const IdxT delta = last_arrived - branch; + const IdxT newest = last_arrived - delta % NUM_CHAN; + const int32_t phase = static_cast( + (channel + t * decimation_factor) % NUM_CHAN); + for (int32_t p = 0; p < P; p++) { + const IdxT sample = newest - static_cast(p) * NUM_CHAN; + const input_t iv = smem_input[sample - first_raw]; + detail::channelize_cmac( + filtered[q], detail::channelize_cast_operand( + smem_filter[p * NUM_CHAN + phase]), + detail::channelize_cast_operand(iv)); + } + } } - const filter_t h_val = filter_acc(h_ind); - detail::channelize_cmac(accum[chan], - detail::channelize_cast_filter(h_val), - detail::channelize_cast_input(input_b(sample_idx))); - h_ind++; - sample_idx--; } } - - // For complex inputs, the DFT will not generally be conjugate symmetric, so compute all - // terms. For real inputs, we only compute the unique (up to conjugate symmetry) components. - if constexpr (is_complex_v || is_complex_half_v) { - for (int chan = 0; chan < NUM_CHAN; chan++) { - complex_accum_t dft { 0 }; - for (int j = 0; j < NUM_CHAN; j++) { - dft += accum[j] * smem_eij[chan][j]; + if (active) { + int32_t fft_channel = channel; + if constexpr (RADIX_1 == 1) { + // Store pure-power-of-two FIR branches in FFT input order. + // No radix-K stage or cross twiddle is needed in this case. + fft_channel = detail::cpoly::BitReverse(channel); + } + #pragma unroll + for (int q = 0; q < FIR_OUTPUTS_PER_THREAD; q++) { + const int32_t row = fir_group * FIR_OUTPUTS_PER_THREAD + q; + if constexpr (is_complex_v) { + smem_work[row * NUM_CHAN + fft_channel] = + static_cast(filtered[q]); + } else { + smem_work[row * NUM_CHAN + fft_channel] = { + static_cast(filtered[q]), static_cast(0)}; } - output_b(t, static_cast(chan)) = static_cast(dft); } - } else { - constexpr int mid = NUM_CHAN/2 + 1; - if constexpr (NUM_CHAN % 2 == 0) { - // Channel 0, DC. There is no conjugate symmetric component for this value. - { - complex_accum_t dft { 0 }; - for (int j = 0; j < NUM_CHAN; j++) { - dft += accum[j] * smem_eij[0][j]; + } + __syncthreads(); + if constexpr (config::WarpFft) { + for (int32_t logical = tid; + logical < NROWS * RADIX_2; logical += THREADS) { + const int32_t row = logical / RADIX_2; + const int32_t lane = logical % RADIX_2; + const int32_t bit_reversed = detail::cpoly::BitReverse(lane); + complex_accum_t value[RADIX_1]; + if constexpr (RADIX_1 == 1) { + // Pure-power-of-two FIR stores are already bit reversed. + value[0] = smem_work[row * NUM_CHAN + lane]; + } else { + complex_accum_t branch[RADIX_1]; + #pragma unroll + for (int k = 0; k < RADIX_1; ++k) { + branch[k] = smem_work[row * NUM_CHAN + bit_reversed + k * RADIX_2]; + } + ChannelizePolySmallDFT(branch, value); + #pragma unroll + for (int k = 1; k < RADIX_1; ++k) { + complex_accum_t twiddled{}; + detail::channelize_cmac(twiddled, + twiddle_cross[k * RADIX_2 + bit_reversed], value[k]); + value[k] = twiddled; } - output_b(t, static_cast(0)) = static_cast(dft); } - // Channel mid-1, Nyquist. There is no conjugate symmetric component for this value. - { - complex_accum_t dft { 0 }; - for (int j = 0; j < NUM_CHAN; j++) { - dft += accum[j] * smem_eij[mid-1][j]; + constexpr unsigned mask = 0xffffffffU; + #pragma unroll + for (int len = 2; len <= RADIX_2; len *= 2) { + const int j = lane % (len / 2); + const bool lower = (lane & (len / 2)) == 0; + const complex_accum_t twiddle = twiddle_radix2[len + j]; + #pragma unroll + for (int k = 0; k < RADIX_1; ++k) { + const complex_accum_t partner{ + __shfl_xor_sync(mask, value[k].real(), len / 2, RADIX_2), + __shfl_xor_sync(mask, value[k].imag(), len / 2, RADIX_2)}; + const complex_accum_t lo = lower ? value[k] : partner; + const complex_accum_t hi = lower ? partner : value[k]; + complex_accum_t product{}; + detail::channelize_cmac(product, twiddle, hi); + value[k] = lower ? lo + product : lo - product; } - output_b(t, static_cast(mid-1)) = static_cast(dft); } - // Conjugate symmetric components - for (int chan = 1; chan < mid-1; chan++) { - complex_accum_t dft { 0 }; - for (int j = 0; j < NUM_CHAN; j++) { - dft += accum[j] * smem_eij[chan][j]; + if constexpr (RADIX_1 > 1) { + // Warp shuffles do not order shared loads before reuse. + __syncwarp(mask); + } + #pragma unroll + for (int k = 0; k < RADIX_1; ++k) { + if constexpr (RADIX_1 == 1) { + const IdxT t = next_start + row; + if (t <= last_elem) { + output_b(t, static_cast(lane)) = static_cast(value[k]); + } + } else { + // Transpose the register-owned odd-radix outputs for + // coalesced global stores across all CTA threads. + smem_work[row * NUM_CHAN + lane * RADIX_1 + k] = value[k]; } - output_b(t, static_cast(chan)) = static_cast(dft); - output_b(t, static_cast(NUM_CHAN - chan)) = static_cast(conj(dft)); } - } else { - // Channel 0, DC. There is no conjugate symmetric component for this value. - { - complex_accum_t dft { 0 }; - for (int j = 0; j < NUM_CHAN; j++) { - dft += accum[j] * smem_eij[0][j]; + } + if constexpr (RADIX_1 > 1) { + __syncthreads(); + } + } else { + if constexpr (RADIX_1 > 1) { + for (int32_t logical = tid; + logical < NROWS * RADIX_2; logical += THREADS) { + const int32_t row = logical / RADIX_2; + const int32_t n1 = logical % RADIX_2; + const int32_t base = row * NUM_CHAN + n1; + const int32_t stage_base = row * NUM_CHAN + n1 * RADIX_1; + // Array store loops increased register/stack usage on L4. + if constexpr (RADIX_1 == 3) { + const complex_accum_t x0 = smem_work[base]; + const complex_accum_t x1 = smem_work[base + RADIX_2]; + const complex_accum_t x2 = smem_work[base + 2 * RADIX_2]; + complex_accum_t y0, y1, y2; + ChannelizePolyRadix3(x0, x1, x2, y0, y1, y2); + smem_stage1[stage_base] = y0; + smem_stage1[stage_base + 1] = y1; + smem_stage1[stage_base + 2] = y2; + } else { + const complex_accum_t x0 = smem_work[base]; + const complex_accum_t x1 = smem_work[base + RADIX_2]; + const complex_accum_t x2 = smem_work[base + 2 * RADIX_2]; + const complex_accum_t x3 = smem_work[base + 3 * RADIX_2]; + const complex_accum_t x4 = smem_work[base + 4 * RADIX_2]; + complex_accum_t y0, y1, y2, y3, y4; + ChannelizePolyRadix5(x0, x1, x2, x3, x4, y0, y1, y2, y3, y4); + smem_stage1[stage_base] = y0; + smem_stage1[stage_base + 1] = y1; + smem_stage1[stage_base + 2] = y2; + smem_stage1[stage_base + 3] = y3; + smem_stage1[stage_base + 4] = y4; } - output_b(t, static_cast(0)) = static_cast(dft); } - // Conjugate symmetric components - for (int chan = 1; chan < mid; chan++) { - complex_accum_t dft { 0 }; - for (int j = 0; j < NUM_CHAN; j++) { - dft += accum[j] * smem_eij[chan][j]; + __syncthreads(); + + for (int32_t logical = tid; + logical < NROWS * NUM_CHAN; logical += THREADS) { + const int32_t row = logical / NUM_CHAN; + const int32_t fft_channel = logical % NUM_CHAN; + const int32_t n1 = fft_channel / RADIX_1; + const int32_t k2 = fft_channel % RADIX_1; + const int32_t bit_reversed = detail::cpoly::BitReverse(n1); + complex_accum_t value{}; + detail::channelize_cmac( + value, twiddle_cross[k2 * RADIX_2 + n1], + smem_stage1[row * NUM_CHAN + n1 * RADIX_1 + k2]); + smem_work[row * NUM_CHAN + bit_reversed * RADIX_1 + k2] = value; + } + __syncthreads(); + } + + for (int32_t len = 2; len <= RADIX_2; len *= 2) { + for (int32_t logical = tid; + logical < NROWS * NUM_CHAN; logical += THREADS) { + const int32_t row = logical / NUM_CHAN; + const int32_t fft_channel = logical % NUM_CHAN; + const int32_t k1 = fft_channel / RADIX_1; + const int32_t k2 = fft_channel % RADIX_1; + const int32_t j = k1 % len; + if (j < len / 2) { + const int32_t base = k1 - j; + const int32_t lo = row * NUM_CHAN + (base + j) * RADIX_1 + k2; + const int32_t hi = lo + (len / 2) * RADIX_1; + const complex_accum_t u = smem_work[lo]; + complex_accum_t v{}; + detail::channelize_cmac(v, twiddle_radix2[len + j], smem_work[hi]); + smem_work[lo] = u + v; + smem_work[hi] = u - v; } - output_b(t, static_cast(chan)) = static_cast(dft); - output_b(t, static_cast(NUM_CHAN - chan)) = static_cast(conj(dft)); } + __syncthreads(); } } + if constexpr (!config::WarpFft || RADIX_1 > 1) { + for (int32_t logical = tid; + logical < NROWS * NUM_CHAN; logical += THREADS) { + const int32_t row = logical / NUM_CHAN; + const int32_t fft_channel = logical % NUM_CHAN; + const IdxT t = next_start + row; + if (t <= last_elem) { + output_b(t, static_cast(fft_channel)) = + static_cast(smem_work[logical]); + } + } + } + __syncthreads(); + + if constexpr (MaximallyDecimated) { + if (next_start < last_start) { + const IdxT first_global_row = next_start + out_elem_offset + NROWS; + const IdxT input_base = first_global_row * NUM_CHAN; + for (int32_t logical = tid; + logical < NROWS * NUM_CHAN; logical += THREADS) { + const int32_t row = logical / NUM_CHAN; + const int32_t col = logical % NUM_CHAN; + int32_t smem_row = reload_row + row; + if constexpr (config::WarpFft) { + if (smem_row >= height) smem_row -= height; + } else { + const IdxT global_row = first_global_row + row; + smem_row = static_cast(global_row % height); + } + smem_input[smem_row * NUM_CHAN + col] = + detail::cpoly::LoadRelativeInput( + input_b, input_base, input_len, logical); + } + if constexpr (config::WarpFft) { + newest_row += NROWS; + if (newest_row >= height) newest_row -= height; + reload_row += NROWS; + if (reload_row >= height) reload_row -= height; + } + __syncthreads(); + } + } } } +// One source-level kernel family covers row-owned small channel counts, pure +// powers of two, and supported K*2^n channel counts. Runtime dispatch +// deliberately instantiates only a small supported set of NUM_CHAN values. +template +__device__ __forceinline__ void ChannelizePolyFusedRadixBody( + OutType &output, const InType &input, const FilterType &filter, + IdxT elems_per_channel_per_cta, IdxT decimation_factor, IdxT out_elem_offset) +{ + extern __shared__ __align__(16) uint8_t smem_raw[]; + using config = detail::cpoly::FusedRadixConfig< + NUM_CHAN, AccumType, MaximallyDecimated, typename InType::value_type>; + static_assert(THREADS == config::Threads); + constexpr int OUTPUTS_PER_THREAD = config::SmallOutputsPerThread; + if constexpr (NUM_CHAN <= 6 && MaximallyDecimated) { + using traits = ChannelizePolyFusedSmallTraits< + NUM_CHAN, OutType, InType, FilterType, AccumType>; + if constexpr (NUM_CHAN == 2 && traits::UseM2PairLeaf) { + ChannelizePolyFusedM2D2Body< + THREADS, OUTPUTS_PER_THREAD, IsUnitStride, OutType, InType, FilterType, AccumType>( + output, input, filter, out_elem_offset, elems_per_channel_per_cta, smem_raw); + } else { + const IdxT P = (filter.Size(FilterType::Rank() - 1) + NUM_CHAN - 1) / NUM_CHAN; + if (detail::cpoly::FusedSmallUseDirect< + NUM_CHAN, OutType, InType, FilterType, AccumType>( + P, output.Size(OutType::Rank() - 2))) { + ChannelizePolyFusedSmallDirectBody< + THREADS, NUM_CHAN, OUTPUTS_PER_THREAD, IsUnitStride, + OutType, InType, FilterType, AccumType>( + output, input, filter, out_elem_offset); + } else { + ChannelizePolyFusedSmallCachedBody< + THREADS, NUM_CHAN, OUTPUTS_PER_THREAD, IsUnitStride, + OutType, InType, FilterType, AccumType>( + output, input, filter, out_elem_offset, smem_raw); + } + } + } else if constexpr (NUM_CHAN == 2) { + ChannelizePolyFusedM2D1Body< + THREADS, OUTPUTS_PER_THREAD, IsUnitStride, OutType, InType, FilterType, AccumType>( + output, input, filter, out_elem_offset, smem_raw); + } else { + ChannelizePolyFusedRadixPow2Body< + THREADS, NUM_CHAN, NROWS, MaximallyDecimated, IsUnitStride, + OutType, InType, FilterType, AccumType>( + output, input, filter, elems_per_channel_per_cta, + decimation_factor, out_elem_offset, smem_raw); + } +} + +// Mutually exclusive overloads omit the optional bound instead of passing zero. +// Select index width explicitly, not from the types of integer launch literals. +template ::MinBlocksPerSm : 0> + requires (MIN_BLOCKS == 0) +__launch_bounds__(THREADS) +__global__ void ChannelizePoly1D_FusedRadixPow2( + OutType output, InType input, FilterType filter, + cuda::std::type_identity_t elems_per_channel_per_cta, + cuda::std::type_identity_t decimation_factor, + cuda::std::type_identity_t out_elem_offset) +{ + ChannelizePolyFusedRadixBody( + output, input, filter, elems_per_channel_per_cta, decimation_factor, out_elem_offset); +} + +template ::MinBlocksPerSm : 0> + requires (MIN_BLOCKS > 0) +__launch_bounds__(THREADS, MIN_BLOCKS) +__global__ void ChannelizePoly1D_FusedRadixPow2( + OutType output, InType input, FilterType filter, + cuda::std::type_identity_t elems_per_channel_per_cta, + cuda::std::type_identity_t decimation_factor, + cuda::std::type_identity_t out_elem_offset) +{ + ChannelizePolyFusedRadixBody( + output, input, filter, elems_per_channel_per_cta, decimation_factor, out_elem_offset); +} + // Unpack the compressed representation of the spectrum after a real-to-complex FFT. // Because the input was real, the spectrum is conjugate symmetric and fft() will // return a packed version of the output that includes only the unique elements @@ -1121,27 +2204,27 @@ __global__ void ChannelizePoly1DUnpackDFT(DataType inout) constexpr int ElemRank = Rank-2; using value_t = typename DataType::value_type; - const int tid = blockIdx.y * blockDim.x + threadIdx.x; - const index_t num_elem_per_channel = inout.Size(ElemRank); const index_t num_channels = inout.Size(ChannelRank); const index_t mid = num_channels/2 + 1; - if (tid >= num_elem_per_channel) { - return; - } // Bind batch coords; remaining dims are (elem, channel). Access with - // inout_b(tid, chan). + // inout_b(elem, chan). detail::TensorAccessor inout_acc(inout); const auto batch_idx = BlockToIdx(inout, blockIdx.x, 2); auto inout_b = detail::bind_first_n(inout_acc, batch_idx); const index_t upper = (num_channels % 2 == 0) ? (mid - 1) : mid; - for (index_t i = 1; i < upper; i++) { - const value_t val = inout_b(static_cast(tid), i); - inout_b(static_cast(tid), i) = conj(val); - inout_b(static_cast(tid), num_channels - i) = val; + // The launch caps grid.y at the grid.y limit; stride over the remainder. + const index_t stride = static_cast(gridDim.y) * blockDim.x; + for (index_t elem = static_cast(blockIdx.y) * blockDim.x + threadIdx.x; + elem < num_elem_per_channel; elem += stride) { + for (index_t i = 1; i < upper; i++) { + const value_t val = inout_b(elem, i); + inout_b(elem, i) = conj(val); + inout_b(elem, num_channels - i) = val; + } } } diff --git a/include/matx/transforms/channelize_poly.h b/include/matx/transforms/channelize_poly.h index 6f4318604..682f10e19 100644 --- a/include/matx/transforms/channelize_poly.h +++ b/include/matx/transforms/channelize_poly.h @@ -53,11 +53,6 @@ namespace matx { namespace detail { namespace cpoly { -// Any channel count at or below FusedChanThreshold will -// use the fused-channel kernel. If this is increased, then the switch statement -// in the fused kernel wrapper below must be adjusted to include the additional -// channel counts. -constexpr index_t FusedChanThreshold = 6; // Number of output samples per channel per iteration for the kernel that stores // the input data in shared memory. Ideally, this value would be determined dynamically @@ -67,23 +62,58 @@ constexpr index_t FullSmemKernelNoutPerIter = 4; // Maximum dynamic shared memory (bytes) for caching filter taps in ChannelizePoly1D. constexpr size_t GenericMaxFilterSmemBytes = 6 * 1024; -// Constants for the SmemTiled kernel. Two CTILE sizes are instantiated: -// CTILE=32 for num_channels <= 32: single channel tile, 100% thread -// utilization. Block size: 32 * NOUT = 128 threads. -// CTILE=64 for num_channels > 32: original shape. Block size: 64 * NOUT -// = 256 threads. -// NOUT=4 across both variants; NOUT=8 for the CTILE=32 path was evaluated -// and was ~5% slower (doubles block size without increasing elem_per_block, -// so fewer blocks/SM and no occupancy win). -constexpr int SmemTiledCtile = 64; -constexpr int SmemTiledCtileSmall = 32; -constexpr int SmemTiledNout = 4; -constexpr size_t SmemTiledMaxBytes = 48 * 1024; -// Maximum filter smem budget. A filter is only considered for smem -// residency if its footprint (in whichever layout we're considering, Full -// or Rotated) fits under this cap. Raising this costs occupancy by -// inflating per-block smem, so keep it small. -constexpr size_t SmemTiledMaxFilterBytes = 4096; +// Constants for the tiled shared-memory kernel, which processes CTILE channels +// per CTA. NOUT=16 is the maximally decimated tile: four time lanes each +// compute four rows, reusing every filter tap across those rows. NOUT=4 is the +// oversampled tile and the critical fallback when the grouped ring does not fit. +constexpr int SmemTiledCtile = 32; +constexpr int SmemTiledNout = 4; +constexpr int SmemTiledMaxDecNout = 16; + +// Dispatch policy. Up to this many channels, a FIR that is not fused uses the +// whole-channel Smem kernel (critical) or the Generic kernel (oversampled); +// wider channel counts use the channel-tiled kernel. +constexpr index_t SmallChannelLimit = 16; +// Oversampled tiles cache the Full filter layout only up to this size. Every CTA stages the whole +// filter, so larger filters cost more to load than they save and reduce occupancy. +constexpr size_t SmemTiledOversampledFilterBytes = 4 * 1024; + +// Memory bus width per SM at or above this marks an HBM-class GPU (roughly +// 40-55 bits per SM, versus about 3-16 for GDDR and LPDDR parts). +constexpr int HighBandwidthBusBitsPerSm = 32; + +// Device attributes used by dispatch and launch sizing. Each is queried at most +// once per channelizer call, and only when a decision needs it. +class DeviceAttrs { +public: + int SmCount() { return Get(sms_, cudaDevAttrMultiProcessorCount); } + size_t L2Bytes() + { + return static_cast(Get(l2_bytes_, cudaDevAttrL2CacheSize)); + } + // GPUs whose FP64 rate is far below FP32 make FP64 channelizers arithmetic + // bound, so padded or idle lanes cost time directly. + bool LowFp64Throughput() + { + return Get(fp64_ratio_, cudaDevAttrSingleToDoublePrecisionPerfRatio) > 8; + } + // Bus width is a cheap proxy for bandwidth; memory clock queries are slow. + bool HighMemoryBandwidth() + { + return Get(bus_bits_, cudaDevAttrGlobalMemoryBusWidth) >= HighBandwidthBusBitsPerSm * SmCount(); + } + +private: + static int Get(int &value, cudaDeviceAttr attr) + { + if (value < 0) value = GetDeviceAttr(attr); + return value; + } + int sms_ = -1; + int l2_bytes_ = -1; + int fp64_ratio_ = -1; + int bus_bits_ = -1; +}; // Filter-smem placement strategy chosen by the dispatcher and forwarded to // the kernel via template parameters. @@ -96,14 +126,67 @@ constexpr size_t SmemTiledMaxFilterBytes = 4096; // Global: filter stays in GMEM, loaded through L1 in the inner loop. enum class SmemTiledFilterLayout { Full, Rotated, Global }; -// Pick the CTILE to use for this launch: 32 when num_channels fits in a -// single 32-wide tile, 64 otherwise. Centralized so dispatch, sizing, and -// launch agree. -constexpr int SmemTiledCtileFor(index_t num_channels) +// A tiled launch computes SmemTiledCtile channels by nout output rows per +// iteration: SmemTiledMaxDecNout (critical only) or SmemTiledNout. +struct SmemTiledPlan { + int nout; + SmemTiledFilterLayout filter_layout; + size_t bytes; +}; + +// Tiled grids target a number of CTAs per SM so they scale with the GPU. +constexpr int SmemTiledCtasPerSm = 8; + +constexpr int SmemTiledChannelTiles(index_t num_channels) +{ + return static_cast((num_channels + SmemTiledCtile - 1) / SmemTiledCtile); +} + +// A tiled plan launches when its shared memory fits the default limit and its +// channel tiles fit the grid.y limit. +inline bool SmemTiledFits(const SmemTiledPlan &plan, index_t num_channels) +{ + return plan.bytes <= MaxDefaultDynamicSmemBytes && + SmemTiledChannelTiles(num_channels) <= 65535; +} + +// Rows per CTA for the tiled kernel. A single batch spreads target_ctas over +// its channel tiles; batches share the grid, bounding the rows each CTA adds. +inline index_t SmemTiledElemsPerBlock( + index_t nout_per_channel, index_t num_channels, index_t batches, index_t target_ctas) +{ + const index_t channel_tiles = SmemTiledChannelTiles(num_channels); + const index_t time_blocks = (target_ctas + channel_tiles - 1) / channel_tiles; + const index_t single = (nout_per_channel + time_blocks - 1) / time_blocks; + const index_t spatial_batches = std::max(1, channel_tiles * batches); + const index_t time_targets = std::max(1, + (2 * target_ctas + spatial_batches - 1) / spatial_batches); + const index_t span = (nout_per_channel + time_targets - 1) / time_targets; + return std::max(single, std::min(span, 128)); +} + +// Whether 32-bit indices can address a tiled launch whose window ends before +// global output row window_end. The input ring loads through the end of the row +// holding the window's newest sample, which can pass input_len + M when D < M. +inline bool SmemTiledFitsInt32( + index_t input_len, index_t num_channels, index_t decimation_factor, index_t window_end) { - return (num_channels <= SmemTiledCtileSmall) - ? SmemTiledCtileSmall - : SmemTiledCtile; + constexpr int64_t max_index = std::numeric_limits::max(); + const int64_t newest_row = + (static_cast(window_end) * decimation_factor - 1) / num_channels; + return static_cast(input_len) + num_channels <= max_index && + (newest_row + 1) * num_channels - 1 <= max_index; +} + +// The Rotated layout stores K = M / gcd(M, D) phases per tile channel; the +// kernel supports at most SmemTiledMaxRotations of them. +inline bool SmemTiledRotatedLayoutSupported(index_t num_channels, index_t decimation_factor) +{ + if (decimation_factor == num_channels) { + return true; + } + const index_t gcd_val = std::gcd(num_channels, decimation_factor); + return num_channels / gcd_val <= SmemTiledMaxRotations; } // Compute the shared memory footprint for the SmemTiled filter taps in the @@ -111,23 +194,18 @@ constexpr int SmemTiledCtileFor(index_t num_channels) // across channels, zero-padded when CTILE > M). // For D==M: one phase per tile channel -> CTILE * P elements. // For D CTILE * K * P elements. -// Returns a value exceeding the filter budget when K exceeds the rotation -// limit, which causes FilterInSmem=false in the dispatch. template inline size_t SmemTiledFilterBytesRotated( - index_t num_channels, index_t filter_len, index_t decimation_factor, int ctile) + index_t num_channels, index_t filter_len, index_t decimation_factor) { using filter_t = typename FilterType::value_type; const index_t P = (filter_len + num_channels - 1) / num_channels; if (decimation_factor == num_channels) { - return static_cast(ctile) * P * sizeof(filter_t); + return static_cast(SmemTiledCtile) * P * sizeof(filter_t); } const index_t gcd_val = std::gcd(num_channels, decimation_factor); const index_t K = num_channels / gcd_val; - if (K > SmemTiledMaxRotations) { - return SmemTiledMaxFilterBytes + 1; // exceeds budget, so FilterInSmem=false - } - return static_cast(ctile) * K * P * sizeof(filter_t); + return static_cast(SmemTiledCtile) * K * P * sizeof(filter_t); } // Shared memory footprint for the Full filter layout: M * P unique taps, @@ -144,165 +222,81 @@ inline size_t SmemTiledFilterBytesFull( return static_cast(num_channels) * P * sizeof(filter_t); } -// Pick a filter-smem layout for this dispatch. Of the candidates that fit -// under the filter budget (Full, Rotated), choose the one with the smaller -// footprint to maximize occupancy. Fall back to Global when neither fits. -// The caller must separately verify that filter + input both fit under -// MAX_BYTES. -template -inline SmemTiledFilterLayout SmemTiledChooseFilterLayout( - const OutType &o, const FilterType &filter, index_t decimation_factor, int ctile) +template +inline size_t SmemTiledInputBytes( + index_t num_channels, index_t filter_len, index_t decimation, int nout) { - const index_t num_channels = o.Size(OutType::Rank() - 1); - const index_t filter_len = filter.Size(FilterType::Rank() - 1); - const size_t full_bytes = SmemTiledFilterBytesFull( - num_channels, filter_len); - const size_t rotated_bytes = SmemTiledFilterBytesRotated( - num_channels, filter_len, decimation_factor, ctile); - const bool full_fits = full_bytes <= SmemTiledMaxFilterBytes; - const bool rotated_fits = rotated_bytes <= SmemTiledMaxFilterBytes; - if (full_fits && rotated_fits) { - return (full_bytes <= rotated_bytes) - ? SmemTiledFilterLayout::Full - : SmemTiledFilterLayout::Rotated; - } - if (full_fits) { - return SmemTiledFilterLayout::Full; - } - if (rotated_fits) { - return SmemTiledFilterLayout::Rotated; - } - return SmemTiledFilterLayout::Global; + const index_t P = (filter_len + num_channels - 1) / num_channels; + return static_cast( + SmemTiledInputHeight(P, num_channels, decimation, nout)) * SmemTiledCtile * + sizeof(typename InType::value_type); } +// Size an explicit tile height and filter layout. The plan is launchable when +// bytes <= MaxDefaultDynamicSmemBytes and, for Rotated, the rotation count is supported. template -inline size_t SmemTiledSizeBytes( - const OutType &o, const InType &, const FilterType &filter, index_t decimation_factor, int ctile) +inline SmemTiledPlan SmemTiledPlanWithLayout( + const OutType &o, const InType &, const FilterType &filter, + index_t decimation_factor, int nout, SmemTiledFilterLayout layout) { - using input_t = typename InType::value_type; - - constexpr int NOUT = SmemTiledNout; - - const index_t M = o.Size(OutType::Rank() - 1); + using input_t = typename InType::value_type; + const index_t num_channels = o.Size(OutType::Rank() - 1); const index_t filter_len = filter.Size(FilterType::Rank() - 1); - const index_t P = (filter_len + M - 1) / M; - const size_t input_smem = static_cast(P + NOUT - 1) * ctile * sizeof(input_t); - - const auto layout = SmemTiledChooseFilterLayout( - o, filter, decimation_factor, ctile); - if (layout == SmemTiledFilterLayout::Global) { - return input_smem; + size_t filter_bytes = 0; + if (layout == SmemTiledFilterLayout::Full) { + filter_bytes = SmemTiledFilterBytesFull(num_channels, filter_len); + } else if (layout == SmemTiledFilterLayout::Rotated) { + filter_bytes = SmemTiledFilterBytesRotated( + num_channels, filter_len, decimation_factor); } - - const size_t filter_smem = (layout == SmemTiledFilterLayout::Full) - ? SmemTiledFilterBytesFull(M, filter_len) - : SmemTiledFilterBytesRotated(M, filter_len, decimation_factor, ctile); - - const size_t filter_smem_aligned = filter_smem + - ((filter_smem % sizeof(input_t)) ? (sizeof(input_t) - filter_smem % sizeof(input_t)) : 0); - return filter_smem_aligned + input_smem; + return {nout, layout, + SmemTiledInputBytes(num_channels, filter_len, decimation_factor, nout) + + MATX_ROUND_UP(filter_bytes, sizeof(input_t))}; } +// Critical tiles cache the smaller of the Full and Rotated layouts that fits, falling back from the +// grouped 32x16 tile to 32x4 when its ring does not fit. Oversampled tiles use 32x4 and cache only +// small Full filters. Global keeps the filter in device memory. The caller checks SmemTiledFits and +// uses Generic when no tile fits. template -inline bool ShouldUseSmemTiled( +inline SmemTiledPlan SelectTiledPlan( const OutType &o, const InType &in, const FilterType &filter, index_t decimation_factor) { const index_t num_channels = o.Size(OutType::Rank() - 1); - const int ctile = SmemTiledCtileFor(num_channels); - // The input circular buffer must fit in smem. The filter may or may not - // be included (FilterInSmem is decided separately). - return SmemTiledSizeBytes(o, in, filter, decimation_factor, ctile) - <= SmemTiledMaxBytes; -} - -template -__MATX_HOST__ __MATX_INLINE__ auto HostChannelizeCastFilter(FilterT v) -{ - if constexpr (is_complex_v) { - return static_cast(v); - } else if constexpr (is_complex_v) { - using accum_scalar_t = typename inner_op_type_t::type; - return static_cast(v); - } else { - return static_cast(v); + auto plan_for = [&](int nout, SmemTiledFilterLayout layout) { + return SmemTiledPlanWithLayout(o, in, filter, decimation_factor, nout, layout); + }; + if (decimation_factor == num_channels) { + SmemTiledPlan global{}; + for (int nout : {SmemTiledMaxDecNout, SmemTiledNout}) { + const auto full = plan_for(nout, SmemTiledFilterLayout::Full); + const auto rotated = plan_for(nout, SmemTiledFilterLayout::Rotated); + const auto &cached = full.bytes <= rotated.bytes ? full : rotated; + if (cached.bytes <= MaxDefaultDynamicSmemBytes) return cached; + global = plan_for(nout, SmemTiledFilterLayout::Global); + if (global.bytes <= MaxDefaultDynamicSmemBytes) return global; + } + return global; } -} - -template -__MATX_HOST__ __MATX_INLINE__ auto HostChannelizeCastInput(InputT v) -{ - if constexpr (is_complex_v) { - return static_cast(v); - } else if constexpr (is_complex_v) { - using accum_scalar_t = typename inner_op_type_t::type; - return static_cast(v); - } else { - return static_cast(v); + const auto full = plan_for(SmemTiledNout, SmemTiledFilterLayout::Full); + const size_t filter_bytes = SmemTiledFilterBytesFull( + num_channels, filter.Size(FilterType::Rank() - 1)); + if (filter_bytes <= SmemTiledOversampledFilterBytes && full.bytes <= MaxDefaultDynamicSmemBytes) { + return full; } + return plan_for(SmemTiledNout, SmemTiledFilterLayout::Global); } -template -__MATX_HOST__ __MATX_INLINE__ void HostChannelizeCmac( - AccumT &accum, FilterValT hv, InputValT iv) +template +inline void DispatchUnitStride(Launch &&launch, const Ops &...ops) { - if constexpr (is_complex_v && is_complex_v && is_complex_v) { - auto h_re = hv.real(), h_im = hv.imag(); - auto i_re = iv.real(), i_im = iv.imag(); - auto a_re = accum.real(), a_im = accum.imag(); - a_re = h_re * i_re + a_re; - a_re = -(h_im * i_im) + a_re; - a_im = h_re * i_im + a_im; - a_im = h_im * i_re + a_im; - accum = {a_re, a_im}; - } else if constexpr (is_complex_v && !is_complex_v && is_complex_v) { - auto a_re = accum.real(), a_im = accum.imag(); - a_re = hv * iv.real() + a_re; - a_im = hv * iv.imag() + a_im; - accum = {a_re, a_im}; - } else if constexpr (is_complex_v && is_complex_v && !is_complex_v) { - auto a_re = accum.real(), a_im = accum.imag(); - a_re = hv.real() * iv + a_re; - a_im = hv.imag() * iv + a_im; - accum = {a_re, a_im}; - } else { - accum += hv * iv; + if constexpr ((is_tensor_view_v && ...)) { + if (((ops.Stride(Ops::Rank() - 1) == 1) && ...)) { + launch(cuda::std::bool_constant{}); + return; + } } -} - -template -__MATX_HOST__ __MATX_INLINE__ decltype(auto) HostReadSignalImpl( - const Op &op, const Arr &batch_idx, index_t sample_idx, - cuda::std::index_sequence) -{ - return op(batch_idx[Is]..., sample_idx); -} - -template -__MATX_HOST__ __MATX_INLINE__ decltype(auto) HostReadSignal( - const Op &op, const Arr &batch_idx, index_t sample_idx) -{ - return HostReadSignalImpl( - op, batch_idx, sample_idx, - cuda::std::make_index_sequence(Op::Rank() - 1)>{}); -} - -template -__MATX_HOST__ __MATX_INLINE__ void HostWriteOutputImpl( - OutType &out, const Arr &batch_idx, index_t output_idx, index_t channel, - const ValueT &value, cuda::std::index_sequence) -{ - out(batch_idx[Is]..., output_idx, channel) = - static_cast(value); -} - -template -__MATX_HOST__ __MATX_INLINE__ void HostWriteOutput( - OutType &out, const Arr &batch_idx, index_t output_idx, index_t channel, - const ValueT &value) -{ - HostWriteOutputImpl( - out, batch_idx, output_idx, channel, value, - cuda::std::make_index_sequence(OutType::Rank() - 2)>{}); + launch(cuda::std::bool_constant{}); } template @@ -328,15 +322,59 @@ __MATX_HOST__ __MATX_INLINE__ ComplexAccumT HostTwiddle(index_t channel, index_t static_cast(std::sin(arg))}; } -template +template +inline void ValidateChannelizePolyArgs( + [[maybe_unused]] const OutType &out, const InType &in, const FilterType &, + [[maybe_unused]] index_t num_channels, index_t decimation_factor, + [[maybe_unused]] index_t out_elem_offset) +{ + using output_t = typename OutType::value_type; + constexpr int IN_RANK = InType::Rank(); + constexpr int OUT_RANK = OutType::Rank(); + + static_assert(!is_complex_v, + "channelize_poly: accumulator type must be real; " + "it will be treated as complex when necessary"); + MATX_STATIC_ASSERT_STR(OUT_RANK == IN_RANK + 1, matxInvalidDim, + "channelize_poly: output rank should be 1 higher than input"); + MATX_STATIC_ASSERT_STR(is_complex_v, matxInvalidType, + "channelize_poly: output type must be complex"); + MATX_STATIC_ASSERT_STR(FilterType::Rank() == 1, matxInvalidDim, + "channelize_poly: currently only support 1D filters"); + + MATX_ASSERT_STR(num_channels > 0, matxInvalidParameter, + "channelize_poly: num_channels must be positive"); + MATX_ASSERT_STR(decimation_factor > 0, matxInvalidParameter, + "channelize_poly: decimation_factor must be positive"); + MATX_ASSERT_STR(decimation_factor <= num_channels, matxInvalidParameter, + "channelize_poly: decimation_factor must be <= num_channels"); + + for (int i = 0; i < IN_RANK - 1; i++) { + MATX_ASSERT_STR(out.Size(i) == in.Size(i), matxInvalidDim, + "channelize_poly: input/output must have matched batch sizes"); + } + + [[maybe_unused]] const index_t num_elem_per_channel = + (in.Size(IN_RANK - 1) + decimation_factor - 1) / decimation_factor; + MATX_ASSERT_STR(out.Size(OUT_RANK - 1) == num_channels, matxInvalidDim, + "channelize_poly: output size OUT_RANK-1 mismatch"); + // A window may cover any part of the full output-element grid, including a + // prefix beginning at zero that is shorter than the full grid. + MATX_ASSERT_STR( + out.Size(OUT_RANK - 2) + out_elem_offset <= num_elem_per_channel, matxInvalidDim, + "channelize_poly: output-element window exceeds the full output size"); +} + +template inline void SmemTiledImpl( OutType o, const InType &i, const FilterType &filter, - index_t decimation_factor, cudaStream_t stream, index_t out_elem_offset = 0) + index_t decimation_factor, const SmemTiledPlan &plan, + cudaStream_t stream, index_t out_elem_offset, DeviceAttrs &attrs) { #ifdef __CUDACC__ MATX_NVTX_START("", matx::MATX_NVTX_LOG_INTERNAL) - - constexpr int NOUT = SmemTiledNout; + constexpr int CTILE = SmemTiledCtile; const index_t num_channels = o.Size(OutType::Rank() - 1); const index_t nout_per_channel = o.Size(OutType::Rank() - 2); @@ -345,45 +383,36 @@ inline void SmemTiledImpl( const index_t gcd_val = std::gcd(num_channels, decimation_factor); const int32_t K = static_cast(num_channels / gcd_val); - const int channel_tiles = static_cast((num_channels + CTILE - 1) / CTILE); - // Target ~1024 spatial blocks (time * channel tiles) to saturate the GPU. - // For large channel counts, fewer time blocks are needed. - const int target_time_blocks = cuda::std::max(1, (1024 + channel_tiles - 1) / channel_tiles); - const int elem_per_block = static_cast( - (nout_per_channel + target_time_blocks - 1) / target_time_blocks); + const int channel_tiles = SmemTiledChannelTiles(num_channels); + // Target spatial blocks across the whole batch, with bounded CTA work. + const index_t target_ctas = SmemTiledCtasPerSm * attrs.SmCount(); + int elem_per_block = static_cast(SmemTiledElemsPerBlock( + nout_per_channel, num_channels, num_batches, target_ctas)); + // Use the full time dimension of oversampled FIR tiles when possible. + if (decimation_factor < num_channels) { + elem_per_block = std::max(elem_per_block, NOUT); + } + if constexpr (MaximallyDecimated && NOUT > SmemTiledMaxDecYThreads) { + // Each iteration computes a full group even in a partial CTA. Pack output + // rows into complete groups so only the final CTA can waste that work. + elem_per_block = ((elem_per_block + NOUT - 1) / NOUT) * NOUT; + } const int time_blocks = static_cast( (nout_per_channel + elem_per_block - 1) / elem_per_block); - dim3 block(CTILE, NOUT); + dim3 block(CTILE, MaximallyDecimated ? SmemTiledMaxDecYThreads : NOUT); dim3 grid(time_blocks, channel_tiles, num_batches); - const auto filter_layout = SmemTiledChooseFilterLayout( - o, filter, decimation_factor, CTILE); - const size_t smem_size = SmemTiledSizeBytes(o, i, filter, decimation_factor, CTILE); - - // Use int32_t for intra-kernel index arithmetic when all tensor dimensions - // fit, avoiding 64-bit IMAD.WIDE instructions in the inner loops. - const index_t input_len = i.Size(i.Rank() - 1); + // Use int32_t for intra-kernel index arithmetic when all indices fit, + // avoiding 64-bit IMAD.WIDE instructions in the inner loops. const bool use_32bit = (sizeof(index_t) <= sizeof(int32_t)) || - (static_cast(input_len) + num_channels <= std::numeric_limits::max() && - nout_per_channel <= std::numeric_limits::max() && - num_channels <= std::numeric_limits::max()); + SmemTiledFitsInt32(i.Size(i.Rank() - 1), num_channels, decimation_factor, + out_elem_offset + nout_per_channel); // Dispatch on MaximallyDecimated x (FilterInSmem, FilterFullLayout) x // IndexType x IsUnitStride. Filter-smem layout selection: Full (smallest // footprint, phase compute at access), Rotated (direct indexing, redundant // storage), or Global (filter stays in GMEM). - [[maybe_unused]] constexpr bool kMaxDec = true; - [[maybe_unused]] constexpr bool kOversampled = false; - - // Unit-stride fast path eligibility + runtime check: same pattern as - // sar_bp / ChannelizePoly1D. Computed ops without .Data() / .Stride() fall - // through to the slow-path (operator()) instantiation. - constexpr bool fast_path_eligible = - is_tensor_view_v && - is_tensor_view_v && - is_tensor_view_v; - auto launch = [&](auto idx_tag, auto is_unit_c) { using IdxT = decltype(idx_tag); constexpr bool IsUnitStride = decltype(is_unit_c)::value; @@ -394,16 +423,13 @@ inline void SmemTiledImpl( auto launch_with_layout = [&](auto in_smem_c, auto full_c) { constexpr bool FIS = decltype(in_smem_c)::value; constexpr bool FFL = decltype(full_c)::value; - if (decimation_factor == num_channels) { - ChannelizePoly1D_SmemTiled - <<>>(o, i, filter, epb, df, K, oeo); - } else { - ChannelizePoly1D_SmemTiled - <<>>(o, i, filter, epb, df, K, oeo); - } + ChannelizePoly1D_SmemTiled + <<>>(o, i, filter, epb, df, K, oeo); }; - switch (filter_layout) { + switch (plan.filter_layout) { case SmemTiledFilterLayout::Full: launch_with_layout(cuda::std::bool_constant{}, cuda::std::bool_constant{}); break; @@ -417,19 +443,9 @@ inline void SmemTiledImpl( }; auto dispatch = [&](auto idx_tag) { - if constexpr (fast_path_eligible) { - const bool is_unit_stride = - o.Stride(OutType::Rank() - 1) == 1 && - i.Stride(InType::Rank() - 1) == 1 && - filter.Stride(FilterType::Rank() - 1) == 1; - if (is_unit_stride) { - launch(idx_tag, cuda::std::bool_constant{}); - } else { - launch(idx_tag, cuda::std::bool_constant{}); - } - } else { - launch(idx_tag, cuda::std::bool_constant{}); - } + DispatchUnitStride([&](auto is_unit_c) { + launch(idx_tag, is_unit_c); + }, o, i, filter); }; if constexpr (sizeof(index_t) <= sizeof(int32_t)) { @@ -444,25 +460,25 @@ inline void SmemTiledImpl( #endif } -// Wrapper that picks CTILE=32 (single-tile, full thread utilization) for -// num_channels <= 32 and CTILE=64 otherwise. Previously this path was -// rejected for num_channels <= 48 and fell through to the generic -// ChannelizePoly1D kernel, which does uncoalesced strided global loads and -// is ~3x slower. +// Launch the tile height and filter layout selected by the shared-memory fit check. template inline void SmemTiled( OutType o, const InType &i, const FilterType &filter, - index_t decimation_factor, cudaStream_t stream, index_t out_elem_offset = 0) + index_t decimation_factor, const SmemTiledPlan &plan, + cudaStream_t stream, index_t out_elem_offset, DeviceAttrs &attrs) { const index_t num_channels = o.Size(OutType::Rank() - 1); - if (num_channels <= SmemTiledCtileSmall) { - SmemTiledImpl< - SmemTiledCtileSmall, - OutType, InType, FilterType, AccumType>(o, i, filter, decimation_factor, stream, out_elem_offset); + auto launch = [&]() { + SmemTiledImpl( + o, i, filter, decimation_factor, plan, stream, out_elem_offset, attrs); + }; + if (plan.nout == SmemTiledMaxDecNout && decimation_factor == num_channels) { + launch.template operator()(); + } else if (plan.nout == SmemTiledNout) { + // Critical 32x4 fallbacks reuse the oversampled leaf's K=1 case. + launch.template operator()(); } else { - SmemTiledImpl< - SmemTiledCtile, - OutType, InType, FilterType, AccumType>(o, i, filter, decimation_factor, stream, out_elem_offset); + MATX_THROW(matxInvalidParameter, "channelize_poly: unsupported tiled row count"); } } @@ -489,53 +505,45 @@ inline void Generic(OutType o, const InType &i, const int THREADS = 256; const index_t ELTS_PER_THREAD = ElemsPerThread * THREADS; - const int elem_blocks = static_cast( - (nout_per_channel + ELTS_PER_THREAD - 1) / ELTS_PER_THREAD); - dim3 grid(elem_blocks, static_cast(num_channels), num_batches); - - // Unit-stride fast path: only viable when every hot tensor is a - // storage-backed view (Data()/Stride() callable) AND each one's last-dim - // stride is 1. Runtime-check those strides and dispatch to the specialized - // kernel via a bool_constant lambda. - constexpr bool fast_path_eligible = - is_tensor_view_v && - is_tensor_view_v && - is_tensor_view_v; + const index_t elem_blocks = (nout_per_channel + ELTS_PER_THREAD - 1) / ELTS_PER_THREAD; + // Channels vary fastest (grid.x) for input reuse through L2. Time blocks use + // grid.y, split into launches of at most 65535 blocks. + constexpr index_t max_grid_y = 65535; + MATX_ASSERT_STR(elem_blocks <= std::numeric_limits::max(), matxInvalidSize, + "channelize_poly: output too long for the generic kernel"); auto launch = [&](auto is_unit_c) { constexpr bool IsUnitStride = decltype(is_unit_c)::value; - if (decimation_factor == num_channels) { - // For M == D, cache one filter phase in dynamic shared memory if it fits. - const index_t filter_phase_len = (filter_len + num_channels - 1) / num_channels; - const size_t smem_needed = static_cast(filter_phase_len) * sizeof(filter_t); - const uint32_t smem_bytes = (smem_needed <= GenericMaxFilterSmemBytes) - ? static_cast(smem_needed) : 0; - ChannelizePoly1D - <<>>(o, i, filter, decimation_factor, smem_bytes, out_elem_offset); - } else { - ChannelizePoly1D - <<>>(o, i, filter, decimation_factor, 0, out_elem_offset); + for (index_t offset = 0; offset < elem_blocks; offset += max_grid_y) { + const dim3 grid(static_cast(num_channels), + static_cast(std::min(max_grid_y, elem_blocks - offset)), + num_batches); + if (decimation_factor == num_channels) { + // For M == D, cache one filter phase in dynamic shared memory if it fits. + const index_t filter_phase_len = (filter_len + num_channels - 1) / num_channels; + const size_t smem_needed = static_cast(filter_phase_len) * sizeof(filter_t); + const uint32_t smem_bytes = (smem_needed <= GenericMaxFilterSmemBytes) + ? static_cast(smem_needed) : 0; + ChannelizePoly1D + <<>>( + o, i, filter, decimation_factor, smem_bytes, out_elem_offset, + static_cast(offset)); + } else { + ChannelizePoly1D + <<>>( + o, i, filter, decimation_factor, 0, out_elem_offset, static_cast(offset)); + } } }; - if constexpr (fast_path_eligible) { - const bool is_unit_stride = - o.Stride(OutType::Rank() - 1) == 1 && - i.Stride(InType::Rank() - 1) == 1 && - filter.Stride(FilterType::Rank() - 1) == 1; - if (is_unit_stride) { - launch(cuda::std::bool_constant{}); - } else { - launch(cuda::std::bool_constant{}); - } - } else { - launch(cuda::std::bool_constant{}); - } + DispatchUnitStride(launch, o, i, filter); #endif } template -inline size_t SmemSizeBytes(const OutType &o, const InType &, const FilterType &filter) +inline size_t SmemSizeBytes( + const OutType &o, const InType &, const FilterType &filter, + int nout = FullSmemKernelNoutPerIter) { using input_t = typename InType::value_type; using filter_t = typename FilterType::value_type; @@ -546,7 +554,7 @@ inline size_t SmemSizeBytes(const OutType &o, const InType &, const FilterType & const index_t filter_phase_len = (filter_len + num_channels - 1) / num_channels; size_t smem_size = sizeof(filter_t)*(num_channels)*(filter_phase_len) + - sizeof(input_t)*(num_channels)*(filter_phase_len + FullSmemKernelNoutPerIter - 1); + sizeof(input_t)*(num_channels)*(filter_phase_len + nout - 1); const size_t max_sizeof = cuda::std::max(sizeof(filter_t), sizeof(input_t)); if (smem_size % max_sizeof) { smem_size += max_sizeof - (smem_size % max_sizeof); @@ -555,24 +563,21 @@ inline size_t SmemSizeBytes(const OutType &o, const InType &, const FilterType & } template -inline size_t ShouldUseSmem(const OutType &out, const InType &in, const FilterType &filter) +inline bool ShouldUseSmem(const OutType &out, const InType &in, const FilterType &filter) { - // 48 KB is the largest shared memory allocation that does not require - // explicit opt-in via cudaFuncSetAttribute() - const size_t MAX_SMEM_BYTES = 48 * 1024; // The full shared memory kernel uses blocks of size // (num_channels, FullSmemKernelNoutPerIter), so ensure // that the resulting thread per block count will not exceed MAX_NUM_THREADS_PER_BLOCK const int MAX_NUM_THREADS_PER_BLOCK = 1024; const index_t num_channels = out.Size(OutType::Rank()-1); return ( - SmemSizeBytes(out, in, filter) <= MAX_SMEM_BYTES && + SmemSizeBytes(out, in, filter) <= MaxDefaultDynamicSmemBytes && num_channels <= (MAX_NUM_THREADS_PER_BLOCK/FullSmemKernelNoutPerIter)); } template -inline void Smem(OutType o, const InType &i, const FilterType &filter, cudaStream_t stream, - index_t out_elem_offset = 0) +inline void Smem(OutType o, const InType &i, const FilterType &filter, + cudaStream_t stream, index_t out_elem_offset, DeviceAttrs &attrs) { #ifdef __CUDACC__ MATX_NVTX_START("", matx::MATX_NVTX_LOG_INTERNAL) @@ -581,98 +586,374 @@ inline void Smem(OutType o, const InType &i, const FilterType &filter, cudaStrea const index_t nout_per_channel = o.Size(OutType::Rank()-2); const int num_batches = static_cast(TotalSize(i)/i.Size(i.Rank() - 1)); - const int target_num_blocks = 1024; - const int elem_per_block = static_cast( - (nout_per_channel + target_num_blocks - 1) / target_num_blocks); dim3 block(static_cast(num_channels), FullSmemKernelNoutPerIter); - const uint32_t num_blocks = static_cast((nout_per_channel + elem_per_block - 1) / elem_per_block); - dim3 grid(num_blocks, 1, num_batches); - const size_t smem_size = SmemSizeBytes(o, i, filter); - - constexpr bool fast_path_eligible = - is_tensor_view_v && - is_tensor_view_v && - is_tensor_view_v; + size_t smem_size = SmemSizeBytes(o, i, filter); + auto launch = [&](auto is_unit_c) { constexpr bool IsUnitStride = decltype(is_unit_c)::value; - ChannelizePoly1D_Smem + auto kernel = ChannelizePoly1D_Smem; + const index_t sms = attrs.SmCount(); + const index_t l2_bytes = static_cast(attrs.L2Bytes()); + const index_t row_bytes = num_channels * static_cast( + sizeof(typename InType::value_type) + 2 * sizeof(AccumType)); + const size_t batch_bytes = static_cast(nout_per_channel) * row_bytes; + const size_t input_bytes = static_cast(TotalSize(i)) * + sizeof(typename InType::value_type); + // When input fits L2 but the pipeline footprint does not, retain small + // whole-warp tiles and enough CTAs to cover the kernel's residency. + const bool cache_pressure = block.x * block.y <= 3 * 32 && + input_bytes < static_cast(l2_bytes) && batch_bytes >= static_cast(l2_bytes); + // Taller groups amortize barriers in blocks with fewer than four warps. + // Keep at least one CTA per SM and preserve the original fit fallback. + while (block.y < 16 && block.x * block.y <= 3 * 32 && + (!cache_pressure || (block.x * block.y) % 32 != 0)) { + const int next = static_cast(2 * block.y); + const size_t bytes = SmemSizeBytes(o, i, filter, next); + if (bytes > MaxDefaultDynamicSmemBytes || + ((nout_per_channel + next - 1) / next) * num_batches < sms) + break; + block.y = static_cast(next); + smem_size = bytes; + } + const index_t group = block.y; + int resident_blocks = 0; + MATX_CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor( + &resident_blocks, kernel, static_cast(block.x * block.y), smem_size)); + index_t targets = std::max(1, + ((cache_pressure ? 1 : 2) * sms * resident_blocks + num_batches - 1) / + num_batches); + index_t elem_per_block = (nout_per_channel + targets - 1) / targets; + if (group > FullSmemKernelNoutPerIter) + elem_per_block = std::max(elem_per_block, group); + // Bound serial work without rounding longer spans and reducing coverage. + // This is a logical work budget, not exclusive L2 ownership. + if (sms > 0 && l2_bytes > 0) { + const index_t limit = std::max(group, (l2_bytes / sms / row_bytes / group) * group); + elem_per_block = std::min(elem_per_block, limit); + } + // Amortize a mostly idle final iteration without another tile shape. + const index_t rounded = MATX_ROUND_UP(elem_per_block, group); + if (group > FullSmemKernelNoutPerIter && + elem_per_block >= group && 4 * elem_per_block < 3 * rounded) { + elem_per_block = rounded; + } + dim3 grid(static_cast( + (nout_per_channel + elem_per_block - 1) / elem_per_block), 1, num_batches); + kernel <<>>(o, i, filter, elem_per_block, out_elem_offset); }; - if constexpr (fast_path_eligible) { - const bool is_unit_stride = - o.Stride(OutType::Rank() - 1) == 1 && - i.Stride(InType::Rank() - 1) == 1 && - filter.Stride(FilterType::Rank() - 1) == 1; - if (is_unit_stride) { - launch(cuda::std::bool_constant{}); - } else { - launch(cuda::std::bool_constant{}); + DispatchUnitStride(launch, o, i, filter); +#endif +} + +template +inline size_t FusedRadixSizeBytes( + const OutType &out, const InType &, const FilterType &filter, index_t decimation_factor) +{ + using input_t = typename InType::value_type; + using filter_t = typename FilterType::value_type; + using critical_layout = FusedRadixSmemLayout; + using oversampled_layout = FusedRadixSmemLayout; + const index_t P = (filter.Size(FilterType::Rank() - 1) + NUM_CHAN - 1) / NUM_CHAN; + const size_t bytes = decimation_factor == NUM_CHAN + ? critical_layout(P, decimation_factor).bytes + : oversampled_layout(P, decimation_factor).bytes; + + if constexpr (NUM_CHAN >= 3 && NUM_CHAN <= 6) { + if (decimation_factor == NUM_CHAN && + FusedSmallUseDirect( + P, out.Size(OutType::Rank() - 2))) { + return 0; } + } + return bytes; +} + +// Whether the fused leaf's working set fits the default shared-memory limit. +template +inline bool FusedRadixFits( + const OutType &out, const InType &in, const FilterType &filter, index_t decimation_factor) +{ + return FusedRadixSizeBytes( + out, in, filter, decimation_factor) <= MaxDefaultDynamicSmemBytes; +} + +// Fused leaves avoid the intermediate filtered tensor and the separate FFT, +// and are preferred whenever they apply. The exceptions are grouped leaves +// (M > 6) that measure slower than the FIR plus cuFFT path: +// - complex inputs to leaves whose DFT stages go through shared memory (no +// register FFT) or, when critically sampled, that pad each row to more +// than twice the channel count. The separate path writes and rereads the +// filtered intermediate, which costs less than this extra arithmetic while +// the intermediate stays in L2 or memory bandwidth per SM is high. +// Oversampling enlarges the intermediate that fusion removes, which +// outweighs the padding. +// - FP64 on GPUs with low FP64 throughput, where the fused leaf's padded, +// complex-valued arithmetic dominates. Complex double always takes the +// separate path there. Real double does so only when the separate path's +// working set stays in L2 and each batch is long enough for the FIR and FFT +// kernels to run efficiently; short or DRAM-bound calls still fuse. +constexpr index_t SeparateFp64MinInputLength = 128 * 1024; + +template +inline bool FusedRadixPreferred(index_t decimation_factor, index_t input_length, + size_t input_bytes, DeviceAttrs &attrs) +{ + if constexpr (NUM_CHAN <= 6) { + return true; + } else if constexpr (!is_complex_v) { + if constexpr (cuda::std::is_same_v) { + return !(input_length >= SeparateFp64MinInputLength && + attrs.LowFp64Throughput() && input_bytes <= attrs.L2Bytes()); + } + return true; } else { - launch(cuda::std::bool_constant{}); + using Config = FusedRadixConfig; + if (!Config::WarpFft || (decimation_factor == NUM_CHAN && Config::RowStride > 2 * NUM_CHAN)) { + if (input_bytes <= attrs.L2Bytes() || attrs.HighMemoryBandwidth()) { + return false; + } + } + if constexpr (cuda::std::is_same_v) { + return !attrs.LowFp64Throughput(); + } + return true; } -#endif } -template -inline void FusedChan(OutType o, const InType &i, - const FilterType &filter, cudaStream_t stream, - index_t out_elem_offset = 0) +template +inline void FusedRadixImpl( + OutType o, const InType &i, const FilterType &filter, + index_t decimation_factor, cudaStream_t stream, index_t out_elem_offset = 0) { #ifdef __CUDACC__ MATX_NVTX_START("", matx::MATX_NVTX_LOG_INTERNAL) + using config = FusedRadixConfig< + NUM_CHAN, AccumType, MaximallyDecimated, typename InType::value_type>; + constexpr int threads = config::Threads; + constexpr int nrows = config::NRows; + const index_t nout_per_channel = o.Size(OutType::Rank() - 2); + const int num_batches = static_cast(TotalSize(i) / i.Size(i.Rank() - 1)); + constexpr int target_blocks = 512; + index_t elem_per_block; + if constexpr (NUM_CHAN == 2 || (NUM_CHAN <= 6 && MaximallyDecimated)) { + constexpr int block_rows = threads * config::SmallOutputsPerThread; + elem_per_block = block_rows; + if constexpr (NUM_CHAN == 2 && MaximallyDecimated && + cuda::std::is_same_v && + is_complex_v) { + using traits = ChannelizePolyFusedSmallTraits< + NUM_CHAN, OutType, InType, FilterType, AccumType>; + if constexpr (traits::UseM2PairLeaf) { + // Favor grid coverage for short inputs; keep reuse for larger batches. + const index_t blocks = (nout_per_channel + block_rows - 1) / block_rows; + if (blocks <= (target_blocks - 1) / num_batches) elem_per_block = threads; + } + } + } else { + const index_t target_elems = (nout_per_channel + target_blocks - 1) / target_blocks; + // Each loop iteration computes a complete NROWS tile. Rounding the CTA's span prevents hundreds + // of CTAs from doing mostly-invalid tile work for short signals or large channel counts. + elem_per_block = MATX_ROUND_UP(target_elems, nrows); + } + const int time_blocks = static_cast( + (nout_per_channel + elem_per_block - 1) / elem_per_block); + const dim3 grid(time_blocks, 1, num_batches); + const size_t smem_size = + FusedRadixSizeBytes( + o, i, filter, decimation_factor); - const index_t num_channels = o.Size(OutType::Rank()-1); - const index_t nout_per_channel = o.Size(OutType::Rank()-2); - const int num_batches = static_cast(TotalSize(i)/i.Size(i.Rank() - 1)); - - const int THREADS = 256; - const index_t ELTS_PER_THREAD = ElemsPerThread * THREADS; - const int elem_blocks = static_cast( - (nout_per_channel + ELTS_PER_THREAD - 1) / ELTS_PER_THREAD); - dim3 grid(elem_blocks, 1, num_batches); - - constexpr bool fast_path_eligible = - is_tensor_view_v && - is_tensor_view_v && - is_tensor_view_v; auto launch = [&](auto is_unit_c) { - constexpr bool IsUnitStride = decltype(is_unit_c)::value; - // Dispatch on num_channels in [2, FusedChanThreshold]. Generated at - // compile time from the threshold so raising FusedChanThreshold - // automatically grows the dispatch table; no per-N switch case to - // update. A runtime value outside the range falls through to MATX_THROW. - constexpr int kMinChan = 2; - [[maybe_unused]] constexpr int kMaxChan = static_cast(FusedChanThreshold); - const bool matched = [&](cuda::std::integer_sequence) { - return ((num_channels == kMinChan + Is - ? (ChannelizePoly1D_FusedChan - <<>>(o, i, filter, out_elem_offset), - true) - : false) - || ...); - }(cuda::std::make_integer_sequence{}); - if (!matched) { - MATX_THROW(matxInvalidDim, "channelize_poly: channel count not supported with fused kernel"); - } - }; - if constexpr (fast_path_eligible) { - const bool is_unit_stride = - o.Stride(OutType::Rank() - 1) == 1 && - i.Stride(InType::Rank() - 1) == 1 && - filter.Stride(FilterType::Rank() - 1) == 1; - if (is_unit_stride) { - launch(cuda::std::bool_constant{}); + constexpr bool is_unit_stride = decltype(is_unit_c)::value; + auto launch_indexed = [&](auto idx_tag) { + using IdxT = decltype(idx_tag); + ChannelizePoly1D_FusedRadixPow2< + threads, NUM_CHAN, nrows, MaximallyDecimated, is_unit_stride, + OutType, InType, FilterType, AccumType, IdxT> + <<>>(o, i, filter, + static_cast(elem_per_block), static_cast(decimation_factor), + static_cast(out_elem_offset)); + }; + // Narrow only grouped mixed-radix warp FFTs; retain all other index paths. + if constexpr (sizeof(index_t) <= sizeof(int32_t) || !config::WarpFft || config::Radix1 == 1) { + launch_indexed(index_t{}); } else { - launch(cuda::std::bool_constant{}); + // Keep strides and batch addressing wide; reserve padded-group headroom. + const index_t padding_rows = elem_per_block + nrows; + const index_t max_index = std::numeric_limits::max(); + const index_t input_len = i.Size(InType::Rank() - 1); + const index_t filter_len = filter.Size(FilterType::Rank() - 1); + const bool padding_fits = padding_rows < max_index / NUM_CHAN; + const index_t limit = padding_fits ? max_index - padding_rows * NUM_CHAN : 0; + const bool use_32bit = padding_fits && input_len <= limit && + filter_len <= max_index - NUM_CHAN && nout_per_channel <= limit / decimation_factor && + out_elem_offset <= limit / decimation_factor - nout_per_channel; + if (use_32bit) { + launch_indexed(int32_t{}); + } else { + launch_indexed(index_t{}); + } } + }; + if constexpr (NUM_CHAN <= 6 && MaximallyDecimated) { + DispatchUnitStride(launch, o, i, filter); } else { - launch(cuda::std::bool_constant{}); + // General fused launches instantiate only unit-stride tensor paths. The public + // dispatcher checks this precondition. + auto check_stride = [](const auto &op) { + using op_t = cuda::std::remove_cvref_t; + if constexpr (is_tensor_view_v) { + if (op.Stride(op_t::Rank() - 1) != 1) { + MATX_THROW(matxInvalidParameter, + "channelize_poly: general fused launch requires unit last-dimension stride"); + } + } + }; + check_stride(o); + check_stride(i); + check_stride(filter); + launch(cuda::std::bool_constant{}); } #endif } +// Type-level eligibility of the fused FIR+DFT kernels. Critical M=2..6 +// leaves accept expressions, strided views, half precision, and mixed types. +// The grouped leaves require matched float/double types and a tensor filter. +template +struct FusedRadixTypes { + using output_t = typename OutType::value_type; + using input_t = typename InType::value_type; + using filter_t = typename FilterType::value_type; + // Input expressions are evaluated while the fused kernel stages its input + // tile. This preserves narrow source loads and avoids materializing results + // such as converted, phase-corrected complex samples. + static constexpr bool Small = + is_tensor_view_v && is_matx_op() && + is_matx_op() && is_complex_v && + (cuda::std::is_same_v || + cuda::std::is_same_v || is_matx_half_v); + static constexpr bool General = + is_tensor_view_v && is_matx_op() && is_tensor_view_v && + (cuda::std::is_same_v || + cuda::std::is_same_v) && + cuda::std::is_same_v> && + (cuda::std::is_same_v || + cuda::std::is_same_v>) && + cuda::std::is_same_v; +}; + +// Invoke fn with the compile-time channel count of a fused leaf and return its result, or return +// false when no leaf is instantiated for num_channels. M=2...6 replace the older FusedChan family; +// the remaining entries are the measured power-of-two and K-by-power-of-two cases. +template +inline bool VisitFusedRadixChannels(index_t num_channels, Fn &&fn) +{ + switch (num_channels) { + case 2: return fn(cuda::std::integral_constant{}); + case 3: return fn(cuda::std::integral_constant{}); + case 4: return fn(cuda::std::integral_constant{}); + case 5: return fn(cuda::std::integral_constant{}); + case 6: return fn(cuda::std::integral_constant{}); + default: break; + } + if constexpr (!SmallOnly) { + switch (num_channels) { + case 8: return fn(cuda::std::integral_constant{}); + case 10: return fn(cuda::std::integral_constant{}); + case 16: return fn(cuda::std::integral_constant{}); + case 20: return fn(cuda::std::integral_constant{}); + case 32: return fn(cuda::std::integral_constant{}); + case 40: return fn(cuda::std::integral_constant{}); + case 64: return fn(cuda::std::integral_constant{}); + case 80: return fn(cuda::std::integral_constant{}); + default: break; + } + } + return false; +} + +// Whether a fused leaf exists for these types, strides, and channel count and +// its working set fits in shared memory. This is a launchability check only. +template +inline bool FusedRadixFeasible( + const OutType &out, const InType &in, const FilterType &filter, + index_t num_channels, index_t decimation_factor) +{ + using types = FusedRadixTypes; + auto fits = [&](auto channels_c) { + constexpr int channels = decltype(channels_c)::value; + return FusedRadixFits( + out, in, filter, decimation_factor); + }; + if constexpr (types::Small) { + if (decimation_factor == num_channels && num_channels >= 2 && num_channels <= 6) { + return VisitFusedRadixChannels(num_channels, fits); + } + } + if constexpr (types::General) { + if (out.Stride(OutType::Rank() - 1) != 1 || filter.Stride(FilterType::Rank() - 1) != 1) { + return false; + } + if constexpr (is_tensor_view_v) { + if (in.Stride(InType::Rank() - 1) != 1) return false; + } + return VisitFusedRadixChannels(num_channels, fits); + } + return false; +} + +// Performance preference for a feasible fused launch. A critical small-channel +// leaf whose FIR tile does not fit in shared memory runs the uncached direct FIR, +// which dispatch reserves for outputs that cuFFT cannot produce. +template +inline bool PreferFusedRadix(const InType &in, const FilterType &filter, index_t num_channels, + index_t decimation_factor, DeviceAttrs &attrs) +{ + using types = FusedRadixTypes; + const index_t input_length = in.Size(InType::Rank() - 1); + const size_t input_bytes = static_cast(TotalSize(in)) * + sizeof(typename InType::value_type); + const index_t taps = (filter.Size(FilterType::Rank() - 1) + num_channels - 1) / num_channels; + return VisitFusedRadixChannels(num_channels, [&](auto channels_c) { + constexpr int channels = decltype(channels_c)::value; + if constexpr (channels >= 3 && channels <= 6) { + if (decimation_factor == channels && + !FusedSmallCachedFits(taps)) { + return false; + } + } + return FusedRadixPreferred( + decimation_factor, input_length, input_bytes, attrs); + }); +} + +// Launch the fused leaf for num_channels. The caller checks feasibility. +template +inline void RunFusedRadix( + OutType out, const InType &in, const FilterType &filter, + index_t num_channels, index_t decimation_factor, cudaStream_t stream, index_t out_elem_offset) +{ + using types = FusedRadixTypes; + if constexpr (types::Small || types::General) { + VisitFusedRadixChannels(num_channels, [&](auto channels_c) { + constexpr int channels = decltype(channels_c)::value; + if (decimation_factor == channels) { + FusedRadixImpl( + out, in, filter, decimation_factor, stream, out_elem_offset); + return true; + } + if constexpr (types::General) { + FusedRadixImpl( + out, in, filter, decimation_factor, stream, out_elem_offset); + return true; + } + return false; + }); + } +} + template inline void UnpackDFT(DataType inout, cudaStream_t stream) { @@ -683,25 +964,234 @@ inline void UnpackDFT(DataType inout, cudaStream_t stream) const index_t num_channels = inout.Size(DataType::Rank()-1); const int num_batches = static_cast(TotalSize(inout)/ (num_channels * num_elem_per_channel)); - const int gy = static_cast((num_elem_per_channel + THREADS - 1) / THREADS); + const int gy = static_cast(std::min( + (num_elem_per_channel + THREADS - 1) / THREADS, 65535)); const dim3 grid(num_batches, gy); - constexpr bool fast_path_eligible = is_tensor_view_v; auto launch = [&](auto is_unit_c) { constexpr bool IsUnitStride = decltype(is_unit_c)::value; ChannelizePoly1DUnpackDFT<<>>(inout); }; - if constexpr (fast_path_eligible) { - const bool is_unit_stride = inout.Stride(DataType::Rank() - 1) == 1; - if (is_unit_stride) { - launch(cuda::std::bool_constant{}); - } else { - launch(cuda::std::bool_constant{}); + DispatchUnitStride(launch, inout); +#endif +} + +// CUDA backends. Fused computes the FIR and DFT in one kernel. The others run +// a FIR kernel followed by a cuFFT transform (and an unpack for real inputs). +enum class Backend { Fused, Smem, Tiled, Generic }; + +struct Plan { + Backend backend = Backend::Generic; + // Tiled: tile height and filter layout. + SmemTiledPlan tiled{}; +}; + +// The FIR backends run the DFT with cuFFT, whose half-precision transforms +// require a power-of-two size. +template +inline bool SeparateFftSupported(index_t num_channels) +{ + if constexpr (is_complex_half_v) { + return (num_channels & (num_channels - 1)) == 0; + } + return true; +} + +// Dispatch policy: +// 1. A fused leaf whenever one applies and is preferred (PreferFusedRadix), or +// whenever cuFFT cannot produce the output (SeparateFftSupported). +// 2. Otherwise a FIR kernel followed by cuFFT. Up to SmallChannelLimit +// channels, the whole-channel Smem kernel (critical) or Generic +// (oversampled); wider channel counts use the tiled kernel +// (SelectTiledPlan), except that critical FP64 on low-FP64 GPUs keeps +// Smem when the last channel tile would be partly idle. +// 3. Generic when no tile fits the default shared-memory limit. +// Without a fused leaf, outputs that cuFFT cannot produce get an unlaunchable +// plan, which ExecutePlan rejects. +template +inline Plan SelectPlan( + const OutType &out, const InType &in, const FilterType &filter, + index_t num_channels, index_t decimation_factor, DeviceAttrs &attrs) +{ + Plan plan; + if (FusedRadixFeasible( + out, in, filter, num_channels, decimation_factor) && + (!SeparateFftSupported(num_channels) || + PreferFusedRadix( + in, filter, num_channels, decimation_factor, attrs))) { + plan.backend = Backend::Fused; + return plan; + } + const bool critical = decimation_factor == num_channels; + // With low FP64 throughput, the whole-channel Smem kernel also beats a tiled + // kernel whose last channel tile is partly idle. + bool smem_preferred = num_channels <= SmallChannelLimit; + if constexpr (cuda::std::is_same_v) { + smem_preferred = smem_preferred || + (num_channels % SmemTiledCtile != 0 && attrs.LowFp64Throughput()); + } + if (critical && smem_preferred && ShouldUseSmem(out, in, filter)) { + plan.backend = Backend::Smem; + return plan; + } + if (!critical && num_channels <= SmallChannelLimit) { + return plan; + } + plan.tiled = SelectTiledPlan(out, in, filter, decimation_factor); + if (SmemTiledFits(plan.tiled, num_channels)) { + plan.backend = Backend::Tiled; + } + return plan; +} + +template +inline Plan SelectPlan( + const OutType &out, const InType &in, const FilterType &filter, + index_t num_channels, index_t decimation_factor) +{ + DeviceAttrs attrs; + return SelectPlan( + out, in, filter, num_channels, decimation_factor, attrs); +} + +// Whether a plan can be launched for these operands. Selected plans can unless +// no backend can produce the output; this also guards explicitly constructed plans. +template +inline bool PlanIsLaunchable( + const Plan &plan, const OutType &out, const InType &in, + const FilterType &filter, index_t num_channels, index_t decimation_factor) +{ + if (plan.backend != Backend::Fused && !SeparateFftSupported(num_channels)) { + return false; + } + switch (plan.backend) { + case Backend::Fused: + return FusedRadixFeasible( + out, in, filter, num_channels, decimation_factor); + case Backend::Smem: + return decimation_factor == num_channels && ShouldUseSmem(out, in, filter); + case Backend::Tiled: { + const int nout = plan.tiled.nout; + const bool nout_ok = nout == SmemTiledNout || + (nout == SmemTiledMaxDecNout && decimation_factor == num_channels); + const auto sized = SmemTiledPlanWithLayout( + out, in, filter, decimation_factor, nout, plan.tiled.filter_layout); + return nout_ok && sized.bytes == plan.tiled.bytes && + SmemTiledFits(sized, num_channels) && + (plan.tiled.filter_layout != SmemTiledFilterLayout::Rotated || + SmemTiledRotatedLayoutSupported(num_channels, decimation_factor)); + } + case Backend::Generic: + return true; + } + return false; +} + +template +inline void ExecutePlan( + const Plan &plan, OutType out, const InType &in, const FilterType &f, + index_t num_channels, index_t decimation_factor, cudaStream_t stream, + index_t out_elem_offset, DeviceAttrs &attrs) +{ + using input_t = typename InType::value_type; + using filter_t = typename FilterType::value_type; + using output_t = typename OutType::value_type; + constexpr int OUT_RANK = OutType::Rank(); + + if (!PlanIsLaunchable( + plan, out, in, f, num_channels, decimation_factor)) { + MATX_THROW(matxInvalidParameter, SeparateFftSupported(num_channels) + ? "channelize_poly: backend plan is not supported for these operands" + : "channelize_poly: half-precision outputs need a power-of-two channel count " + "unless a fused kernel applies (critical sampling with 3, 5, or 6 channels)"); + } + + if (plan.backend == Backend::Fused) { + RunFusedRadix( + out, in, f, num_channels, decimation_factor, stream, out_elem_offset); + return; + } + + auto run_filter_backend = [&](auto filtered_output) { + using filtered_output_t = decltype(filtered_output); + switch (plan.backend) { + case Backend::Smem: + Smem( + filtered_output, in, f, stream, out_elem_offset, attrs); + break; + case Backend::Tiled: + SmemTiled( + filtered_output, in, f, decimation_factor, plan.tiled, stream, out_elem_offset, attrs); + break; + default: + Generic( + filtered_output, in, f, decimation_factor, stream, out_elem_offset); + break; } + }; + + // If neither the input nor the filter is complex, then the filtered samples will be real-valued + // and we will use an R2C transform. Otherwise, we will use a C2C transform. + if constexpr (!is_complex_v && !is_complex_v) { + index_t start_dims[OUT_RANK], stop_dims[OUT_RANK]; + std::fill_n(start_dims, OUT_RANK, 0); + std::fill_n(stop_dims, OUT_RANK, matxEnd); + + // The first kernel below needs a buffer of type input_t (known to be real in this constexpr + // branch) into which we store filtered data prior to the real-to-complex FFT. If the output + // buffer is contiguous, then we use an aliased tensor view of type input_t for that buffer + // where the last dimension is twice as large (because input_t is real and output_t is complex). + // We then use a slice to maintain the expected dimensions. If the output buffer is not + // contiguous, then we async allocate a temporary buffer. There is one caveat with this + // allocate: the batched fft implementation currently requires that all input pointers must be + // aligned to the corresponding complex type, which cannot be guaranteed to always be true for a + // real-valued tensor. This was not an issue for the reused output buffer because the output + // tensor is complex-valued, so we always have an even stride from one batch to the next. As a + // temporary workaround for the FFT alignment issue, we add one channel in the odd-channel case + // and use a slice to create a tensor view of only [0, num_channels-1]. This guarantees that we + // always stride by an even number of elements from one batch to the next while exposing a + // tensor view of appropriate dimensions. + using post_filter_t = typename inner_op_type_t::type; + auto fft_in_slice = [&out, &start_dims, &stop_dims, num_channels, stream]() -> auto { + auto fft_in_shape = out.Shape(); + if (out.IsContiguous()) { + fft_in_shape[OUT_RANK-1] *= 2; + auto fft_in = make_tensor( + reinterpret_cast(out.Data()), fft_in_shape); + stop_dims[OUT_RANK-1] = num_channels; + return slice(fft_in, start_dims, stop_dims); + } else { + if (num_channels % 2 == 1) { + fft_in_shape[OUT_RANK-1]++; + stop_dims[OUT_RANK-1] = num_channels; + } + auto tmp = make_tensor(fft_in_shape, MATX_ASYNC_DEVICE_MEMORY, stream); + return slice(tmp, start_dims, stop_dims); + } + }(); + + run_filter_backend(fft_in_slice); + stop_dims[OUT_RANK-1] = (num_channels/2) + 1; + auto out_packed = slice(out, start_dims, stop_dims); + (out_packed = fft(fft_in_slice, num_channels)).run(stream); + UnpackDFT(out, stream); } else { - launch(cuda::std::bool_constant{}); + run_filter_backend(out); + // Specify FORWARD here to prevent any normalization after the ifft. We do not + // want any extra scaling on the output values. + (out = ifft(out, num_channels, FFTNorm::FORWARD)).run(stream); } -#endif +} + +template +inline void ExecutePlan( + const Plan &plan, OutType out, const InType &in, const FilterType &f, + index_t num_channels, index_t decimation_factor, cudaStream_t stream, + index_t out_elem_offset = 0) +{ + DeviceAttrs attrs; + ExecutePlan( + plan, out, in, f, num_channels, decimation_factor, stream, out_elem_offset, attrs); } } // end namespace cpoly @@ -745,134 +1235,20 @@ inline void channelize_poly_impl(OutType out, const InType &in, const FilterType using OutputOp = cuda::std::remove_cv_t>; using InputOp = cuda::std::remove_cv_t>; using FilterOp = cuda::std::remove_cv_t>; - using input_t = typename InputOp::value_type; - using filter_t = typename FilterOp::value_type; - using output_t = typename OutputOp::value_type; - - constexpr int IN_RANK = InputOp::Rank(); - constexpr int OUT_RANK = OutputOp::Rank(); - // The last dimension of the input becomes [num_channels, num_elem_per_channel] in the last - // two dimensions of the output - MATX_STATIC_ASSERT_STR(OUT_RANK == IN_RANK+1, matxInvalidDim, "channelize_poly: output rank should be 1 higher than input"); - - MATX_STATIC_ASSERT_STR(is_complex_v || is_complex_half_v, - matxInvalidType, "channelize_poly: output type must be complex"); - - // Currently only support 1D filters. - MATX_STATIC_ASSERT_STR(FilterType::Rank() == 1, matxInvalidDim, "channelize_poly: currently only support 1D filters"); - - MATX_ASSERT_STR(num_channels > 0, matxInvalidParameter, - "channelize_poly: num_channels must be positive"); - MATX_ASSERT_STR(decimation_factor > 0, matxInvalidParameter, - "channelize_poly: decimation_factor must be positive"); - MATX_ASSERT_STR(decimation_factor <= num_channels, matxInvalidParameter, - "channelize_poly: decimation_factor must be <= num_channels"); - - for(int i = 0 ; i < IN_RANK-1; i++) { - MATX_ASSERT_STR(out.Size(i) == in.Size(i), matxInvalidDim, "channelize_poly: input/output must have matched batch sizes"); + detail::cpoly::ValidateChannelizePolyArgs( + out, in, f, num_channels, decimation_factor, out_elem_offset); + // An empty signal, batch dimension, or output window leaves nothing to compute. + if (TotalSize(out) == 0) { + return; } - [[maybe_unused]] const index_t num_elem_per_channel = (in.Size(IN_RANK-1) + decimation_factor - 1) / decimation_factor; - - MATX_ASSERT_STR(out.Size(OUT_RANK-1) == num_channels, matxInvalidDim, - "channelize_poly: output size OUT_RANK-1 mismatch"); - // The output is a window [out_elem_offset, out_elem_offset + rows) of the - // full output-element grid; a full (non-windowed) call is the special case - // rows == num_elem_per_channel, offset == 0. A zero offset does NOT imply the - // full grid: the streaming channelizer's first feed emits only the - // fully-covered prefix rows (floor(chunk/D)) at offset 0, while the full - // local grid has ceil(chunk/D) rows. The only requirement is that the window - // fits within the grid. - MATX_ASSERT_STR(out.Size(OUT_RANK-2) + out_elem_offset <= num_elem_per_channel, - matxInvalidDim, - "channelize_poly: output-element window exceeds the full output size"); - - // If neither the input nor the filter is complex, then the filtered samples will be real-valued - // and we will use an R2C transform. Otherwise, we will use a C2C transform. - if constexpr (! is_complex_v && ! is_complex_half_v && ! is_complex_v && ! is_complex_half_v) { - // The fused-DFT kernel only supports the maximally decimated case (D == M). - // num_channels == 1 is degenerate (no channelization, trivial DFT) and - // the fused kernel's switch starts at N=2. Let num_channels==1 fall - // through to Smem / SmemTiled / Generic, all of which handle it - // correctly as a plain FIR. All four filter kernels honor out_elem_offset, - // so the windowed path uses the same kernel selection as the full grid. - if (decimation_factor == num_channels && num_channels >= 2 && - num_channels <= detail::cpoly::FusedChanThreshold) { - detail::cpoly::FusedChan(out, in, f, stream, out_elem_offset); - } else { - index_t start_dims[OUT_RANK], stop_dims[OUT_RANK]; - std::fill_n(start_dims, OUT_RANK, 0); - std::fill_n(stop_dims, OUT_RANK, matxEnd); - - // The first kernel below needs a buffer of type input_t (known to be real in this - // constexpr branch) into which we store filtered data prior to the real-to-complex - // FFT. If the output buffer is contiguous, then we use an aliased tensor view of type input_t - // for that buffer where the last dimension is twice as large (because input_t is real - // and output_t is complex). We then use a slice to maintain the expected dimensions. - // If the output buffer is not contiguous, then we async allocate a temporary buffer. - // There is one caveat with this allocate: the batched fft implementation currently - // requires that all input pointers must be aligned to the corresponding complex type, - // which cannot be guaranteed to always be true for a real-valued tensor. This was - // not an issue for the reused output buffer because the output tensor is complex-valued, - // so we always have an even stride from one batch to the next. As a temporary workaround - // for the FFT alignment issue, we add one channel in the odd-channel case and use a - // slice to create a tensor view of only [0, num_channels-1]. This guarantees that we - // always stride by an even number of elements from one batch to the next while exposing - // a tensor view of appropriate dimensions. - using post_filter_t = typename inner_op_type_t::type; - auto fft_in_slice = [&out, &start_dims, &stop_dims, num_channels, stream]() -> auto { - auto fft_in_shape = out.Shape(); - if (out.IsContiguous()) { - fft_in_shape[OUT_RANK-1] *= 2; - auto fft_in = make_tensor(reinterpret_cast(out.Data()), fft_in_shape); - stop_dims[OUT_RANK-1] = num_channels; - return slice(fft_in, start_dims, stop_dims); - } else { - if (num_channels % 2 == 1) { - fft_in_shape[OUT_RANK-1]++; - stop_dims[OUT_RANK-1] = num_channels; - } - auto tmp = make_tensor(fft_in_shape, MATX_ASYNC_DEVICE_MEMORY, stream); - return slice(tmp, start_dims, stop_dims); - } - }(); - - if (decimation_factor == num_channels && detail::cpoly::ShouldUseSmem(out, in, f)) { - detail::cpoly::Smem(fft_in_slice, in, f, stream, out_elem_offset); - } else if (detail::cpoly::ShouldUseSmemTiled(out, in, f, decimation_factor)) { - detail::cpoly::SmemTiled(fft_in_slice, in, f, decimation_factor, stream, out_elem_offset); - } else { - detail::cpoly::Generic(fft_in_slice, in, f, decimation_factor, stream, out_elem_offset); - } - stop_dims[OUT_RANK-1] = (num_channels/2) + 1; - auto out_packed = slice(out, start_dims, stop_dims); - (out_packed = fft(fft_in_slice, num_channels)).run(stream); - detail::cpoly::UnpackDFT(out, stream); - } - } else { - // The fused-DFT kernel only supports the maximally decimated case (D == M). - // num_channels == 1 is degenerate (no channelization, trivial DFT) and - // the fused kernel's switch starts at N=2. Let num_channels==1 fall - // through to Smem / SmemTiled / Generic, all of which handle it - // correctly as a plain FIR. All four filter kernels honor out_elem_offset, - // so the windowed path uses the same kernel selection as the full grid. - if (decimation_factor == num_channels && num_channels >= 2 && - num_channels <= detail::cpoly::FusedChanThreshold) { - detail::cpoly::FusedChan(out, in, f, stream, out_elem_offset); - } else { - if (decimation_factor == num_channels && detail::cpoly::ShouldUseSmem(out, in, f)) { - detail::cpoly::Smem(out, in, f, stream, out_elem_offset); - } else if (detail::cpoly::ShouldUseSmemTiled(out, in, f, decimation_factor)) { - detail::cpoly::SmemTiled(out, in, f, decimation_factor, stream, out_elem_offset); - } else { - detail::cpoly::Generic(out, in, f, decimation_factor, stream, out_elem_offset); - } - // Specify FORWARD here to prevent any normalization after the ifft. We do not - // want any extra scaling on the output values. - (out = ifft(out, num_channels, FFTNorm::FORWARD)).run(stream); - } - } + detail::cpoly::DeviceAttrs attrs; + const auto plan = detail::cpoly::SelectPlan< + OutputOp, InputOp, FilterOp, AccumType>( + out, in, f, num_channels, decimation_factor, attrs); + detail::cpoly::ExecutePlan( + plan, out, in, f, num_channels, decimation_factor, stream, out_elem_offset, attrs); } /** @@ -913,52 +1289,19 @@ inline void channelize_poly_impl(OutType out, const InType &in, const FilterType using FilterOp = cuda::std::remove_cv_t>; using input_t = typename InputOp::value_type; using filter_t = typename FilterOp::value_type; - using output_t = typename OutputOp::value_type; using filtering_accum_t = cuda::std::conditional_t< is_complex_v || is_complex_v, typename detail::scalar_to_complex::ctype, AccumType>; using complex_accum_t = typename detail::scalar_to_complex::ctype; - static_assert(!is_complex_v, - "channelize_poly: accumulator type must be real; it will be treated as complex when necessary"); - constexpr int IN_RANK = InputOp::Rank(); constexpr int OUT_RANK = OutputOp::Rank(); + detail::cpoly::ValidateChannelizePolyArgs( + out, in, f, num_channels, decimation_factor, out_elem_offset); - MATX_STATIC_ASSERT_STR(OUT_RANK == IN_RANK+1, matxInvalidDim, - "channelize_poly: output rank should be 1 higher than input"); - MATX_STATIC_ASSERT_STR(is_complex_v || is_complex_half_v, - matxInvalidType, "channelize_poly: output type must be complex"); - MATX_STATIC_ASSERT_STR(FilterType::Rank() == 1, matxInvalidDim, - "channelize_poly: currently only support 1D filters"); - - MATX_ASSERT_STR(num_channels > 0, matxInvalidParameter, - "channelize_poly: num_channels must be positive"); - MATX_ASSERT_STR(decimation_factor > 0, matxInvalidParameter, - "channelize_poly: decimation_factor must be positive"); - MATX_ASSERT_STR(decimation_factor <= num_channels, matxInvalidParameter, - "channelize_poly: decimation_factor must be <= num_channels"); - - for (int i = 0; i < IN_RANK-1; i++) { - MATX_ASSERT_STR(out.Size(i) == in.Size(i), matxInvalidDim, - "channelize_poly: input/output must have matched batch sizes"); - } - - const index_t input_len = in.Size(IN_RANK-1); - [[maybe_unused]] const index_t num_elem_per_channel = - (input_len + decimation_factor - 1) / decimation_factor; - // Rows actually computed: the window [out_elem_offset, out_elem_offset + - // out_rows) of the full output-element grid (the full grid when not - // windowed). A zero offset does NOT imply the full grid -- the streaming - // channelizer's first feed emits only the fully-covered prefix rows - // (floor(chunk/D)) at offset 0 while the full local grid has ceil(chunk/D) - // rows -- so require only that the window fits within the grid. - const index_t out_rows = out.Size(OUT_RANK-2); - MATX_ASSERT_STR(out.Size(OUT_RANK-1) == num_channels, matxInvalidDim, - "channelize_poly: output size OUT_RANK-1 mismatch"); - MATX_ASSERT_STR(out_rows + out_elem_offset <= num_elem_per_channel, matxInvalidDim, - "channelize_poly: output-element window exceeds the full output size"); + const index_t input_len = in.Size(IN_RANK - 1); + const index_t out_rows = out.Size(OUT_RANK - 2); const index_t filter_full_len = f.Size(FilterOp::Rank()-1); const index_t filter_phase_len = (filter_full_len + num_channels - 1) / num_channels; @@ -984,8 +1327,7 @@ inline void channelize_poly_impl(OutType out, const InType &in, const FilterType // that drives the input footprint and polyphase phase (tg == t when not // windowed). const index_t tg = t + out_elem_offset; - const auto in_batch_idx = detail::BlockToIdx(in, batch, 1); - const auto out_batch_idx = detail::BlockToIdx(out, batch, 2); + auto [input_b, output_b, filter_acc] = detail::cpoly::MakeAccessors(out, in, f, batch); index_t thread_index = 0; #ifdef MATX_EN_OMP if (num_thread_buffers > 1) { @@ -1049,11 +1391,11 @@ inline void channelize_poly_impl(OutType out, const InType &in, const FilterType } for (index_t i = 0; i < niter; i++) { - const input_t in_val = detail::cpoly::HostReadSignal(in, in_batch_idx, sample_idx); - const filter_t h_val = f(h_ind); - detail::cpoly::HostChannelizeCmac(accum, - detail::cpoly::HostChannelizeCastFilter(h_val), - detail::cpoly::HostChannelizeCastInput(in_val)); + const input_t in_val = input_b(sample_idx); + const filter_t h_val = filter_acc(h_ind); + detail::channelize_cmac( + accum, detail::channelize_cast_operand(h_val), + detail::channelize_cast_operand(in_val)); h_ind += num_channels; sample_idx -= num_channels; } @@ -1068,7 +1410,7 @@ inline void channelize_poly_impl(OutType out, const InType &in, const FilterType filtered[static_cast(branch)]) * twiddles[static_cast(channel * num_channels + branch)]; } - detail::cpoly::HostWriteOutput(out, out_batch_idx, t, channel, dft); + output_b(t, channel) = static_cast(dft); } }; diff --git a/test/00_transform/ChannelizePoly.cu b/test/00_transform/ChannelizePoly.cu index 9747d11ef..e6fb3f1f9 100644 --- a/test/00_transform/ChannelizePoly.cu +++ b/test/00_transform/ChannelizePoly.cu @@ -37,11 +37,496 @@ #include "gtest/gtest.h" #include #include +#include #include +#include #include +#include using namespace matx; +TEST(ChannelizePoly, BatchedAndPartialRowsMatchHost) +{ + MATX_ENTER_HANDLER(); + cudaExecutor cuda_exec{}; + SingleThreadedHostExecutor host_exec{}; + auto check = [&](auto input_tag, auto real_tag) { + using Input = decltype(input_tag); + using Real = decltype(real_tag); + using Complex = cuda::std::complex; + struct Config { index_t m, d, p, n, b, tail = 1; }; + const index_t batches = sizeof(Input) <= 8 ? 128 : 64; + const Config cases[] = { + {128, 128, 13, 1027, 8}, + {40, 40, 13, 65537, batches}, + {40, 20, 13, 65537, batches}, + // Near-critical oversampling needs history across output-group reloads. + {64, 63, 1, 2051, 2}, + {64, 63, 4, 2051, 2}, + {64, 63, 13, 2051, 2}, + {31, 31, 13, 65537, 8}, + {16, 8, 128, 4099, 8}, + {16, 16, 13, 1023, 4}, + {16, 16, 13, 1039, 4}, + // Packed power-of-two rows: short filters, partial tiles, and batching. + {8, 8, 1, 1031, 3, 0}, + {8, 8, 4, 1031, 3}, + {8, 8, 21, 16385, 4}, + {8, 4, 21, 1031, 3}, + {16, 16, 21, 1031, 3}, + {16, 8, 21, 1031, 3}, + // Long signals with modest filters, partial rows, and batching. + {12, 12, 13, 5000003, 1}, + {16, 16, 13, 5000003, 1}, + {24, 24, 21, 5000003, 2}, + // Long batched filters and power-of-two shared-memory fit boundaries. + {32, 32, sizeof(Input) <= 4 ? 375 : sizeof(Input) <= 8 ? 183 : 87, + 65537, 64}, + {32, 32, 64, 65537, batches}, + {32, 16, 64, 65537, batches}, + {64, 64, 64, 65537, batches}, + {64, 32, 64, 65537, batches}, + // Short two-channel filters, partial rows, and batched output boundaries. + {2, 2, 1, 1027, 3}, + {2, 2, 2, 1027, 3}, + {2, 2, 4, 1027, 3}, + {2, 2, 5, 1027, 3}, + {2, 2, 13, 1027, 3}, + {2, 2, 2, 1027, 3, 0}, + {2, 2, 4, 1027, 3, 0}, + {2, 1, 2, 1027, 3}, + {2, 2, 939, 1027, 2}, + {2, 2, 940, 1027, 2}, + // Long real-double filters span the cached and fallback regimes. + {2, 2, 1280, 1027, 2}, + {2, 2, 1281, 1027, 2}, + {2, 2, 1472, 1027, 2}, + {2, 2, 1473, 1027, 2}, + {3, 3, 13, 65537, 2}, + {4, 4, 21, 16383, 1}, + {5, 5, 21, 16387, 4}, + {6, 6, 48, 16383, 1}, + // Direct small-channel FIRs: causal starts and full/partial filter tails. + {5, 5, 1, 1027, 3}, + {5, 5, 4, 1027, 3}, + {5, 5, 13, 1027, 3}, + {5, 5, 1, 1027, 3, 0}, + {5, 5, 4, 1027, 3, 0}, + {5, 5, 13, 1027, 3, 0}, + {6, 6, 1, 1027, 3}, + {6, 6, 4, 1027, 3}, + {6, 6, 13, 1027, 3}, + {6, 6, 1, 1027, 3, 0}, + {6, 6, 4, 1027, 3, 0}, + {6, 6, 13, 1027, 3, 0}, + }; + for (const auto &tc : cases) { + SCOPED_TRACE(::testing::Message() << "M=" << tc.m << " D=" << tc.d + << " P=" << tc.p << " N=" << tc.n << " B=" << tc.b + << " bytes=" << sizeof(Input) << " complex=" << is_complex_v); + const index_t rows = (tc.n + tc.d - 1) / tc.d; + auto input = make_tensor({tc.b, tc.n}); + auto filter = make_tensor({tc.m * tc.p - tc.tail}); + auto output = make_tensor({tc.b, rows, tc.m}); + auto reference = make_tensor({tc.b, 1, tc.m}); + for (index_t b = 0; b < tc.b; b++) { + for (index_t n = 0; n < tc.n; n++) { + const auto re = static_cast((n + 17 * b) % 97 - 48) / 64; + if constexpr (is_complex_v) { + const auto im = static_cast((n + 31 * b) % 89 - 44) / 64; + input(b, n) = Input(re, im); + } else { + input(b, n) = re; + } + } + } + for (index_t k = 0; k < filter.Size(0); k++) { + filter(k) = static_cast((5 * k) % 17 - 8) / 32; + } + (output = channelize_poly(input, filter, tc.m, tc.d)).run(cuda_exec); + MATX_CUDA_CHECK_LAST_ERROR(); + cuda_exec.sync(); + for (index_t row : {index_t{0}, index_t{1}, index_t{3}, index_t{4}, + index_t{7}, index_t{8}, index_t{15}, index_t{16}, + index_t{31}, index_t{32}, index_t{63}, index_t{64}, + index_t{127}, index_t{128}, rows / 2, rows - 1}) { + if (row >= rows) continue; + channelize_poly_impl(reference, input, filter, tc.m, tc.d, + host_exec, row); + for (index_t b = 0; b < tc.b; b++) { + for (index_t c = 0; c < tc.m; c++) { + const auto expected = reference(b, 0, c); + const double tolerance = (sizeof(Real) == 4 ? 2e-4 : 2e-11) * + (1.0 + cuda::std::abs(expected)); + EXPECT_NEAR(output(b, row, c).real(), expected.real(), tolerance); + EXPECT_NEAR(output(b, row, c).imag(), expected.imag(), tolerance); + } + } + } + } + }; + check(float{}, float{}); + check(double{}, double{}); + check(cuda::std::complex{}, float{}); + check(cuda::std::complex{}, double{}); + MATX_EXIT_HANDLER(); +} + +TEST(ChannelizePoly, LongPartialTileMatchesHost) +{ + MATX_ENTER_HANDLER(); + cudaExecutor cuda_exec{}; + SingleThreadedHostExecutor host_exec{}; + auto check = [&](auto input_tag, auto real_tag) { + using Input = decltype(input_tag); + using Real = decltype(real_tag); + using Complex = cuda::std::complex; + struct Config { index_t channels, rows; }; + const Config cases[] = { + {24, 32769}, {30, 32769}, {31, 32769}, {32, 32769}, + // Multiple short batches still cannot amortize per-block tile setup. + {31, 2049}, + }; + for (const auto &tc : cases) { + const index_t channels = tc.channels; + const index_t rows = tc.rows; + SCOPED_TRACE(::testing::Message() + << "M=" << channels << " rows=" << rows << " input_bytes=" << sizeof(Input) + << " complex=" << is_complex_v); + constexpr index_t batches = 2; + const index_t length = channels * rows - 7; + const index_t taps = sizeof(Input) == 4 ? 375 : + (sizeof(Input) == 8 ? 183 : 87); + auto input = make_tensor({batches, length}); + auto filter = make_tensor({channels * taps - 1}); + auto output = make_tensor({batches, rows, channels}); + auto reference = make_tensor({batches, 1, channels}); + for (index_t b = 0; b < batches; b++) { + for (index_t n = 0; n < length; n++) { + const auto re = static_cast(static_cast(n % 97 - 48 + b) / 64.0); + if constexpr (is_complex_v) { + input(b, n) = Input(re, static_cast(static_cast(n % 89 - 44) / 64.0)); + } else { + input(b, n) = re; + } + } + } + for (index_t k = 0; k < filter.Size(0); k++) { + filter(k) = static_cast(static_cast(k % 13 - 6) / 64.0); + } + (output = channelize_poly(input, filter, channels, channels)).run(cuda_exec); + MATX_CUDA_CHECK_LAST_ERROR(); + cuda_exec.sync(); + for (index_t row : {index_t{0}, index_t{15}, index_t{16}, + rows / 2, rows - 2, rows - 1}) { + channelize_poly_impl(reference, input, filter, channels, + channels, host_exec, row); + for (index_t b = 0; b < batches; b++) { + for (index_t c = 0; c < channels; c++) { + const auto expected = reference(b, 0, c); + const double tolerance = (sizeof(Real) == 4 ? 2e-4 : 2e-11) * + (1.0 + cuda::std::abs(expected)); + EXPECT_NEAR(output(b, row, c).real(), expected.real(), tolerance); + EXPECT_NEAR(output(b, row, c).imag(), expected.imag(), tolerance); + } + } + } + } + }; + check(float{}, float{}); + check(double{}, double{}); + check(cuda::std::complex{}, float{}); + check(cuda::std::complex{}, double{}); + MATX_EXIT_HANDLER(); +} + +// Every launchable backend plan must match the host implementation, whether or +// not the dispatcher selects it, for the full grid and for a window at a nonzero +// out_elem_offset. Plans that are not launchable must throw. +TEST(ChannelizePoly, ExplicitPlansMatchHost) +{ + MATX_ENTER_HANDLER(); + namespace cp = detail::cpoly; + cudaExecutor cuda_exec{}; + SingleThreadedHostExecutor host_exec{}; + auto check = [&](auto input_tag, auto real_tag) { + using Input = decltype(input_tag); + using Real = decltype(real_tag); + using Complex = cuda::std::complex; + struct Config { index_t m, d, p, n, b; }; + const Config cases[] = { + {4, 4, 13, 4099, 2}, + {5, 3, 7, 4099, 2}, + {8, 6, 9, 4099, 2}, + {10, 10, 21, 8195, 2}, + {16, 16, 21, 8195, 2}, + {20, 15, 13, 8195, 2}, + {33, 22, 6, 16387, 3}, + {40, 40, 13, 16387, 2}, + {64, 64, 8, 16387, 2}, + {96, 72, 5, 16387, 2}, + {100, 100, 17, 16387, 2}, + // The FIR tile exceeds shared memory, so the fused plan runs the direct FIR. + {5, 5, 2048, 8195, 2}, + }; + for (const auto &tc : cases) { + SCOPED_TRACE(::testing::Message() << "M=" << tc.m << " D=" << tc.d + << " P=" << tc.p << " N=" << tc.n << " B=" << tc.b + << " bytes=" << sizeof(Input) << " complex=" << is_complex_v); + const index_t rows = (tc.n + tc.d - 1) / tc.d; + auto input_storage = make_tensor({tc.b, tc.n}); + auto filter = make_tensor({tc.m * tc.p - 1}); + auto output_storage = make_tensor({tc.b, rows, tc.m}); + auto reference = make_tensor({tc.b, 1, tc.m}); + for (index_t b = 0; b < tc.b; b++) { + for (index_t n = 0; n < tc.n; n++) { + const auto re = static_cast((n + 17 * b) % 97 - 48) / 64; + if constexpr (is_complex_v) { + const auto im = static_cast((n + 31 * b) % 89 - 44) / 64; + input_storage(b, n) = Input(re, im); + } else { + input_storage(b, n) = re; + } + } + } + for (index_t k = 0; k < filter.Size(0); k++) { + filter(k) = static_cast((5 * k) % 17 - 8) / 32; + } + // Match the argument types passed by the channelize_poly operator. + detail::base_type_t input = input_storage; + detail::base_type_t output = output_storage; + using OutOp = decltype(output); + using InOp = decltype(input); + using FilterOp = decltype(filter); + + const index_t sample_rows[] = {0, 1, 2, 15, 16, 17, rows / 2, rows - 2, rows - 1}; + // Windowed launches, as issued by the streaming channelizer, cover rows + // [window_start, rows - 1): past the causal start and before the final row. + const index_t window_start = 17; + const index_t window_rows = rows - 1 - window_start; + auto window_storage = make_tensor({tc.b, window_rows, tc.m}); + detail::base_type_t window = window_storage; + std::vector expected; + for (index_t row : sample_rows) { + channelize_poly_impl(reference, input_storage, filter, tc.m, + tc.d, host_exec, row); + for (index_t b = 0; b < tc.b; b++) { + for (index_t c = 0; c < tc.m; c++) { + expected.push_back(reference(b, 0, c)); + } + } + } + + std::vector plans; + cp::Plan plan; + plan.backend = cp::Backend::Fused; + plans.push_back(plan); + plan.backend = cp::Backend::Smem; + plans.push_back(plan); + // The last row count is not supported and must throw. + for (int nout : {cp::SmemTiledMaxDecNout, cp::SmemTiledNout, 2 * cp::SmemTiledNout}) { + for (auto layout : {cp::SmemTiledFilterLayout::Full, + cp::SmemTiledFilterLayout::Rotated, + cp::SmemTiledFilterLayout::Global}) { + plan.backend = cp::Backend::Tiled; + plan.tiled = cp::SmemTiledPlanWithLayout(output, input, filter, tc.d, nout, layout); + plans.push_back(plan); + } + } + plans.push_back(cp::Plan{}); + const auto selected = cp::SelectPlan( + output, input, filter, tc.m, tc.d); + ASSERT_TRUE((cp::PlanIsLaunchable( + selected, output, input, filter, tc.m, tc.d))); + plans.push_back(selected); + + int launched = 0; + for (const auto &p : plans) { + SCOPED_TRACE(::testing::Message() << "backend=" + << static_cast(p.backend) << " nout=" << p.tiled.nout + << " layout=" << static_cast(p.tiled.filter_layout)); + if (!cp::PlanIsLaunchable( + p, output, input, filter, tc.m, tc.d)) { + EXPECT_THROW((cp::ExecutePlan( + p, output, input, filter, tc.m, tc.d, cuda_exec.getStream())), detail::matxException); + continue; + } + launched++; + for (const bool windowed : {false, true}) { + SCOPED_TRACE(windowed ? "windowed" : "full grid"); + auto &storage = windowed ? window_storage : output_storage; + const index_t first_row = windowed ? window_start : 0; + const index_t row_count = windowed ? window_rows : rows; + (storage = Complex{9, -9}).run(cuda_exec); + cp::ExecutePlan( + p, windowed ? window : output, input, filter, tc.m, tc.d, cuda_exec.getStream(), + first_row); + MATX_CUDA_CHECK_LAST_ERROR(); + cuda_exec.sync(); + size_t i = 0; + for (index_t row : sample_rows) { + for (index_t b = 0; b < tc.b; b++) { + for (index_t c = 0; c < tc.m; c++, i++) { + if (row < first_row || row >= first_row + row_count) continue; + const auto got = storage(b, row - first_row, c); + const double tolerance = (sizeof(Real) == 4 ? 2e-4 : 2e-11) * + (1.0 + cuda::std::abs(expected[i])); + EXPECT_NEAR(got.real(), expected[i].real(), tolerance) << "row=" << row; + EXPECT_NEAR(got.imag(), expected[i].imag(), tolerance) << "row=" << row; + } + } + } + } + } + // Generic and at least one tiled layout are always launchable here. + EXPECT_GE(launched, 2); + } + }; + check(float{}, float{}); + check(double{}, double{}); + check(cuda::std::complex{}, float{}); + check(cuda::std::complex{}, double{}); + MATX_EXIT_HANDLER(); +} + +// The dispatcher must always select a launchable plan. +TEST(ChannelizePoly, SelectedPlansAreLaunchable) +{ + MATX_ENTER_HANDLER(); + namespace cp = detail::cpoly; + auto check = [&](auto input_tag, auto real_tag) { + using Input = decltype(input_tag); + using Real = decltype(real_tag); + using Complex = cuda::std::complex; + // Plan selection reads only shapes and strides; no kernel is launched, + // so non-owning views over a small allocation suffice. + auto storage = make_tensor({1}, MATX_DEVICE_MEMORY); + for (index_t m : {1, 2, 3, 6, 7, 8, 10, 16, 17, 32, 64, 65, 80, 256, 1000}) { + for (index_t d : {m, (m + 1) / 2, index_t{1}}) { + for (index_t p : {1, 4, 64, 400}) { + for (index_t n : {index_t{1000}, index_t{100000}}) { + const index_t rows = (n + d - 1) / d; + auto input = make_tensor(reinterpret_cast(storage.Data()), {2, n}); + auto filter = make_tensor(reinterpret_cast(storage.Data()), {m * p}); + auto output = make_tensor(storage.Data(), {2, rows, m}); + const auto plan = cp::SelectPlan(output, input, filter, m, d); + EXPECT_TRUE((cp::PlanIsLaunchable(plan, output, input, filter, m, d))) + << "M=" << m << " D=" << d << " P=" << p << " N=" << n; + } + } + } + } + }; + check(float{}, float{}); + check(double{}, double{}); + check(cuda::std::complex{}, float{}); + check(cuda::std::complex{}, double{}); + MATX_EXIT_HANDLER(); +} + +// Outputs with more rows than the grid.y limit allows: Generic splits its time +// blocks into several launches and the real-input DFT unpack strides over rows. +// M=3 and M=4 give the unpack conjugate/mirror work (M=2 has none), and the +// checked rows straddle both kernels' grid.y boundaries. +TEST(ChannelizePoly, TallGridsMatchHost) +{ + MATX_ENTER_HANDLER(); + namespace cp = detail::cpoly; + cudaExecutor cuda_exec{}; + SingleThreadedHostExecutor host_exec{}; + constexpr index_t decimation = 1; + constexpr index_t max_grid_y = 65535; + constexpr index_t generic_rows = 256; // Generic time rows per CTA + constexpr index_t unpack_rows = 128; // UnpackDFT rows per CTA + const index_t n = max_grid_y * generic_rows + 3 * generic_rows + 5; + const index_t unpack_edge = max_grid_y * unpack_rows; + const index_t generic_edge = max_grid_y * generic_rows; + for (index_t channels : {index_t{3}, index_t{4}}) { + SCOPED_TRACE(::testing::Message() << "M=" << channels); + auto input = make_tensor({1, n}); + auto filter = make_tensor({channels * 4 - 1}); + auto output = make_tensor>({1, n, channels}); + auto reference = make_tensor>({1, 1, channels}); + for (index_t k = 0; k < n; k++) { + input(0, k) = static_cast(k % 97 - 48) / 64; + } + for (index_t k = 0; k < filter.Size(0); k++) { + filter(k) = static_cast((5 * k) % 17 - 8) / 32; + } + (output = cuda::std::complex{9, -9}).run(cuda_exec); + cp::ExecutePlan( + cp::Plan{}, output, input, filter, channels, decimation, cuda_exec.getStream()); + MATX_CUDA_CHECK_LAST_ERROR(); + cuda_exec.sync(); + for (index_t row : {index_t{0}, unpack_edge - 1, unpack_edge, unpack_edge + 1, + generic_edge - 1, generic_edge, generic_edge + 1, n - 1}) { + channelize_poly_impl(reference, input, filter, channels, decimation, host_exec, row); + for (index_t c = 0; c < channels; c++) { + const auto expected = reference(0, 0, c); + const double tolerance = 2e-4 * (1.0 + cuda::std::abs(expected)); + EXPECT_NEAR(output(0, row, c).real(), expected.real(), tolerance) + << "row=" << row << " c=" << c; + EXPECT_NEAR(output(0, row, c).imag(), expected.imag(), tolerance) + << "row=" << row << " c=" << c; + } + } + } + MATX_EXIT_HANDLER(); +} + +// An empty signal, an empty batch dimension, or a zero-row output window leaves +// nothing to compute and must not launch. The configurations span the fused, +// Smem, tiled, and Generic dispatch families for real and complex input. +TEST(ChannelizePoly, EmptyOutputsAreNoOps) +{ +#ifndef NDEBUG + // Debug builds reject zero-sized tensors at construction, so empty outputs + // can reach channelize_poly only when assertions are compiled out. + GTEST_SKIP() << "Zero-sized tensors require NDEBUG"; +#else + MATX_ENTER_HANDLER(); + cudaExecutor exec{}; + auto check = [&](auto input_tag) { + using Input = decltype(input_tag); + using Complex = cuda::std::complex; + struct Config { index_t m, d; }; + const Config cases[] = {{4, 4}, {8, 8}, {40, 20}, {12, 12}, {64, 64}, {12, 6}}; + for (const auto &tc : cases) { + SCOPED_TRACE(::testing::Message() << "M=" << tc.m << " D=" << tc.d + << " complex=" << is_complex_v); + const index_t n = 4099; + const index_t rows = (n + tc.d - 1) / tc.d; + auto filter = make_tensor({13 * tc.m - 1}); + + auto empty_signal = make_tensor({0}); + auto empty_output = make_tensor({0, tc.m}); + (empty_output = channelize_poly(empty_signal, filter, tc.m, tc.d)).run(exec); + + auto empty_batches = make_tensor({0, n}); + auto empty_batch_output = make_tensor({0, rows, tc.m}); + (empty_batch_output = channelize_poly(empty_batches, filter, tc.m, tc.d)).run(exec); + + auto signal = make_tensor({n}); + auto window = make_tensor({0, tc.m}); + channelize_poly_impl( + window, signal, filter, tc.m, tc.d, exec.getStream(), 3); + + exec.sync(); + ASSERT_EQ(cudaGetLastError(), cudaSuccess); + } + }; + check(float{}); + check(cuda::std::complex{}); + MATX_EXIT_HANDLER(); +#endif +} + template class ChannelizePolyTest : public ::testing::Test { using GTestType = cuda::std::tuple_element_t<0, T>; @@ -329,7 +814,12 @@ TYPED_TEST(ChannelizePolyTestNonHalfFloatTypes, Simple) { 27137, 301*13+3, 13 }, { 27138, 301*14+4, 14 }, { 1000000, 32*16, 32 }, - { 1000000, 40*16, 40 } + { 1000000, 40*16, 40 }, + // Fused radix/power-of-two paths with partial input and filter rows. + { 100003, 20*17+7, 20 }, + { 100003, 40*17+7, 40 }, + { 100003, 64*17+7, 64 }, + { 100003, 80*17+7, 80 } }; for (size_t i = 0; i < sizeof(test_cases)/sizeof(test_cases[0]); i++) { @@ -535,10 +1025,8 @@ TYPED_TEST(ChannelizePolyTestDoubleType, AccumProperty) "00_transforms", "channelize_poly_operators", "channelize", {a_len, f_len, num_channels, num_channels}); auto a64 = make_tensor({a_len}); auto f64 = make_tensor({f_len}); - auto gold_output_hreal = make_tensor>({b_len_per_channel, num_channels}); this->pb->NumpyToTensorView(a64, "a"); this->pb->NumpyToTensorView(f64, "filter_random"); - this->pb->NumpyToTensorView(gold_output_hreal, "b_random_hreal"); auto a32 = make_tensor({a_len}); (a32 = as_float(a64)).run(this->exec); @@ -581,11 +1069,11 @@ TYPED_TEST(ChannelizePolyTestDoubleType, AccumProperty) } this->pb->template InitAndRunTVGenerator>( - "00_transforms", "channelize_poly_operators", "channelize", {a_len, f_len, num_channels, num_channels}); + "00_transforms", "channelize_poly_operators", "channelize_accum", + {a_len, f_len, num_channels, num_channels}); auto ac64 = make_tensor>({a_len}); this->pb->NumpyToTensorView(ac64, "a"); this->pb->NumpyToTensorView(f64, "filter_random_real"); - this->pb->NumpyToTensorView(gold_output_hreal, "b_random_hreal"); // The following cases are all for complex inputs and real filters @@ -593,6 +1081,11 @@ TYPED_TEST(ChannelizePolyTestDoubleType, AccumProperty) (ac32 = as_complex_float(ac64)).run(this->exec); (f32 = as_float(f64)).run(this->exec); + // Python computes FP64 gold from the same rounded FP32 operands, separating + // input quantization from accumulation error without using MatX as gold. + auto quantized_gold = make_tensor>({b_len_per_channel, num_channels}); + this->pb->NumpyToTensorView(quantized_gold, "b_quantized_hreal"); + // Below, we test that using a double-precision accumulator by running with fp32 inputs with and // without a double-precision accumulator demonstrates higher accuracy when an fp64 accumulator. // Note that for the single-precision test, we need to write to a single-precision output or the @@ -600,19 +1093,19 @@ TYPED_TEST(ChannelizePolyTestDoubleType, AccumProperty) auto max_err = make_tensor({}); double max_err_fp32{}; double max_err_fp32_in_fp64_accum{}; - - // All single precision + // All single precision, using tensor filters in both accuracy cases. { auto chan_poly = channelize_poly(ac32, f32, num_channels, decimation_factor); (b32 = chan_poly).run(this->exec); cudaStreamSynchronize(stream); MATX_TEST_ASSERT_COMPARE(this->pb, b32, "b_random_hreal", mixed_thresh_complex_input); - (max_err = matx::max(matx::abs(as_complex_double(b32) - gold_output_hreal), {0,1})).run(this->exec); + (max_err = matx::max( + matx::abs(as_complex_double(b32) - quantized_gold), {0,1})).run(this->exec); cudaStreamSynchronize(stream); max_err_fp32 = max_err(); } - // Single precision complex input, output, and filter, double precision accumulator + // Single precision complex input and filter, double accumulator and output. { auto chan_poly = channelize_poly(ac32, f32, num_channels, decimation_factor) .props, PropOutput>>(); @@ -622,7 +1115,7 @@ TYPED_TEST(ChannelizePolyTestDoubleType, AccumProperty) (b64 = chan_poly + cuda::std::complex(0.0, 0.0)).run(this->exec); cudaStreamSynchronize(stream); MATX_TEST_ASSERT_COMPARE(this->pb, b64, "b_random_hreal", mixed_thresh_complex_input); - (max_err = matx::max(matx::abs(b64 - gold_output_hreal), {0,1})).run(this->exec); + (max_err = matx::max(matx::abs(b64 - quantized_gold), {0,1})).run(this->exec); cudaStreamSynchronize(stream); max_err_fp32_in_fp64_accum = max_err(); } @@ -782,6 +1275,536 @@ TYPED_TEST(ChannelizePolyTestNonHalfFloatTypes, IdentityFilter) MATX_EXIT_HANDLER(); } +TEST(ChannelizePoly, LongCriticalFilterMatchesHost) +{ + MATX_ENTER_HANDLER(); + + constexpr index_t M = 32; + constexpr index_t P = 184; + constexpr index_t input_len = 6400; + constexpr index_t output_len = (input_len + M - 1) / M; + using complex_t = cuda::std::complex; + cudaExecutor cuda_exec{}; + SingleThreadedHostExecutor host_exec{}; + auto input = make_tensor({input_len}); + auto filter = make_tensor({M * P}); + auto cuda_output = make_tensor({output_len, M}); + auto host_output = make_tensor({output_len, M}); + + for (index_t n = 0; n < input_len; n++) { + const double x = static_cast(n); + input(n) = complex_t{static_cast(std::sin(0.037 * x)), + static_cast(std::cos(0.053 * x))}; + } + for (index_t k = 0; k < M * P; k++) { + const double x = static_cast(k); + filter(k) = static_cast(std::cos(0.071 * x) * std::exp(-0.0013 * x)); + } + + (cuda_output = channelize_poly(input, filter, M, M)).run(cuda_exec); + cuda_exec.sync(); + (host_output = channelize_poly(input, filter, M, M)).run(host_exec); + + for (index_t t = 0; t < output_len; t++) { + for (index_t c = 0; c < M; c++) { + const auto expected = host_output(t, c); + const auto got = cuda_output(t, c); + const double scale = 1.0 + cuda::std::abs(expected); + EXPECT_NEAR(got.real(), expected.real(), 2e-3 * scale); + EXPECT_NEAR(got.imag(), expected.imag(), 2e-3 * scale); + } + } + + MATX_EXIT_HANDLER(); +} + +TEST(ChannelizePoly, OperatorInput) +{ + MATX_ENTER_HANDLER(); + + using complex_t = cuda::std::complex; + constexpr index_t batches = 2; + constexpr index_t input_len = 1027; + constexpr index_t taps_per_channel = 13; + struct TestCase { + index_t num_channels; + index_t decimation_factor; + }; + const TestCase test_cases[] = {{4, 4}, {5, 3}}; + cudaExecutor exec{}; + + for (const auto &tc : test_cases) { + const index_t filter_len = taps_per_channel * tc.num_channels - 3; + const index_t output_len = (input_len + tc.decimation_factor - 1) / tc.decimation_factor; + auto input_iq = make_tensor({batches, 2 * input_len}); + auto phase_ramp = make_tensor({batches, input_len}); + auto filter = make_tensor({filter_len}); + auto materialized = make_tensor({batches, input_len}); + auto output = make_tensor({batches, output_len, tc.num_channels}); + auto reference = make_tensor({batches, output_len, tc.num_channels}); + + for (index_t b = 0; b < batches; b++) { + for (index_t n = 0; n < input_len; n++) { + input_iq(b, 2 * n) = static_cast((n + 3 * b) % 31 - 15); + input_iq(b, 2 * n + 1) = static_cast((2 * n + 5 * b) % 29 - 14); + const float angle = 0.013f * static_cast(n + 7 * b); + phase_ramp(b, n) = complex_t{std::cos(angle), std::sin(angle)}; + } + } + for (index_t k = 0; k < filter_len; k++) { + const float x = static_cast(k); + filter(k) = std::cos(0.071f * x) * std::exp(-0.013f * x); + } + + auto i_values = slice(input_iq, {0, 0}, {batches, 2 * input_len}, {1, 2}); + auto q_values = slice(input_iq, {0, 1}, {batches, 2 * input_len}, {1, 2}); + auto input_op = phase_ramp * as_complex_float(i_values, q_values); + + (materialized = input_op).run(exec); + (output = channelize_poly(input_op, filter, tc.num_channels, tc.decimation_factor)).run(exec); + // Compare with materialized input and an equivalent filter expression. + (reference = channelize_poly( + materialized, filter * 1.0f, tc.num_channels, tc.decimation_factor)).run(exec); + exec.sync(); + + for (index_t b = 0; b < batches; b++) { + for (index_t t = 0; t < output_len; t++) { + for (index_t c = 0; c < tc.num_channels; c++) { + const complex_t expected = reference(b, t, c); + const complex_t got = output(b, t, c); + const float scale = 1.0f + cuda::std::abs(expected); + EXPECT_NEAR(got.real(), expected.real(), 5e-4f * scale); + EXPECT_NEAR(got.imag(), expected.imag(), 5e-4f * scale); + } + } + } + } + + MATX_EXIT_HANDLER(); +} + +// Compare explicit fused launches with the public channelizer using an equivalent filter +// expression. For critical M<=6, the reference can also use the fused family. Cover both +// precisions, complex input with a matched real filter, critical sampling, and arbitrary +// integer/rational decimation. Exercise M=2 through the public dispatcher; launch the remaining +// family directly so every leaf remains covered even on GPUs where the working-set or FP64 +// performance heuristic prefers the fallback. +TEST(ChannelizePoly, FusedRadixPow2MatchedRealFilter) +{ + MATX_ENTER_HANDLER(); + + struct TestCase { + index_t num_channels; + index_t decimation_factor; + index_t taps_per_channel; + }; + std::vector test_cases = { + { 2, 2, 64 }, + { 2, 1, 64 }, + // Short critical filters exercise the direct small-M leaf. + { 3, 3, 4 }, + { 4, 4, 4 }, + { 5, 5, 4 }, + { 6, 6, 4 }, + // Longer critical filters exercise its shared-memory-cached leaf. + { 3, 3, 17 }, + { 3, 2, 17 }, + { 4, 4, 17 }, + { 4, 3, 17 }, + { 5, 5, 17 }, + { 5, 2, 17 }, + { 6, 6, 17 }, + { 6, 5, 17 }, + { 8, 8, 17 }, + { 8, 3, 17 }, + { 10, 10, 17 }, + { 10, 4, 17 }, + { 16, 16, 17 }, + { 16, 7, 17 }, + { 20, 20, 17 }, + { 20, 8, 17 }, + { 32, 32, 17 }, + { 32, 13, 17 }, + { 40, 40, 17 }, + { 40, 17, 17 }, + { 64, 64, 17 }, + { 64, 27, 17 }, + { 80, 80, 17 }, + { 80, 33, 17 }, + }; + // Single-tap branches cover circular-buffer wrap without a filter halo. + for (index_t channels : {8, 10, 16, 20, 32, 40, 64, 80}) { + test_cases.push_back({channels, channels, 1}); + test_cases.push_back({channels, channels / 2, 1}); + } + // Cover every small-channel oversampling factor with modest filter lengths. + for (index_t channels : {3, 4, 5, 6}) { + for (index_t decimation = 1; decimation < channels; decimation++) { + for (index_t taps : {4, 13, 21}) { + test_cases.push_back({channels, decimation, taps}); + } + } + } + + auto run = [&]() { + using Complex = cuda::std::complex; + constexpr index_t batches = 2; + constexpr index_t input_len = 2051; + const double tolerance = cuda::std::is_same_v + ? 3e-10 : 4e-4; + cudaExecutor exec{}; + + for (const auto &tc : test_cases) { + const index_t filter_len = tc.taps_per_channel * tc.num_channels - 3; + const index_t output_len = (input_len + tc.decimation_factor - 1) / tc.decimation_factor; + auto input = make_tensor({batches, input_len}); + auto filter = make_tensor({filter_len}); + auto fused = make_tensor({batches, output_len, tc.num_channels}); + auto reference = make_tensor({batches, output_len, tc.num_channels}); + + for (index_t b = 0; b < batches; b++) { + for (index_t n = 0; n < input_len; n++) { + const double x = static_cast(n + 11 * b); + input(b, n) = Complex{ + static_cast(std::sin(0.037 * x)), static_cast(std::cos(0.053 * x))}; + } + } + for (index_t k = 0; k < filter_len; k++) { + const double x = static_cast(k); + filter(k) = static_cast(std::cos(0.071 * x) * std::exp(-0.013 * x)); + } + + if (tc.num_channels == 2) { + (fused = channelize_poly(input, filter, tc.num_channels, tc.decimation_factor)).run(exec); + } else { + // Bypass performance heuristics, but retain the launchability checks. + using FusedOp = decltype(fused); + using InputOp = decltype(input); + using FilterOp = decltype(filter); + ASSERT_TRUE((detail::cpoly::FusedRadixFeasible( + fused, input, filter, tc.num_channels, tc.decimation_factor))); + detail::cpoly::RunFusedRadix( + fused, input, filter, tc.num_channels, tc.decimation_factor, exec.getStream(), 0); + } + ASSERT_EQ(cudaGetLastError(), cudaSuccess); + ASSERT_EQ(cudaStreamSynchronize(exec.getStream()), cudaSuccess); + if (tc.num_channels <= 6 && tc.decimation_factor == tc.num_channels) { + // Expression filters also fuse for small critical channelizers. + (reference = channelize_poly( + input, filter, tc.num_channels, tc.decimation_factor)) + .run(SingleThreadedHostExecutor{}); + } else { + (reference = channelize_poly( + input, filter * static_cast(1), tc.num_channels, + tc.decimation_factor)).run(exec); + ASSERT_EQ(cudaGetLastError(), cudaSuccess); + ASSERT_EQ(cudaStreamSynchronize(exec.getStream()), cudaSuccess); + } + + for (index_t b = 0; b < batches; b++) { + for (index_t t = 0; t < output_len; t++) { + for (index_t c = 0; c < tc.num_channels; c++) { + const auto expected = reference(b, t, c); + const auto got = fused(b, t, c); + const double scale = 1.0 + cuda::std::abs(expected); + EXPECT_NEAR(static_cast(got.real()), + static_cast(expected.real()), + tolerance * scale) + << "real mismatch: M=" << tc.num_channels + << " D=" << tc.decimation_factor << " t=" << t + << " c=" << c; + EXPECT_NEAR(static_cast(got.imag()), + static_cast(expected.imag()), + tolerance * scale) + << "imag mismatch: M=" << tc.num_channels + << " D=" << tc.decimation_factor << " t=" << t + << " c=" << c; + } + } + } + } + }; + + run.template operator()(); + run.template operator()(); + + MATX_EXIT_HANDLER(); +} + +TEST(ChannelizePoly, FusedRadixPow2Strided) +{ + MATX_ENTER_HANDLER(); + constexpr index_t batches = 2; + constexpr index_t input_len = 1231; + constexpr index_t offset = 3; + constexpr index_t output_len = 73; + cudaExecutor cuda_exec{}; + SingleThreadedHostExecutor host_exec{}; + + auto run = [&]() { + static_assert(offset + output_len <= (input_len + Decimation - 1) / Decimation); + using Complex = cuda::std::complex; + using Reference = cuda::std::complex; + using config = detail::cpoly::FusedRadixConfig; + constexpr index_t filter_len = 13 * Channels - 1; + constexpr index_t rows_per_cta = 2 * config::NRows; + const double tolerance = sizeof(Scalar) == sizeof(float) ? 1e-4 : 1e-10; + for (const auto strides : {std::array{1, 1, 1}, + {2, 1, 1}, {1, 2, 1}, {1, 1, 2}, {2, 2, 2}}) { + const auto [input_stride, filter_stride, output_stride] = strides; + SCOPED_TRACE(::testing::Message() + << "M=" << Channels << " D=" << Decimation + << " strides=" << input_stride << ',' << filter_stride << ',' + << output_stride << " scalar bytes=" << sizeof(Scalar)); + auto input_storage = make_tensor({batches, input_stride * input_len}); + auto filter_storage = make_tensor({filter_stride * filter_len}); + auto output_storage = make_tensor({batches, output_len, output_stride * Channels}); + auto input = slice(input_storage, {0, 0}, + {batches, input_stride * input_len}, {1, input_stride}); + auto filter = slice(filter_storage, {0}, {filter_stride * filter_len}, {filter_stride}); + auto output = slice(output_storage, {0, 0, 0}, + {batches, output_len, output_stride * Channels}, + {1, 1, output_stride}); + auto reference = make_tensor({batches, output_len, index_t{Channels}}); + (input_storage = Complex{9, -9}).run(host_exec); + (filter_storage = Scalar{9}).run(host_exec); + for (index_t b = 0; b < batches; b++) { + for (index_t n = 0; n < input_len; n++) { + input(b, n) = {Scalar(((n + b) % 17) - 8) / Scalar{32}, + Scalar(((n + b) % 13) - 6) / Scalar{32}}; + } + } + for (index_t k = 0; k < filter_len; k++) { + filter(k) = Scalar((k % 11) - 5) / Scalar{64}; + } + channelize_poly_impl(reference, input, filter, Channels, + Decimation, host_exec, offset); + + using OutOp = decltype(output); + using InOp = decltype(input); + using FilterOp = decltype(filter); + if (strides != std::array{1, 1, 1}) { + EXPECT_THROW((detail::cpoly::FusedRadixImpl< + Channels, Decimation == Channels, OutOp, InOp, FilterOp, Scalar>( + output, input, filter, Decimation, cuda_exec.getStream(), + offset)), detail::matxException); + } + for (bool direct : {true, false}) { + (output_storage = Complex{9, -9}).run(host_exec); + if (direct) { + // Explicit strided kernels must honor their accessor flag even + // though the public grouped dispatcher retains its FIR/FFT fallback. + ASSERT_TRUE((detail::cpoly::FusedRadixFits< + Channels, OutOp, InOp, FilterOp, Scalar>( + output, input, filter, Decimation))); + const size_t smem = detail::cpoly::FusedRadixSizeBytes< + Channels, OutOp, InOp, FilterOp, Scalar>( + output, input, filter, Decimation); + const dim3 grid(static_cast( + (output_len + rows_per_cta - 1) / rows_per_cta), 1, batches); + ChannelizePoly1D_FusedRadixPow2 + <<>>( + output, input, filter, rows_per_cta, Decimation, offset); + } else { + channelize_poly_impl( + output, input, filter, Channels, Decimation, cuda_exec.getStream(), offset); + } + ASSERT_EQ(cudaGetLastError(), cudaSuccess); + ASSERT_EQ(cudaStreamSynchronize(cuda_exec.getStream()), cudaSuccess); + for (index_t b = 0; b < batches; b++) { + for (index_t t = 0; t < output_len; t++) { + for (index_t c = 0; c < Channels; c++) { + const auto expected = reference(b, t, c); + const auto got = output(b, t, c); + const double scale = 1.0 + cuda::std::abs(expected); + EXPECT_NEAR(static_cast(got.real()), expected.real(), tolerance * scale); + EXPECT_NEAR(static_cast(got.imag()), expected.imag(), tolerance * scale); + if (output_stride == 2) { + EXPECT_EQ(output_storage(b, t, 2 * c + 1), (Complex{9, -9})); + } + } + } + } + } + } + }; + run.template operator()(); + run.template operator()(); + run.template operator()(); + run.template operator()(); + run.template operator()(); + run.template operator()(); + run.template operator()(); + run.template operator()(); + run.template operator()(); + run.template operator()(); + run.template operator()(); + run.template operator()(); + MATX_EXIT_HANDLER(); +} + +TEST(ChannelizePoly, CriticalSmallChannelsStrided) +{ + MATX_ENTER_HANDLER(); + + cudaExecutor cuda_exec{}; + SingleThreadedHostExecutor host_exec{}; + + auto run = [&](index_t batches = 2, + index_t input_len = 1027, index_t offset = 0, index_t count = 0) { + using Complex = typename detail::scalar_to_complex::ctype; + using Input = cuda::std::conditional_t; + using Reference = cuda::std::complex; + const double tolerance = is_matx_half_v ? 2e-2 : + (cuda::std::is_same_v ? 1e-10 : 1e-4); + for (index_t num_channels : {2, 3, 4, 5, 6}) { + if (count != 0 && num_channels != 2) continue; + for (index_t taps_per_channel : {3, 13, 4096}) { + // Also cover half transforms that cannot fall back to cuFFT when + // their filter is too large to cache in shared memory. + if (taps_per_channel == 4096 && + (!is_matx_half_v || num_channels == 2 || + num_channels == 4)) continue; + for (index_t stride : {1, 2}) { + SCOPED_TRACE(::testing::Message() + << "M=" << num_channels << " P=" << taps_per_channel + << " stride=" << stride << " scalar bytes=" << sizeof(Scalar) + << " B=" << batches << " offset=" << offset << " count=" << count); + // Odd M=2 taps-per-channel counts exercise input-alignment padding. + const index_t filter_len = taps_per_channel * num_channels - 1; + const index_t output_len = (input_len + num_channels - 1) / num_channels; + const index_t window_len = count == 0 ? output_len : count; + ASSERT_LE(offset + window_len, output_len); + auto input_storage = make_tensor({batches, stride * input_len}); + auto filter_storage = make_tensor({stride * filter_len}); + auto output_storage = make_tensor({batches, output_len, stride * num_channels}); + auto input = slice(input_storage, {0, 0}, + {batches, stride * input_len}, {1, stride}); + auto filter = slice(filter_storage, {0}, {stride * filter_len}, {stride}); + auto output = slice(output_storage, {0, offset, 0}, + {batches, offset + window_len, stride * num_channels}, + {1, 1, stride}); + auto input_ref = make_tensor({batches, input_len}); + auto filter_ref = make_tensor({filter_len}); + auto reference = make_tensor({batches, output_len, num_channels}); + + (input_storage = Input{9.0f}).run(host_exec); + (filter_storage = Scalar{9.0f}).run(host_exec); + (output_storage = Complex{9.0f, -9.0f}).run(host_exec); + for (index_t b = 0; b < batches; b++) { + for (index_t n = 0; n < input_len; n++) { + const float real = static_cast(((n + b) % 17) - 8) / 32.0f; + const float imag = static_cast(((n + b) % 13) - 6) / 32.0f; + if constexpr (ComplexInput) { + input(b, n) = Complex{real, imag}; + input_ref(b, n) = {real, imag}; + } else { + input(b, n) = real; + input_ref(b, n) = {real, 0.0}; + } + } + } + for (index_t k = 0; k < filter_len; k++) { + const float value = static_cast((k % 11) - 5) / 64.0f; + filter(k) = Scalar{value}; + filter_ref(k) = value; + } + + if (count == 0) { + auto op = channelize_poly(input, filter, num_channels, num_channels); + if constexpr (DoubleAccum) { + (output = op.template props>()).run(cuda_exec); + } else { + (output = op).run(cuda_exec); + } + } else if constexpr (cuda::std::is_same_v && + ComplexInput && !DoubleAccum) { + // Reuse the existing complex-float window specialization only. + channelize_poly_impl(output, input, filter, 2, 2, + cuda_exec.getStream(), offset); + } + MATX_CUDA_CHECK_LAST_ERROR(); + cuda_exec.sync(); + (reference = channelize_poly( + input_ref, filter_ref, num_channels, num_channels)).run(host_exec); + + for (index_t b = 0; b < batches; b++) { + for (index_t t = 0; t < output_len; t++) { + for (index_t c = 0; c < stride * num_channels; c++) { + const auto got = output_storage(b, t, c); + if (t < offset || t >= offset + window_len || c % stride != 0) { + EXPECT_EQ(static_cast(got.real()), 9.0); + EXPECT_EQ(static_cast(got.imag()), -9.0); + } else { + const auto expected = reference(b, t, c / stride); + const double scale = 1.0 + cuda::std::abs(expected); + EXPECT_NEAR(static_cast(got.real()), expected.real(), tolerance * scale); + EXPECT_NEAR(static_cast(got.imag()), expected.imag(), tolerance * scale); + } + } + } + } + } + } + } + }; + run.template operator()(); + // M=2 windows: a partial trailing row and batched windows of adjacent sizes. + run.template operator()(1, 1031, 3, 513); + run.template operator()(64, 8195, 3, 3584); + run.template operator()(64, 8195, 3, 3585); + run.template operator()(); + run.template operator()(); + run.template operator()(); + run.template operator()(); + run.template operator()(); + run.template operator()(); + run.template operator()(); + run.template operator()(); + run.template operator()(); + + MATX_EXIT_HANDLER(); +} + +// cuFFT has no half-precision transforms of non-power-of-two size. Critical +// channelizers with 3, 5, or 6 channels always use the fused kernel (see +// CriticalSmallChannelsStrided); other such configurations are rejected before any +// kernel runs. +TEST(ChannelizePoly, HalfNonPowerOfTwoWithoutFusedKernelThrows) +{ + MATX_ENTER_HANDLER(); + using Complex = matxFp16Complex; + cudaExecutor exec{}; + struct Config { index_t m, d; }; + const Config cases[] = {{3, 2}, {5, 4}, {6, 3}, {7, 7}, {12, 12}}; + for (const auto &tc : cases) { + SCOPED_TRACE(::testing::Message() << "M=" << tc.m << " D=" << tc.d); + constexpr index_t batches = 2; + constexpr index_t input_len = 1027; + const index_t filter_len = 13 * tc.m - 1; + const index_t rows = (input_len + tc.d - 1) / tc.d; + // Strided views of the same types as CriticalSmallChannelsStrided. + auto input_storage = make_tensor({batches, input_len}); + auto filter_storage = make_tensor({filter_len}); + auto output_storage = make_tensor({batches, rows, tc.m}); + auto input = slice(input_storage, {0, 0}, {batches, input_len}, {1, 1}); + auto filter = slice(filter_storage, {0}, {filter_len}, {1}); + auto output = slice(output_storage, {0, 0, 0}, {batches, rows, tc.m}, {1, 1, 1}); + try { + (output = channelize_poly(input, filter, tc.m, tc.d)).run(exec); + ADD_FAILURE() << "expected matxInvalidParameter"; + } catch (const detail::matxException &e) { + EXPECT_EQ(e.e, matxInvalidParameter); + } + } + MATX_EXIT_HANDLER(); +} + // Tests that involve non-trivial input and output operators. TYPED_TEST(ChannelizePolyTestNonHalfFloatTypes, Operators) { @@ -1260,6 +2283,11 @@ TYPED_TEST(ChannelizePolyTestNonHalfFloatTypes, OversampledInteger) { 2500, 187, 10, 5 }, { 1800, 120, 8, 4 }, { 37193, 41*8+4, 8, 4 }, + // Fused radix/power-of-two paths with partial filter rows. + { 10003, 17*20+7, 20, 10 }, + { 10003, 17*40+7, 40, 20 }, + { 10003, 17*64+7, 64, 32 }, + { 10003, 17*80+7, 80, 40 }, // 5x oversampling (D = M/5) { 2500, 170, 10, 2 }, // Large channel count @@ -1318,6 +2346,11 @@ TYPED_TEST(ChannelizePolyTestNonHalfFloatTypes, OversampledRational) { 2500, 187, 11, 4 }, // Large channel count, rational { 35000, 5*181+17, 181, 60 }, + // Fused radix/power-of-two rational-oversampling paths. + { 10003, 17*20+7, 20, 7 }, + { 10003, 17*40+7, 40, 32 }, + { 10003, 17*64+7, 64, 27 }, + { 10003, 17*80+7, 80, 33 }, }; for (size_t i = 0; i < sizeof(test_cases)/sizeof(test_cases[0]); i++) { @@ -1420,6 +2453,11 @@ TYPED_TEST(ChannelizePolyTestNonHalfFloatTypes, SmemTiledMaximallyDecimated) { 256000, 10*512+37, 512 }, // M=1024, P=4 — large channel count, small P { 102400, 4*1024, 1024 }, + // M=512, P=64 — exercises the grouped long-filter path and CTILE=32 + { 128000, 64*512, 512 }, + // M=512, P=192 — exercises the CTILE=32 path when it fits and the + // established generic fallback for wider element types + { 131072, 192*512, 512 }, }; for (size_t i = 0; i < sizeof(test_cases)/sizeof(test_cases[0]); i++) { @@ -1542,9 +2580,9 @@ TYPED_TEST(ChannelizePolyTestNonHalfFloatTypes, SmemTiledOversampledLargeK) MATX_EXIT_HANDLER(); } -// SmemTiled oversampled with FilterInSmem=true. -// Requires CTILE * K * P * sizeof(filter_t) <= 2048. -// With K=2 (2x integer oversample) and small P, the filter fits in smem. +// SmemTiled oversampled with FilterInSmem=true. Oversampled tiles cache only the +// Full filter layout, and only when its M * P taps fit in +// SmemTiledOversampledFilterBytes (4 KiB), so these cases keep M * P <= 256. TYPED_TEST(ChannelizePolyTestNonHalfFloatTypes, SmemTiledOversampledFilterInSmem) { MATX_ENTER_HANDLER(); @@ -1558,12 +2596,14 @@ TYPED_TEST(ChannelizePolyTestNonHalfFloatTypes, SmemTiledOversampledFilterInSmem index_t num_channels; index_t decimation_factor; } test_cases[] = { - // M=1024, D=512, K=2, P=4: filter_smem = 64*2*4*sizeof(filter_t) = 2048 bytes (float) - { 102400, 4*1024, 1024, 512 }, - // M=512, D=256, K=2, P=3: filter_smem = 64*2*3*sizeof(filter_t) = 1536 bytes (float) - { 51200, 3*512, 512, 256 }, - // M=1024, D=512, K=2, P=2: filter_smem = 64*2*2*sizeof(filter_t) = 1024 bytes (float) - { 102400, 2*1024, 1024, 512 }, + // M=24, D=12, K=2, P=4 with a partial final polyphase row + { 10007, 4*24-5, 24, 12 }, + // M=36, D=18, K=2, P=3 + { 10007, 3*36, 36, 18 }, + // M=48, D=36, K=4, P=2 (rational oversampling) + { 10007, 2*48, 48, 36 }, + // M=48, D=24, K=2, P=4: a 3 KiB complex filter, above the former 2 KiB limit + { 10007, 4*48, 48, 24 }, }; for (size_t i = 0; i < sizeof(test_cases)/sizeof(test_cases[0]); i++) { @@ -1582,6 +2622,12 @@ TYPED_TEST(ChannelizePolyTestNonHalfFloatTypes, SmemTiledOversampledFilterInSmem this->pb->NumpyToTensorView(a, "a"); this->pb->NumpyToTensorView(f, "filter_random"); + // Fail if dispatch changes stop these cases from caching the filter. + using Accum = typename inner_op_type_t::type; + const auto plan = detail::cpoly::SelectPlan( + b, a, f, num_channels, decimation_factor); + ASSERT_EQ(plan.backend, detail::cpoly::Backend::Tiled); + ASSERT_EQ(plan.tiled.filter_layout, detail::cpoly::SmemTiledFilterLayout::Full); (b = channelize_poly(a, f, num_channels, decimation_factor)).run(this->exec); this->exec.sync(); @@ -1606,7 +2652,9 @@ TYPED_TEST(ChannelizePolyTestNonHalfFloatTypes, ComplexFilter) index_t num_channels; index_t decimation_factor; } test_cases[] = { - // Complex filter, D==M, small M (FusedChan path) + // Complex filter, D==M, specialized two-point DFT + { 1800, 120, 2, 2 }, + // Complex filter, D==M, small M (fused FIR and DFT) { 1800, 120, 4, 4 }, // Complex filter, D==M, medium M (_Smem path) { 2500, 170, 10, 10 }, @@ -1703,9 +2751,10 @@ TYPED_TEST(ChannelizePolyTestNonHalfFloatTypes, GenericOversampledFallback) namespace cpoly_window_test { // Run one windowed-vs-full comparison. InT is the input sample type (real -> -// R2C path, complex -> C2C path); the filter is real. -template -bool run_case(cudaExecutor &exec, index_t N, index_t M, index_t D, index_t L, +// R2C path, complex -> C2C path); the filter is real. ExplicitFused launches +// the fused window regardless of the device's performance heuristics. +template +void run_case(cudaExecutor &exec, index_t N, index_t M, index_t D, index_t L, index_t offset, index_t count) { using OutT = cuda::std::complex; @@ -1713,7 +2762,7 @@ bool run_case(cudaExecutor &exec, index_t N, index_t M, index_t D, index_t L, const index_t T = (N + D - 1) / D; // full number of output elements per channel if (offset < 0 || count <= 0 || offset + count > T) { - return true; // skip degenerate windows + return; // skip degenerate windows } auto h = make_tensor({L}); @@ -1736,12 +2785,22 @@ bool run_case(cudaExecutor &exec, index_t N, index_t M, index_t D, index_t L, // Full one-shot reference (uses whichever kernel the heuristics pick). auto out_full = make_tensor({T, M}); (out_full = channelize_poly(in, h, M, D)).run(exec); + ASSERT_EQ(cudaGetLastError(), cudaSuccess); - // Windowed compute of only rows [offset, offset+count) via normal dispatch. + // Windowed compute of only rows [offset, offset+count). auto out_win = make_tensor({count, M}); - channelize_poly_impl( - out_win, in, h, M, D, exec.getStream(), offset); - exec.sync(); + if constexpr (ExplicitFused) { + using OutOp = decltype(out_win); + ASSERT_TRUE((detail::cpoly::FusedRadixFeasible( + out_win, in, h, M, D))); + detail::cpoly::RunFusedRadix( + out_win, in, h, M, D, exec.getStream(), offset); + } else { + channelize_poly_impl( + out_win, in, h, M, D, exec.getStream(), offset); + } + ASSERT_EQ(cudaGetLastError(), cudaSuccess); + ASSERT_EQ(cudaStreamSynchronize(exec.getStream()), cudaSuccess); float max_abs = 0.0f, max_err = 0.0f; for (index_t t = 0; t < count; ++t) { @@ -1756,7 +2815,6 @@ bool run_case(cudaExecutor &exec, index_t N, index_t M, index_t D, index_t L, EXPECT_TRUE(ok) << "M=" << M << " D=" << D << " L=" << L << " offset=" << offset << " count=" << count << " max_err=" << max_err << " max_abs=" << max_abs; - return ok; } template @@ -1764,7 +2822,8 @@ void sweep(cudaExecutor &exec) { struct Cfg { index_t M, D; }; const index_t N = 4096; - for (Cfg c : {Cfg{4, 4}, Cfg{8, 8}, // maximally decimated + for (Cfg c : {Cfg{4, 4}, Cfg{5, 5}, Cfg{6, 6}, Cfg{8, 8}, + // maximally decimated Cfg{8, 4}, Cfg{8, 2}, // integer oversampled Cfg{6, 4}, Cfg{8, 6}, Cfg{9, 6}}) { // rational oversampled const index_t L = 4 * c.M; // P = 4 taps/branch @@ -1780,22 +2839,56 @@ void sweep(cudaExecutor &exec) // A bounded middle window as well. run_case(exec, N, c.M, c.D, L, T / 3, T / 4); } + + // Small oversampled channelizers must preserve phase at arbitrary window + // offsets, including partial filter rows and partially populated end rows. + for (index_t channels : {3, 4, 5, 6}) { + constexpr index_t input_len = 1031; + const index_t filter_len = 13 * channels - 1; + for (index_t decimation = 1; decimation < channels; decimation++) { + const index_t rows = (input_len + decimation - 1) / decimation; + run_case(exec, input_len, channels, decimation, filter_len, 3, 67); + run_case(exec, input_len, channels, decimation, filter_len, rows - 17, 17); + } + // A longer bounded window also exercises repeated work within a CTA. + run_case(exec, 131075, channels, 1, filter_len, 3, 65539); + } + + // Critical direct FIRs must preserve causal and trailing bounds in windows. + for (index_t channels : {5, 6}) { + constexpr index_t input_len = 1031; + const index_t rows = (input_len + channels - 1) / channels; + for (index_t tail : {0, 1}) { + const index_t filter_len = 13 * channels - tail; + for (index_t offset : {index_t{1}, index_t{7}, rows - 1}) { + run_case(exec, input_len, channels, channels, filter_len, + offset, std::min(index_t{67}, rows - offset)); + } + } + } } -// Target the non-fused CUDA paths with out_elem_offset > 0. P is the number -// of prototype-filter taps per channel. These sizes are chosen from the -// dispatcher thresholds in transforms/channelize_poly.h: -// M=16, D=16, P=8: Smem -// M=64, D=32, P=8: SmemTiled, Full filter layout -// M=256, D=128, P=8: SmemTiled, Rotated filter layout -// M=64, D=32, P=20: SmemTiled, Global filter layout -// M=64, D=64/48, P=192: Generic (input tile exceeds 48 KiB) +// Target CUDA dispatch paths with out_elem_offset > 0. P is the number of +// prototype-filter taps per channel. With the float filter used here, the +// dispatcher in transforms/channelize_poly.h selects (real / complex input): +// M=2, D=1, P=64: fused two-channel oversampled leaf +// M=40, D=20/32, P=17: fused +// M=12, D=12, P=8: Smem +// M=64, D=32, P=8: fused / SmemTiled 32x4, Full filter layout +// M=256, D=128, P=8: SmemTiled 32x4, Global filter layout +// M=64, D=32, P=20: fused / SmemTiled 32x4, Global filter layout +// M=64, D=64, P=192: SmemTiled 32x16, Global filter layout / Generic +// M=64, D=48, P=192: SmemTiled 32x4, Global filter layout / Generic +// ExplicitPlansMatchHost covers every backend and filter layout at a nonzero +// offset regardless of these heuristics. template void large_dispatch_sweep(cudaExecutor &exec) { struct Cfg { index_t M, D, P; }; const index_t N = 2051; // partial trailing block for every D below - for (Cfg c : {Cfg{16, 16, 8}, Cfg{64, 32, 8}, Cfg{256, 128, 8}, + for (Cfg c : {Cfg{2, 1, 64}, Cfg{40, 20, 17}, Cfg{40, 32, 17}, + Cfg{12, 12, 8}, + Cfg{64, 32, 8}, Cfg{256, 128, 8}, Cfg{64, 32, 20}, Cfg{64, 64, 192}, Cfg{64, 48, 192}}) { const index_t L = c.P * c.M - 1; // partial final polyphase row const index_t T = (N + c.D - 1) / c.D; @@ -1804,6 +2897,34 @@ void large_dispatch_sweep(cudaExecutor &exec) std::min(index_t(19), T - offset)); run_case(exec, N, c.M, c.D, L, T - 1, 1); } + + // A long Smem window has enough rows for the launch to use 16-row groups. + run_case(exec, 12 * 8192 + 5, 12, 12, 8 * 12 - 1, 3, 4099); + + // Explicit fused M=40 windows cover critical sampling, oversampling, and a + // partial trailing row. + for (index_t D : {20, 32, 40}) { + const index_t T = (N + D - 1) / D; + run_case(exec, N, 40, D, 17 * 40 - 1, 3, 19); + run_case(exec, N, 40, D, 17 * 40 - 1, T - 1, 1); + } + + // Long, offset windows cover repeated ring reloads and a partial final tile. + // P=1 has no halo; P=13 also wraps within the filtered sample history. + for (index_t P : {1, 13}) { + run_case(exec, 600003, 32, 32, P * 32 - 1, 7, 17003); + run_case(exec, 600003, 40, 40, P * 40 - 1, 7, 14003); + } + + // Nonconstant input, nonzero offsets, and enough rows for grouped + // CTILE=32/NOUT=16 reloads. Both windows end with a partial output group; + // the second also includes the input's partially populated final row. + constexpr index_t grouped_N = 1056 * 512 + 3; + constexpr index_t grouped_count = 1051; + constexpr index_t grouped_T = (grouped_N + 511) / 512; + for (index_t offset : {index_t{3}, grouped_T - grouped_count}) { + run_case(exec, grouped_N, 512, 512, 64 * 512 - 1, offset, grouped_count); + } } } // namespace cpoly_window_test @@ -1821,3 +2942,119 @@ TEST(ChannelizePoly, OutputElemWindowComplexInput) cpoly_window_test::sweep>(exec); cpoly_window_test::large_dispatch_sweep>(exec); } + +TEST(ChannelizePoly, OutputElemWindowLongSignalsAndFilters) +{ + MATX_ENTER_HANDLER(); + using Complex = cuda::std::complex; + cudaExecutor cuda_exec{}; + SingleThreadedHostExecutor host_exec{}; + struct Config { index_t channels, taps, input_len, count; }; + const Config cases[] = { + // Practical filters with long windows and incomplete final groups. + {512, 13, 553987, 1079}, + {256, 21, 553987, 2103}, + // Input tiles fit in 48 KiB, but caching the filter would exceed it. + {32, 80, 131075, 4097}, + {64, 80, 262147, 4097}, + // A trailing window close to the signed 32-bit sample-index limit. + {512, 64, std::numeric_limits::max() - 512, 1079}, + // Fused staging includes padded rows beyond the valid sample range. + {2, 13, std::numeric_limits::max() - 10, 67}, + {3, 13, std::numeric_limits::max() - 10, 67}, + {4, 13, std::numeric_limits::max() - 10, 67}, + {5, 13, std::numeric_limits::max() - 10, 67}, + {6, 13, std::numeric_limits::max() - 10, 67}, + }; + for (const auto &tc : cases) { + SCOPED_TRACE(::testing::Message() + << "M=" << tc.channels << " P=" << tc.taps << " N=" << tc.input_len); + const index_t total_rows = (tc.input_len + tc.channels - 1) / tc.channels; + const index_t offset = total_rows - tc.count; + // Coprime periods expose sample/row shifts without a huge allocation. + auto indices = range<0>({tc.input_len}, index_t{0}, index_t{1}); + auto input = as_complex_double( + (as_double(indices % 1009) - 504.0) / 512.0, (as_double(indices % 1013) - 506.0) / 512.0); + auto filter = make_tensor({tc.taps * tc.channels - 1}); + for (index_t k = 0; k < filter.Size(0); k++) { + filter(k) = static_cast((k % 13) - 6) / 64.0f; + } + auto output = make_tensor({tc.count, tc.channels}); + channelize_poly_impl(output, input, filter, tc.channels, tc.channels, + cuda_exec.getStream(), offset); + MATX_CUDA_CHECK_LAST_ERROR(); + cuda_exec.sync(); + + // Check beginning, interior, and trailing rows against the host FIR/DFT. + auto reference = make_tensor({1, tc.channels}); + for (index_t row : {index_t{0}, tc.count / 2, tc.count - 17, tc.count - 1}) { + channelize_poly_impl(reference, input, filter, + tc.channels, tc.channels, host_exec, offset + row); + for (index_t c = 0; c < tc.channels; c++) { + const auto expected = reference(0, c); + const double tolerance = 1e-10 * (1.0 + cuda::std::abs(expected)); + EXPECT_NEAR(output(row, c).real(), expected.real(), tolerance); + EXPECT_NEAR(output(row, c).imag(), expected.imag(), tolerance); + } + } + } + MATX_EXIT_HANDLER(); +} + +// The tiled kernel uses 32-bit indices only when every input index it forms fits. +// An oversampled ring loads through the end of the input row holding the +// window's newest sample, which can pass INT32_MAX even when input_len + M does +// not. +TEST(ChannelizePoly, TiledIndexWidthNearInt32Limit) +{ + MATX_ENTER_HANDLER(); + namespace cp = detail::cpoly; + using Complex = cuda::std::complex; + constexpr index_t max_index = std::numeric_limits::max(); + constexpr index_t channels = 1022; + constexpr index_t decimation = 1016; + // The newest input row ends 1006 samples past INT32_MAX. + constexpr index_t input_len = 2147482625; + constexpr index_t total_rows = (input_len + decimation - 1) / decimation; + static_assert(input_len + channels <= max_index); + EXPECT_FALSE(cp::SmemTiledFitsInt32(input_len, channels, decimation, total_rows)); + // Critical sampling at the limit keeps 32-bit indices. + constexpr index_t critical_len = max_index - 512; + EXPECT_TRUE(cp::SmemTiledFitsInt32(critical_len, 512, 512, (critical_len + 511) / 512)); + + // A trailing window through the tiled kernel matches the host FIR/DFT. + cudaExecutor cuda_exec{}; + SingleThreadedHostExecutor host_exec{}; + auto indices = range<0>({input_len}, index_t{0}, index_t{1}); + auto input = as_complex_double( + (as_double(indices % 1009) - 504.0) / 512.0, (as_double(indices % 1013) - 506.0) / 512.0); + auto filter = make_tensor({4 * channels - 1}); + for (index_t k = 0; k < filter.Size(0); k++) { + filter(k) = static_cast((k % 13) - 6) / 64.0f; + } + constexpr index_t count = 67; + constexpr index_t offset = total_rows - count; + auto output = make_tensor({count, channels}); + const auto plan = cp::SelectPlan( + output, input, filter, channels, decimation); + ASSERT_EQ(plan.backend, cp::Backend::Tiled); + channelize_poly_impl( + output, input, filter, channels, decimation, cuda_exec.getStream(), offset); + MATX_CUDA_CHECK_LAST_ERROR(); + cuda_exec.sync(); + + auto reference = make_tensor({1, channels}); + for (index_t row : {index_t{0}, count / 2, count - 1}) { + channelize_poly_impl( + reference, input, filter, channels, decimation, host_exec, offset + row); + for (index_t c = 0; c < channels; c++) { + const auto expected = reference(0, c); + const double tolerance = 1e-10 * (1.0 + cuda::std::abs(expected)); + EXPECT_NEAR(output(row, c).real(), expected.real(), tolerance) << "row=" << row; + EXPECT_NEAR(output(row, c).imag(), expected.imag(), tolerance) << "row=" << row; + } + } + MATX_EXIT_HANDLER(); +} diff --git a/test/00_transform/StreamingChannelize.cu b/test/00_transform/StreamingChannelize.cu index 8ebf6bc06..b1b207fbe 100644 --- a/test/00_transform/StreamingChannelize.cu +++ b/test/00_transform/StreamingChannelize.cu @@ -161,7 +161,7 @@ void large_dispatch_sweep(cudaExecutor &exec) { struct Cfg { index_t M, D, P; }; const index_t N = 2051; - for (Cfg c : {Cfg{16, 16, 8}, Cfg{64, 32, 8}, Cfg{256, 128, 8}, + for (Cfg c : {Cfg{12, 12, 8}, Cfg{64, 32, 8}, Cfg{256, 128, 8}, Cfg{64, 32, 20}, Cfg{64, 64, 192}, Cfg{64, 48, 192}}) { const index_t L = c.P * c.M - 1; // partial final polyphase row for (index_t chunk : {index_t(7), index_t(113), index_t(509)}) { diff --git a/test/test_vectors/generators/00_transforms.py b/test/test_vectors/generators/00_transforms.py index ff0a887e2..f093c0a23 100755 --- a/test/test_vectors/generators/00_transforms.py +++ b/test/test_vectors/generators/00_transforms.py @@ -191,12 +191,10 @@ def __init__(self, dtype: str, size: List[int]): channelize_poly_operators.np_random_state = np.random.get_state() - def channelize(self) -> Dict[str, np.ndarray]: + @staticmethod + def _channelize(x, h, num_channels): def idivup(a, b) -> int: return (a+b-1)//b - h = self.res['filter_random'] - num_channels = self.res['num_channels'] - x = self.res['a'] num_taps_per_channel = idivup(h.size, num_channels) if num_channels * num_taps_per_channel > h.size: h = np.pad(h, (0,num_channels*num_taps_per_channel-h.size)) @@ -214,7 +212,7 @@ def idivup(a, b) -> int: return (a+b-1)//b # flipud because samples are inserted into the filter banks in order # M-1, M-2, ..., 0 xf = np.flipud(np.reshape(xpad, (num_channels,x_len_per_channel), order='F')) - buf = np.zeros((num_channels, num_taps_per_channel), dtype=self.dtype) + buf = np.zeros((num_channels, num_taps_per_channel), dtype=x.dtype) # We scale the outputs by num_channels because we use the ifft # and it scales by 1/N for an N-point FFT. We use ifft instead @@ -225,9 +223,8 @@ def idivup(a, b) -> int: return (a+b-1)//b for i in range(x_len_per_channel): buf[:, 1:] = buf[:, 0:num_taps_per_channel-1] buf[:, 0] = xf[:, i] - for j in range(num_channels): - out[batch_ind, j, i] = scale * np.dot(np.squeeze(buf[j,:]), np.squeeze(h[j,:])) - out_hreal[batch_ind, j, i] = scale * np.dot(np.squeeze(buf[j,:]), np.squeeze(np.real(h[j,:]))) + out[batch_ind, :, i] = scale * np.sum(buf * h, axis=1) + out_hreal[batch_ind, :, i] = scale * np.sum(buf * np.real(h), axis=1) out[batch_ind,:,:] = ifft(out[batch_ind,:,:], axis=0) out_hreal[batch_ind,:,:] = ifft(out_hreal[batch_ind,:,:], axis=0) if num_batches > 1: @@ -242,8 +239,20 @@ def idivup(a, b) -> int: return (a+b-1)//b else: out = np.transpose(np.reshape(out, out.shape[1:]), axes=[1,0]) out_hreal = np.transpose(np.reshape(out_hreal, out_hreal.shape[1:]), axes=[1,0]) - self.res['b_random'] = out - self.res['b_random_hreal'] = out_hreal + return out, out_hreal + + def channelize(self) -> Dict[str, np.ndarray]: + self.res['b_random'], self.res['b_random_hreal'] = self._channelize( + self.res['a'], self.res['filter_random'], self.res['num_channels']) + return self.res + + def channelize_accum(self) -> Dict[str, np.ndarray]: + self.channelize() + # Isolate accumulation error using the exact FP32 operands in FP64. + x = self.res['a'].astype(np.complex64).astype(np.complex128) + h = self.res['filter_random_real'].astype(np.float32).astype(np.float64) + _, self.res['b_quantized_hreal'] = self._channelize( + x, h, self.res['num_channels']) return self.res def channelize_oversampled(self) -> Dict[str, np.ndarray]: @@ -631,4 +640,4 @@ def normalize_center(self) -> Dict[str, np.ndarray]: return { 'in_m': seq, 'out_m': (seq - np.mean(seq, axis=0, keepdims=True)) - } \ No newline at end of file + }