import torch
import torchaudio
from transformers import AutoProcessor, AutoModelForCTC
from model.decoder import GreedyDecoder
from model.token_set import TokenSet


class AudioTranscriber:
    def __init__(self, model: AutoModelForCTC, processor: AutoProcessor):
        self.model = model
        self.processor = processor
        self.token_set = TokenSet.from_processor(processor)
        self.greedy_decoder = GreedyDecoder(self.token_set)

    def transcribe_audio(self, audio_waveform, sample_rate):
        # If the audio sample rate is not 16kHz, resample it
        if sample_rate != 16000:
            audio_waveform = torchaudio.functional.resample(
                audio_waveform, orig_freq=sample_rate, new_freq=16000
            )
            sample_rate = 16000

        # If stereo, convert to mono by averaging the channels
        if audio_waveform.shape[0] > 1:
            audio_waveform = torch.mean(audio_waveform, dim=0)

        # Prepare the inputs for the model
        inputs = self.processor(
            audio_waveform.squeeze(), sampling_rate=sample_rate,
            return_tensors="pt"
        )

        # Perform inference
        with torch.no_grad():
            logits = self.model(inputs.input_values).logits
        result = self.greedy_decoder(logits)[0]
        return result
