philippe83260/seamless-streaming
0
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 