Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
125 changes: 125 additions & 0 deletions src/common/audio/qwen3_asr_preprocessor.cpp
Original file line number Diff line number Diff line change
@@ -0,0 +1,125 @@
/// \file qwen3_asr_preprocessor.cpp
/// \brief Qwen3-ASR waveform-to-log-mel preprocessing.

#include "audio/qwen3_asr_preprocessor.hpp"
#include "audio_process_utils/audioproc.hpp"

#include <algorithm>
#include <cmath>
#include <stdexcept>

Qwen3ASRPreprocessor::Qwen3ASRPreprocessor() : Qwen3ASRPreprocessor(Config{}) {}

Qwen3ASRPreprocessor::Qwen3ASRPreprocessor(Config config) : config_(config) {
if (config_.sampling_rate <= 0 || config_.n_fft <= 0 || config_.hop_length <= 0 ||
config_.num_mel_bins <= 0 || config_.n_window < 0 || config_.min_length < 0) {
throw std::invalid_argument("invalid Qwen3-ASR preprocessor configuration");
}
}

int Qwen3ASRPreprocessor::feature_output_length(int mel_frames) {
if (mel_frames <= 0) return 0;

// Exact integer equivalent of Transformers' three stride-2 CNN length
// transforms. A complete 100-frame chunk produces 13 output tokens.
const int remainder = mel_frames % 100;
const int remainder_output = (remainder + 7) / 8;
return remainder_output + (mel_frames / 100) * 13;
}

qwen3_asr_features_t Qwen3ASRPreprocessor::extract(const float* samples, std::size_t num_samples) const {
if (num_samples > 0 && samples == nullptr) {
throw std::invalid_argument("samples must not be null when num_samples is non-zero");
}

const std::size_t padded_samples = std::max<std::size_t>(num_samples, config_.min_length);
std::vector<float> waveform(padded_samples, 0.0f);
if (num_samples > 0) {
std::copy(samples, samples + num_samples, waveform.begin());
}

const int num_frequency_bins = config_.n_fft / 2 + 1;
const int allocated_frames = audioproc::stft_num_frames(
static_cast<int>(waveform.size()), config_.n_fft, config_.hop_length, /*center=*/true);
if (allocated_frames <= 1) {
return {};
}

const std::vector<float> window = audioproc::window_function_optimized(
config_.n_fft, "hann", /*periodic=*/true);
const std::vector<float> mel_filters = audioproc::mel_filter_bank_optimized(
num_frequency_bins,
config_.num_mel_bins,
0.0f,
static_cast<float>(config_.sampling_rate) / 2.0f,
config_.sampling_rate,
/*apply_slaney_norm=*/true,
/*slaney_mel_scale=*/true);

std::vector<float> power_spec(
static_cast<std::size_t>(allocated_frames) * num_frequency_bins);
const int stft_frames = audioproc::stft_power_optimized(
waveform.data(),
static_cast<int>(waveform.size()),
window.data(),
config_.n_fft,
config_.hop_length,
/*center=*/true,
audioproc::StftPadMode::reflect,
power_spec.data());

// Qwen3ASRFeatureExtractor uses stft[..., :-1].
const int valid_frames = stft_frames - 1;
if (valid_frames <= 0) {
return {};
}

// audioproc emits [frames, mel_bins].
std::vector<float> frame_major(
static_cast<std::size_t>(valid_frames) * config_.num_mel_bins);
audioproc::mel_spectrogram_optimized(
power_spec.data(),
mel_filters.data(),
frame_major.data(),
valid_frames,
num_frequency_bins,
config_.num_mel_bins);

const int feature_count = valid_frames * config_.num_mel_bins;
std::vector<float> log_mel(feature_count);
// Use the scalar log10 path here. The AVX512 helper intentionally uses a
// low-order logarithm approximation which is fast but introduces errors up
// to ~2e-2 after Qwen's normalization, large enough to affect ASR parity.
audioproc::log_mel_floor</*UseClamp=*/true, /*Base=*/10>(
frame_major.data(), log_mel.data(), feature_count, 1e-10f);

const float max_value = audioproc::reduce_max(log_mel.data(), feature_count);
audioproc::clamp_below_max(log_mel.data(), feature_count, max_value, 8.0f);
audioproc::affine_scale(log_mel.data(), feature_count, 4.0f, 4.0f);

int padded_frames = valid_frames;
const int frame_multiple = config_.n_window * 2;
if (frame_multiple > 1) {
const int remainder = padded_frames % frame_multiple;
if (remainder != 0) padded_frames += frame_multiple - remainder;
}

qwen3_asr_features_t result;
result.num_mel_bins = config_.num_mel_bins;
result.num_frames = padded_frames;
result.valid_frames = valid_frames;
result.input_features.assign(
static_cast<std::size_t>(config_.num_mel_bins) * padded_frames, 0.0f);
result.attention_mask.assign(padded_frames, 0);
std::fill(result.attention_mask.begin(), result.attention_mask.begin() + valid_frames, 1);

// Transformers layout is [mel_bins, frames].
for (int frame = 0; frame < valid_frames; ++frame) {
for (int mel = 0; mel < config_.num_mel_bins; ++mel) {
result.input_features[static_cast<std::size_t>(mel) * padded_frames + frame] =
log_mel[static_cast<std::size_t>(frame) * config_.num_mel_bins + mel];
}
}

return result;
}
137 changes: 137 additions & 0 deletions src/common/audio/qwen3_asr_utils.cpp
Original file line number Diff line number Diff line change
@@ -0,0 +1,137 @@
/// \file qwen3_asr_utils.cpp
/// \brief Prompt and output helpers for Qwen3-ASR.

