main
6f8fc97 · 1 month ago 10 commits
  1#include <chrono>
  2#include <fstream>
  3#include <iostream>
  4#include <memory>
  5#include <vector>
  6
  7#include "audio_buffer.h"
  8#include "gain_controller2.h"
  9
 10#define MAX_PROCESSABLE_MS 10
 11#define SUPPORTED_SAMPLE_RATES {8000, 16000, 32000, 48000}
 12
 13struct WAVHeader {
 14    char chunkID[4];  // "RIFF"
 15    uint32_t chunkSize;
 16    char format[4];      // "WAVE"
 17    char subchunkID[4];  // "fmt "
 18    uint32_t subchunkSize;
 19    uint16_t audioFormat;
 20    uint16_t numChannels;
 21    uint32_t sampleRate;
 22    uint32_t byteRate;
 23    uint16_t blockAlign;
 24    uint16_t bitsPerSample;
 25    char dataID[4];  // "data"
 26    uint32_t dataSize;
 27
 28    void print() {
 29        std::cout << "chunkID: " << std::string(chunkID, 4) << std::endl;
 30        std::cout << "chunkSize: " << chunkSize << std::endl;
 31        std::cout << "format: " << std::string(format, 4) << std::endl;
 32        std::cout << "subchunkID: " << std::string(subchunkID, 4) << std::endl;
 33        std::cout << "subchunkSize: " << subchunkSize << std::endl;
 34        std::cout << "audioFormat: " << audioFormat << std::endl;
 35        std::cout << "numChannels: " << numChannels << std::endl;
 36        std::cout << "sampleRate: " << sampleRate << std::endl;
 37        std::cout << "byteRate: " << byteRate << std::endl;
 38        std::cout << "blockAlign: " << blockAlign << std::endl;
 39        std::cout << "bitsPerSample: " << bitsPerSample << std::endl;
 40        std::cout << "dataID: " << std::string(dataID, 4) << std::endl;
 41        std::cout << "dataSize: " << dataSize << std::endl;
 42    }
 43};
 44
 45struct AudioData {
 46    WAVHeader header;
 47    std::vector<int16_t> samples;
 48    double duration;
 49};
 50
 51class AGC {
 52   private:
 53    std::unique_ptr<webrtc::GainController2> gainController_m;  // Use smart pointer
 54    std::unique_ptr<webrtc::AudioBuffer> audioBuffer_m;         // Use smart pointer
 55    webrtc::StreamConfig streamConfig_m;
 56    int sampleRate;
 57
 58   public:
 59    AGC(int sampleRate) : streamConfig_m(sampleRate, 1), sampleRate(sampleRate) {
 60        audioBuffer_m =
 61            std::make_unique<webrtc::AudioBuffer>(sampleRate, 1, sampleRate, 1, sampleRate, 1);
 62        initialize();
 63    }
 64    ~AGC() = default;
 65
 66    void initialize() {
 67        webrtc::AudioProcessing::Config::GainController2 config;
 68        webrtc::InputVolumeController::Config inputVolumeControllerConfig;
 69        config.enabled = true;
 70
 71        // config.input_volume_controller.enabled = true;  enable only if direct mic input is used
 72
 73        config.adaptive_digital.enabled = true;
 74        config.adaptive_digital.headroom_db = 5.0f;
 75        config.adaptive_digital.max_gain_db = 30.0f;
 76        config.adaptive_digital.initial_gain_db = 10.0f;
 77        config.adaptive_digital.max_gain_change_db_per_second = 5.0f;
 78        config.adaptive_digital.max_output_noise_level_dbfs = -50.f;
 79
 80        config.fixed_digital.gain_db = 2.0f;
 81
 82        gainController_m = std::make_unique<webrtc::GainController2>(
 83            config, inputVolumeControllerConfig, sampleRate, 1, true); 
 84				// Enable the internal VAD so GainController2 computes ↗ 
 85				// speech_probability required by the adaptive digital AGC.
 86				// When disabled, callers must supply speech_probability to Process().
 87    }
 88
 89    void process(std::vector<int16_t> &frame) {
 90        size_t expectedSize = (size_t)(sampleRate * MAX_PROCESSABLE_MS / 1000);
 91        if (frame.empty() || frame.size() != expectedSize) {
 92            std::cerr << "Invalid frame size" << std::endl;
 93            return;
 94        }
 95        audioBuffer_m->CopyFrom(frame.data(), streamConfig_m);
 96        gainController_m->Process(std::nullopt, false, audioBuffer_m.get());
 97        audioBuffer_m->CopyTo(streamConfig_m, frame.data());
 98    }
 99};
100
101class AudioProcessor {
102   private:
103    AudioData audioData;
104    std::string inputFile;
105    std::string outputFile;
106    std::unique_ptr<AGC> agcManager;
107
108    void processFrame(const int16_t *input, int16_t *output, int frameSize);
109    void readAudioWithFFmpeg();
110    void writeAudioWithFFmpeg();
111    void readRawAudioFile(bool headerOnly = false);
112    void writeRawAudioFile();
113    void performAGC();
114    double calculateDuration(const WAVHeader &header) {
115        return static_cast<double>(header.dataSize) / (header.byteRate);
116    }
117
118   public:
119    AudioProcessor(std::string inputFile, std::string outputFile)
120        : inputFile(inputFile), outputFile(outputFile) {
121        std::cout << "Input file: " << inputFile << std::endl;
122        std::cout << "Output file: " << outputFile << std::endl;
123    }
124    ~AudioProcessor() = default;
125    void resample(uint32_t expectedSampleRate);
126
127    void processWithFFmpeg() {
128        readRawAudioFile(true);
129        readAudioWithFFmpeg();
130        performAGC();
131        writeAudioWithFFmpeg();
132    }
133
134    void processWithCustomResampler(uint32_t expectedSampleRate = 48000) {
135        readRawAudioFile();
136        resample(expectedSampleRate);
137        performAGC();
138        writeRawAudioFile();
139    }
140};