From e0e52953570caea66096cda08dc814f34735f7af Mon Sep 17 00:00:00 2001 From: steven-na Date: Wed, 5 Aug 2026 16:03:11 -0700 Subject: Fixed byte/element inconsistencies --- src/common.h | 2 ++ src/dft.c | 2 +- src/dft.h | 10 ++++++++-- tests/dft.c | 18 ++++++++++++++++++ 4 files changed, 29 insertions(+), 3 deletions(-) diff --git a/src/common.h b/src/common.h index abca66d..1154bdf 100644 --- a/src/common.h +++ b/src/common.h @@ -19,6 +19,8 @@ #define F64_EQ(x, y, eps) (fabs((x) - (y)) <= (eps)) #define MAX(n, m) ((n > m) ? (n) : (m)) #define MIN(n, m) ((n < m) ? (n) : (m)) +// n and m are expected to be byte counts (m a power of 2); keep call sites in bytes, +// not element counts, so alignment semantics stay consistent across the codebase. #define ALIGN_UP_POW2(n, m) (((u64)(n) + (u64)(m) - 1) & (~((u64)(m) - 1))) typedef int8_t i8; diff --git a/src/dft.c b/src/dft.c index 7b95518..7f90318 100644 --- a/src/dft.c +++ b/src/dft.c @@ -149,7 +149,7 @@ stft_data_t short_time_fourier_transform(smrt_arena_t *arena, u64 window_size, u assert(F64_EQ(round(log2(window_size)), log2(window_size), 1e-9) && "STFT input window_size must be a power of 2"); - sample_count = ALIGN_UP_POW2(sample_count, window_size); + sample_count = ALIGN_UP_POW2(sample_count * sizeof(f64), STFT_SAMPLE_ALIGN_BYTES(window_size)) / sizeof(f64); u64 segment_count = ((sample_count - window_size) / hop_size) + 1; diff --git a/src/dft.h b/src/dft.h index aaa2796..6b63175 100644 --- a/src/dft.h +++ b/src/dft.h @@ -42,9 +42,15 @@ typedef struct { u64 sample_rate; } stft_data_t; +/// Byte alignment required for a samples buffer to safely back +/// short_time_fourier_transform with the given window_size. Pass this as +/// align_up_memoryn to wav_load/load_wav_file when loading samples for STFT use. +#define STFT_SAMPLE_ALIGN_BYTES(window_size) ((u64)(window_size) * sizeof(f64)) + /// Run STFT algorithm on samples, sliding a window_size window by hop_size each step -/// Warning: This function assumes that samples' allocation is large enough to fit -/// ALIGN_UP_POW2(sample_count, window_size) elements. +/// Warning: This function assumes that samples' allocation is at least +/// STFT_SAMPLE_ALIGN_BYTES(window_size)-byte aligned (e.g. by passing +/// STFT_SAMPLE_ALIGN_BYTES(window_size) as align_up_memoryn to wav_load/load_wav_file). stft_data_t short_time_fourier_transform(smrt_arena_t * arena , u64 window_size , u64 hop_size , diff --git a/tests/dft.c b/tests/dft.c index 5a581bd..9c406f5 100644 --- a/tests/dft.c +++ b/tests/dft.c @@ -218,6 +218,24 @@ Test(dft, stft_segments_match_direct_transform) { smrt_arena_destroy(arena); } +Test(dft, stft_handles_unaligned_sample_count_with_wav_load_style_buffer) { + smrt_arena_t *arena = smrt_arena_create(KiB(64), KiB(4), false); + + u64 sample_count = 13, sample_rate = 16, window_size = 8, hop_size = 4; + + // Mirrors how wav_load/load_wav_file size their sample buffer: byte-aligned to + // STFT_SAMPLE_ALIGN_BYTES(window_size) and zero-initialized, not element-aligned. + f64 *samples = smrt_arena_push(arena, ALIGN_UP_POW2(sample_count * sizeof(f64), STFT_SAMPLE_ALIGN_BYTES(window_size)), true); + for (u64 i = 0; i < sample_count; i++) samples[i] = (f64)(i + 1); + + stft_data_t stft = short_time_fourier_transform(arena, window_size, hop_size, samples, sample_count, sample_rate); + + cr_assert_eq(stft.segment_count, 3); + cr_assert_neq(stft.segments, NULL); + + smrt_arena_destroy(arena); +} + Test(dft, fast_fourier_transform_matches_direct_transform) { smrt_arena_t *arena = smrt_arena_create(KiB(16), KiB(4), false); -- cgit v1.2.3