#include "audio/qwen3_asr_utils.hpp"

#include <algorithm>
#include <array>
#include <cctype>
#include <stdexcept>

namespace {

constexpr std::array<std::string_view, 30> supported_languages = {
"Chinese", "English", "Cantonese", "Arabic", "German", "French",
"Spanish", "Portuguese", "Indonesian", "Italian", "Korean", "Russian",
"Thai", "Vietnamese", "Japanese", "Turkish", "Hindi", "Malay", "Dutch",
"Swedish", "Danish", "Finnish", "Polish", "Czech", "Filipino", "Persian",
"Greek", "Romanian", "Hungarian", "Macedonian",
};

std::string trim(std::string_view value) {
const auto first = std::find_if_not(value.begin(), value.end(), [](unsigned char c) {
return std::isspace(c) != 0;
});
const auto last = std::find_if_not(value.rbegin(), value.rend(), [](unsigned char c) {
return std::isspace(c) != 0;
}).base();
if (first >= last) return {};
return std::string(first, last);
}

std::string ascii_lower(std::string_view value) {
std::string result(value);
std::transform(result.begin(), result.end(), result.begin(), [](unsigned char c) {
return static_cast<char>(std::tolower(c));
});
return result;
}

} // namespace

std::string qwen3_asr_normalize_language(std::string_view language) {
std::string result = trim(language);
if (result.empty()) return result;

std::transform(result.begin(), result.end(), result.begin(), [](unsigned char c) {
return static_cast<char>(std::tolower(c));
});
result.front() = static_cast<char>(std::toupper(static_cast<unsigned char>(result.front())));
return result;
}

bool qwen3_asr_is_supported_language(std::string_view language) {
const std::string normalized = qwen3_asr_normalize_language(language);
return std::find(supported_languages.begin(), supported_languages.end(), normalized) != supported_languages.end();
}

std::string qwen3_asr_build_prompt(
std::string_view context,
int audio_token_count,
std::optional<std::string_view> forced_language) {
if (audio_token_count <= 0) {
throw std::invalid_argument("audio_token_count must be positive");
}

std::string language;
if (forced_language.has_value()) {
language = qwen3_asr_normalize_language(*forced_language);
if (!qwen3_asr_is_supported_language(language)) {
throw std::invalid_argument("unsupported Qwen3-ASR language: " + language);
}
}

constexpr std::string_view audio_token = "<|audio_pad|>";
std::string prompt;
prompt.reserve(context.size() + static_cast<std::size_t>(audio_token_count) * audio_token.size() + 160);
prompt += "<|im_start|>system\n";
prompt += context;
prompt += "<|im_end|>\n<|im_start|>user\n<|audio_start|>";
for (int i = 0; i < audio_token_count; ++i) prompt += audio_token;
prompt += "<|audio_end|><|im_end|>\n<|im_start|>assistant\n";

if (!language.empty()) {
prompt += "language ";
prompt += language;
prompt += "<asr_text>";
}
return prompt;
}

qwen3_asr_result_t qwen3_asr_parse_output(
std::string_view raw_output,
std::optional<std::string_view> forced_language) {
qwen3_asr_result_t result;
const std::string raw = trim(raw_output);
if (raw.empty()) return result;

if (forced_language.has_value()) {
result.language = qwen3_asr_normalize_language(*forced_language);
result.text = raw;
return result;
}

constexpr std::string_view asr_tag = "<asr_text>";
const std::size_t tag_position = raw.find(asr_tag);
if (tag_position == std::string::npos) {
result.text = raw;
return result;
}

const std::string metadata = trim(std::string_view(raw).substr(0, tag_position));
result.text = trim(std::string_view(raw).substr(tag_position + asr_tag.size()));

const std::string metadata_lower = ascii_lower(metadata);
if (metadata_lower.find("language none") != std::string::npos) {
// Silent/empty audio. Preserve unexpected text but do not claim a language.
return result;
}

std::size_t line_start = 0;
while (line_start <= metadata.size()) {
const std::size_t line_end = metadata.find('\n', line_start);
const std::string line = trim(std::string_view(metadata).substr(
line_start,
line_end == std::string::npos ? std::string::npos : line_end - line_start));
const std::string line_lower = ascii_lower(line);
constexpr std::string_view prefix = "language ";
if (line_lower.starts_with(prefix)) {
result.language = qwen3_asr_normalize_language(std::string_view(line).substr(prefix.size()));
break;
}
if (line_end == std::string::npos) break;
line_start = line_end + 1;
}

return result;
}
52 changes: 52 additions & 0 deletions src/include/audio/qwen3_asr_preprocessor.hpp
Original file line number Diff line number Diff line change
@@ -0,0 +1,52 @@
/// \file qwen3_asr_preprocessor.hpp
/// \brief Qwen3-ASR waveform-to-log-mel preprocessing.
#pragma once

