CoolFace
Apppublic

philippe83260/seamless-streaming

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
simuleval_transcoder.py429 linesDownload Raw Back to src
1from simuleval.utils.agent import build_system_from_dir2from typing import Any, List, Optional, Tuple, Union3import numpy as np4import soundfile5import io6import asyncio7from simuleval.agents.pipeline import TreeAgentPipeline8from simuleval.agents.states import AgentStates9from simuleval.data.segments import Segment, EmptySegment, SpeechSegment10import threading11import math12import logging13import sys14from pathlib import Path15import time16from g2p_en import G2p17import torch18import traceback19import time20import random21import colorlog22 23from .speech_and_text_output import SpeechAndTextOutput24 25MODEL_SAMPLE_RATE = 16_00026 27logger = logging.getLogger(__name__)28# logger.propagate = False29handler = colorlog.StreamHandler(stream=sys.stdout)30formatter = colorlog.ColoredFormatter(31    "%(log_color)s[%(asctime)s][%(levelname)s][%(module)s]:%(reset)s %(message)s",32    reset=True,33    log_colors={34        "DEBUG": "cyan",35        "INFO": "green",36        "WARNING": "yellow",37        "ERROR": "red",38        "CRITICAL": "red,bg_white",39    },40)41handler.setFormatter(formatter)42logger.addHandler(handler)43logger.setLevel(logging.WARNING)44 45 46class OutputSegments:47    def __init__(self, segments: Union[List[Segment], Segment]):48        if isinstance(segments, Segment):49            segments = [segments]50        self.segments: List[Segment] = [s for s in segments]51 52    @property53    def is_empty(self):54        return all(segment.is_empty for segment in self.segments)55 56    @property57    def finished(self):58        return all(segment.finished for segment in self.segments)59 60    def compute_length(self, g2p):61        lengths = []62        for segment in self.segments:63            if segment.data_type == "text":64                lengths.append(len([x for x in g2p(segment.content) if x != " "]))65            elif segment.data_type == "speech":66                lengths.append(len(segment.content) / MODEL_SAMPLE_RATE)67            elif isinstance(segment, EmptySegment):68                continue69            else:70                logger.warning(71                    f"Unexpected data_type: {segment.data_type} not in 'speech', 'text'"72                )73        return max(lengths)74 75    @classmethod76    def join_output_buffer(77        cls, buffer: List[List[Segment]], output: SpeechAndTextOutput78    ):79        num_segments = len(buffer[0])80        for i in range(num_segments):81            segment_list = [82                buffer[j][i]83                for j in range(len(buffer))84                if buffer[j][i].data_type is not None85            ]86            if len(segment_list) == 0:87                continue88            if len(set(segment.data_type for segment in segment_list)) != 1:89                logger.warning(90                    f"Data type mismatch at {i}: {set(segment.data_type for segment in segment_list)}"91                )92                continue93            data_type = segment_list[0].data_type94            if data_type == "text":95                if output.text is not None:96                    logger.warning("Multiple text outputs, overwriting!")97                output.text = " ".join([segment.content for segment in segment_list])98            elif data_type == "speech":99                if output.speech_samples is not None:100                    logger.warning("Multiple speech outputs, overwriting!")101                speech_out = []102                for segment in segment_list:103                    speech_out += segment.content104                output.speech_samples = speech_out105                output.speech_sample_rate = segment.sample_rate106            elif isinstance(segment_list[0], EmptySegment):107                continue108            else:109                logger.warning(110                    f"Invalid output buffer data type: {data_type}, expected 'speech' or 'text"111                )112 113        return output114 115    def __repr__(self) -> str:116        repr_str = str(self.segments)117        return f"{self.__class__.__name__}(\n\t{repr_str}\n)"118 119 120class SimulevalTranscoder:121    def __init__(self, agent, sample_rate, debug, buffer_limit):122        self.agent = agent.agent123        self.has_expressive = agent.has_expressive124        self.input_queue = asyncio.Queue()125        self.output_queue = asyncio.Queue()126        self.states = self.agent.build_states()127        if debug:128            self.get_states_root().debug = True129        self.incoming_sample_rate = sample_rate130        self.close = False131        self.g2p = G2p()132 133        # buffer all outgoing translations within this amount of time134        self.output_buffer_idle_ms = 5000135        self.output_buffer_size_limit = (136            buffer_limit  # phonemes for text, seconds for speech137        )138        self.output_buffer_cur_size = 0139        self.output_buffer: List[List[Segment]] = []140        self.speech_output_sample_rate = None141 142        self.last_output_ts = time.time() * 1000143        self.timeout_ms = (144            30000  # close the transcoder thread after this amount of silence145        )146        self.first_input_ts = None147        self.first_output_ts = None148        self.debug = debug149        self.debug_ts = f"{time.time()}_{random.randint(1000, 9999)}"150        if self.debug:151            debug_folder = Path(__file__).resolve().parent.parent / "debug"152            self.test_incoming_wav = soundfile.SoundFile(153                debug_folder / f"{self.debug_ts}_test_incoming.wav",154                mode="w+",155                format="WAV",156                subtype="PCM_16",157                samplerate=self.incoming_sample_rate,158                channels=1,159            )160            self.get_states_root().test_input_segments_wav = soundfile.SoundFile(161                debug_folder / f"{self.debug_ts}_test_input_segments.wav",162                mode="w+",163                format="WAV",164                samplerate=MODEL_SAMPLE_RATE,165                channels=1,166            )167 168    def get_states_root(self) -> AgentStates:169        if isinstance(self.agent, TreeAgentPipeline):170            # self.states is a dict171            return self.states[self.agent.source_module]172        else:173            # self.states is a list174            return self.states[0]175 176    def reset_states(self):177        if isinstance(self.agent, TreeAgentPipeline):178            states_iter = self.states.values()179        else:180            states_iter = self.states181        for state in states_iter:182            state.reset()183 184    def debug_log(self, *args):185        if self.debug:186            logger.info(*args)187 188    @classmethod189    def build_agent(cls, model_path, config_name):190        logger.info(f"Building simuleval agent: {model_path}, {config_name}")191        agent = build_system_from_dir(192            Path(__file__).resolve().parent.parent / f"models/{model_path}",193            config_name=config_name,194        )195        device = torch.device("cuda" if torch.cuda.is_available() else "cpu")196        agent.to(device, fp16=True)197        logger.info(198            f"Successfully built simuleval agent {model_path} on device {device}"199        )200 201        return agent202 203    def process_incoming_bytes(self, incoming_bytes, dynamic_config):204        # TODO: We probably want to do some validation on dynamic_config to ensure it has what we needs205        segment, sr = self._preprocess_wav(incoming_bytes)206        segment = SpeechSegment(207            content=segment,208            sample_rate=sr,209            tgt_lang=dynamic_config.get("targetLanguage"),210            config=dynamic_config,211        )212        if dynamic_config.get("expressive") is True and self.has_expressive is False:213            logger.warning(214                "Passing 'expressive' but the agent does not support expressive output!"215            )216        # # segment is array([0, 0, 0, ..., 0, 0, 0], dtype=int16)217        self.input_queue.put_nowait(segment)218 219    def get_input_segment(self):220        if self.input_queue.empty():221            return None222        chunk = self.input_queue.get_nowait()223        self.input_queue.task_done()224        return chunk225 226    def convert_waveform(227        self,228        waveform: Union[np.ndarray, torch.Tensor],229        sample_rate: int,230        normalize_volume: bool = False,231        to_mono: bool = False,232        to_sample_rate: Optional[int] = None,233    ) -> Tuple[Union[np.ndarray, torch.Tensor], int]:234        """convert a waveform:235        - to a target sample rate236        - from multi-channel to mono channel237        - volume normalization238 239        Args:240            waveform (numpy.ndarray or torch.Tensor): 2D original waveform241                (channels x length)242            sample_rate (int): original sample rate243            normalize_volume (bool): perform volume normalization244            to_mono (bool): convert to mono channel if having multiple channels245            to_sample_rate (Optional[int]): target sample rate246        Returns:247            waveform (numpy.ndarray): converted 2D waveform (channels x length)248            sample_rate (float): target sample rate249        """250        try:251            import torchaudio.sox_effects as ta_sox252        except ImportError:253            raise ImportError("Please install torchaudio: pip install torchaudio")254 255        effects = []256        if normalize_volume:257            effects.append(["gain", "-n"])258        if to_sample_rate is not None and to_sample_rate != sample_rate:259            effects.append(["rate", f"{to_sample_rate}"])260        if to_mono and waveform.shape[0] > 1:261            effects.append(["channels", "1"])262        if len(effects) > 0:263            is_np_input = isinstance(waveform, np.ndarray)264            _waveform = torch.from_numpy(waveform) if is_np_input else waveform265            converted, converted_sample_rate = ta_sox.apply_effects_tensor(266                _waveform, sample_rate, effects267            )268            if is_np_input:269                converted = converted.numpy()270            return converted, converted_sample_rate271        return waveform, sample_rate272 273    def _preprocess_wav(self, data: Any) -> Tuple[np.ndarray, int]:274        segment, sample_rate = soundfile.read(275            io.BytesIO(data),276            dtype="float32",277            always_2d=True,278            frames=-1,279            start=0,280            format="RAW",281            subtype="PCM_16",282            samplerate=self.incoming_sample_rate,283            channels=1,284        )285        if self.debug:286            self.test_incoming_wav.seek(0, soundfile.SEEK_END)287            self.test_incoming_wav.write(segment)288 289        segment = segment.T290        segment, new_sample_rate = self.convert_waveform(291            segment,292            sample_rate,293            normalize_volume=False,294            to_mono=True,295            to_sample_rate=MODEL_SAMPLE_RATE,296        )297 298        assert MODEL_SAMPLE_RATE == new_sample_rate299        segment = segment.squeeze(axis=0)300        return segment, new_sample_rate301 302    def process_pipeline_impl(self, input_segment):303        try:304            with torch.no_grad():305                output_segment = OutputSegments(306                    self.agent.pushpop(input_segment, self.states)307                )308            if (309                self.get_states_root().first_input_ts is not None310                and self.first_input_ts is None311            ):312                # TODO: this is hacky313                self.first_input_ts = self.get_states_root().first_input_ts314 315            if not output_segment.is_empty:316                self.output_queue.put_nowait(output_segment)317 318            if output_segment.finished:319                self.debug_log("OUTPUT SEGMENT IS FINISHED. Resetting states.")320 321                self.reset_states()322 323                if self.debug:324                    # when we rebuild states, this value is reset to whatever325                    # is in the system dir config, which defaults debug=False.326                    self.get_states_root().debug = True327        except Exception as e:328            logger.error(f"Got exception while processing pipeline: {e}")329            traceback.print_exc()330        return input_segment331 332    def process_pipeline_loop(self):333        if self.close:334            return  # closes the thread335 336        self.debug_log("processing_pipeline")337        while not self.close:338            input_segment = self.get_input_segment()339            if input_segment is None:340                if self.get_states_root().is_fresh_state:  # TODO: this is hacky341                    time.sleep(0.3)342                else:343                    time.sleep(0.03)344                continue345            self.process_pipeline_impl(input_segment)346        self.debug_log("finished processing_pipeline")347 348    def process_pipeline_once(self):349        if self.close:350            return351 352        self.debug_log("processing pipeline once")353        input_segment = self.get_input_segment()354        if input_segment is None:355            return356        self.process_pipeline_impl(input_segment)357        self.debug_log("finished processing_pipeline_once")358 359    def get_output_segment(self):360        if self.output_queue.empty():361            return None362 363        output_chunk = self.output_queue.get_nowait()364        self.output_queue.task_done()365        return output_chunk366 367    def start(self):368        self.debug_log("starting transcoder in a thread")369        threading.Thread(target=self.process_pipeline_loop).start()370 371    def first_translation_time(self):372        return round((self.first_output_ts - self.first_input_ts) / 1000, 2)373 374    def get_buffered_output(self) -> SpeechAndTextOutput:375        now = time.time() * 1000376        self.debug_log(f"get_buffered_output queue size: {self.output_queue.qsize()}")377        while not self.output_queue.empty():378            tmp_out = self.get_output_segment()379            if tmp_out and tmp_out.compute_length(self.g2p) > 0:380                if len(self.output_buffer) == 0:381                    self.last_output_ts = now382                self._populate_output_buffer(tmp_out)383                self._increment_output_buffer_size(tmp_out)384 385                if tmp_out.finished:386                    self.debug_log("tmp_out.finished")387                    res = self._gather_output_buffer_data(final=True)388                    self.debug_log(f"gathered output data: {res}")389                    self.output_buffer = []390                    self.increment_output_buffer_size = 0391                    self.last_output_ts = now392                    self.first_output_ts = now393                    return res394            else:395                self.debug_log("tmp_out.compute_length is not > 0")396 397        if len(self.output_buffer) > 0 and (398            now - self.last_output_ts >= self.output_buffer_idle_ms399            or self.output_buffer_cur_size >= self.output_buffer_size_limit400        ):401            self.debug_log(402                "[get_buffered_output] output_buffer is not empty. getting res to return."403            )404            self.last_output_ts = now405            res = self._gather_output_buffer_data(final=False)406            self.debug_log(f"gathered output data: {res}")407            self.output_buffer = []408            self.output_buffer_phoneme_count = 0409            self.first_output_ts = now410            return res411        else:412            self.debug_log("[get_buffered_output] output_buffer is empty...")413            return None414 415    def _gather_output_buffer_data(self, final):416        output = SpeechAndTextOutput()417        output.final = final418        output = OutputSegments.join_output_buffer(self.output_buffer, output)419        return output420 421    def _increment_output_buffer_size(self, segment: OutputSegments):422        self.output_buffer_cur_size += segment.compute_length(self.g2p)423 424    def _populate_output_buffer(self, segment: OutputSegments):425        self.output_buffer.append(segment.segments)426 427    def _compute_phoneme_count(self, string: str) -> int:428        return len([x for x in self.g2p(string) if x != " "])429