#include <cstddef>
#include <vector>

struct qwen3_asr_features_t {
// Feature-major layout matching Transformers input_features:
// [num_mel_bins, num_frames], contiguous in row-major order.
std::vector<float> input_features;
std::vector<int> attention_mask;
int num_mel_bins = 0;
int num_frames = 0;
int valid_frames = 0;
};

class Qwen3ASRPreprocessor {
public:
struct Config {
int sampling_rate = 16000;
int n_fft = 400;
int hop_length = 160;
int num_mel_bins = 128;
int n_window = 50;
int min_length = 8000;
};

Qwen3ASRPreprocessor();
explicit Qwen3ASRPreprocessor(Config config);

/// Convert mono float32 PCM into Qwen3-ASR log-mel features.
///
/// This follows Hugging Face Qwen3ASRFeatureExtractor:
/// - zero-pad waveforms shorter than min_length;
/// - periodic Hann window and centered reflect-padded STFT;
/// - Slaney mel scale and normalization;
/// - log10 clamp, dynamic-range clamp, and (x + 4) / 4 scaling;
/// - right-pad the mel time axis to a multiple of 2 * n_window.
qwen3_asr_features_t extract(const float* samples, std::size_t num_samples) const;
qwen3_asr_features_t extract(const std::vector<float>& samples) const {
return extract(samples.data(), samples.size());
}

/// Number of soft audio tokens produced by the three stride-2 CNN layers.
static int feature_output_length(int mel_frames);

const Config& config() const { return config_; }

private:
Config config_;
};
33 changes: 33 additions & 0 deletions src/include/audio/qwen3_asr_utils.hpp
Original file line number Diff line number Diff line change
@@ -0,0 +1,33 @@
/// \file qwen3_asr_utils.hpp
/// \brief Prompt and output helpers for Qwen3-ASR.
#pragma once

#include <optional>
#include <string>
#include <string_view>

struct qwen3_asr_result_t {
std::string language;
std::string text;
};

/// Normalize a language name to the canonical Qwen3-ASR spelling.
/// Returns an empty string for an empty input.
std::string qwen3_asr_normalize_language(std::string_view language);

/// Whether a canonical or case-insensitive language name is supported.
bool qwen3_asr_is_supported_language(std::string_view language);

/// Build the exact single-audio prompt used by Qwen3-ASR's chat template.
/// When forced_language is set, generation starts after
/// "language <Language><asr_text>" and decoded output is text-only.
std::string qwen3_asr_build_prompt(
std::string_view context,
int audio_token_count,
std::optional<std::string_view> forced_language = std::nullopt);

/// Parse generated text into language and transcription.
/// If forced_language is provided, raw_output is treated as transcription-only.
qwen3_asr_result_t qwen3_asr_parse_output(
std::string_view raw_output,
std::optional<std::string_view> forced_language = std::nullopt);
30 changes: 30 additions & 0 deletions src/test/qwen3_asr_preprocessor/CMakeLists.txt
Original file line number Diff line number Diff line change
@@ -0,0 +1,30 @@
cmake_minimum_required(VERSION 3.22)
project(qwen3_asr_preprocessor_test LANGUAGES CXX)

set(CMAKE_CXX_STANDARD 20)
set(CMAKE_CXX_STANDARD_REQUIRED ON)

find_package(OpenMP)
find_library(FFTW3F_LIBRARY NAMES fftw3f libfftw3f.so.3 REQUIRED)

add_executable(test_qwen3_asr_preprocessor
test.cpp
../../common/audio/qwen3_asr_preprocessor.cpp
../../common/audio/qwen3_asr_utils.cpp
../../common/audio_process_utils/audioproc.cpp
../../common/audio_process_utils/audioprocAVX512.cpp
)

target_include_directories(test_qwen3_asr_preprocessor PRIVATE ../../include)
target_compile_definitions(test_qwen3_asr_preprocessor PRIVATE USEAVX2=1 USEAVX512=1)
target_compile_options(test_qwen3_asr_preprocessor PRIVATE
-O3 -ffast-math -mavx2 -mfma -mavx512f -mavx512dq -mavx512bw -mavx512vl
)
target_link_libraries(test_qwen3_asr_preprocessor PRIVATE ${FFTW3F_LIBRARY})

if(OpenMP_CXX_FOUND)
target_link_libraries(test_qwen3_asr_preprocessor PRIVATE OpenMP::OpenMP_CXX)
endif()

enable_testing()
add_test(NAME qwen3_asr_preprocessor COMMAND test_qwen3_asr_preprocessor)
Loading