diff --git a/whisper_live/server.py b/whisper_live/server.py index b6efc77..94e6163 100644 --- a/whisper_live/server.py +++ b/whisper_live/server.py @@ -229,8 +229,7 @@ class TranscriptionServer: websocket.close() return False # Indicates that the connection should not continue - if self.backend.is_tensorrt(): - self.vad_detector = VoiceActivityDetector(frame_rate=self.RATE) + self.vad_detector = VoiceActivityDetector(frame_rate=self.RATE) self.initialize_client(websocket, options, faster_whisper_custom_model_path, whisper_tensorrt_path, trt_multilingual) return True @@ -248,17 +247,15 @@ class TranscriptionServer: frame_np = self.get_audio_from_websocket(websocket) client = self.client_manager.get_client(websocket) if frame_np is False: - if self.backend.is_tensorrt(): - client.set_eos(True) + client.set_eos(True) return False - if self.backend.is_tensorrt(): - voice_active = self.voice_activity(websocket, frame_np) - if voice_active: - self.no_voice_activity_chunks = 0 - client.set_eos(False) - if self.use_vad and not voice_active: - return True + voice_active = self.voice_activity(websocket, frame_np) + if voice_active: + self.no_voice_activity_chunks = 0 + client.set_eos(False) + if self.use_vad and not voice_active: + return True client.add_frames(frame_np) return True @@ -331,13 +328,8 @@ class TranscriptionServer: raise ValueError(f"Custom faster_whisper model '{faster_whisper_custom_model_path}' is not a valid path.") if whisper_tensorrt_path is not None and not os.path.exists(whisper_tensorrt_path): raise ValueError(f"TensorRT model '{whisper_tensorrt_path}' is not a valid path.") - if single_model: - if faster_whisper_custom_model_path or whisper_tensorrt_path: - logging.info("Custom model option was provided. Switching to single model mode.") - self.single_model = True - # TODO: load model initially - else: - logging.info("Single model mode currently only works with custom models.") + + self.single_model = single_model if not BackendType.is_valid(backend): raise ValueError(f"{backend} is not a valid backend type. Choose backend from {BackendType.valid_types()}") with serve( @@ -416,6 +408,7 @@ class ServeClientBase(object): self.add_pause_thresh = 3 # add a blank to segment list as a pause(no speech) for 3 seconds self.transcript = [] self.send_last_n_segments = 10 + self.eos = False # text formatting self.pick_previous_segments = 2 @@ -423,6 +416,18 @@ class ServeClientBase(object): # threading self.lock = threading.Lock() + def set_eos(self, eos): + """ + Sets the End of Speech (EOS) flag. + + Args: + eos (bool): The value to set for the EOS flag. + """ + self.lock.acquire() + self.eos = eos + self.lock.release() + + def speech_to_text(self): raise NotImplementedError @@ -543,7 +548,8 @@ class ServeClientBase(object): self.websocket.send( json.dumps({ "uid": self.client_uid, - "segments": segments, + "text": segments, + "eos": self.eos }) ) except Exception as e: @@ -648,17 +654,6 @@ class ServeClientTensorRT(ServeClientBase): for i in range(warmup_steps): self.transcriber.transcribe(mel) - def set_eos(self, eos): - """ - Sets the End of Speech (EOS) flag. - - Args: - eos (bool): The value to set for the EOS flag. - """ - self.lock.acquire() - self.eos = eos - self.lock.release() - def handle_transcription_output(self, last_segment, duration): """ Handle the transcription output, updating the transcript and sending data to the client. @@ -784,7 +779,7 @@ class ServeClientFasterWhisper(ServeClientBase): self.task = task self.initial_prompt = initial_prompt self.vad_parameters = vad_parameters or {"threshold": 0.5} - self.no_speech_thresh = 0.45 + self.no_speech_thresh = 0.35 device = "cuda" if torch.cuda.is_available() else "cpu" @@ -796,6 +791,7 @@ class ServeClientFasterWhisper(ServeClientBase): self.create_model(device) ServeClientFasterWhisper.SINGLE_MODEL = self.transcriber else: + print("Re-using already initialized model.") self.transcriber = ServeClientFasterWhisper.SINGLE_MODEL else: self.create_model(device) @@ -888,8 +884,9 @@ class ServeClientFasterWhisper(ServeClientBase): initial_prompt=self.initial_prompt, language=self.language, task=self.task, - vad_filter=self.use_vad, - vad_parameters=self.vad_parameters if self.use_vad else None) + vad_filter=False, + vad_parameters=self.vad_parameters if self.use_vad else None, + beam_size=5) if ServeClientFasterWhisper.SINGLE_MODEL: ServeClientFasterWhisper.SINGLE_MODEL_LOCK.release() @@ -932,17 +929,16 @@ class ServeClientFasterWhisper(ServeClientBase): result (str): The result from whisper inference i.e. the list of segments. duration (float): Duration of the transcribed audio chunk. """ - segments = [] if len(result): - self.t_start = None last_segment = self.update_segments(result, duration) - segments = self.prepare_segments(last_segment) - else: - # show previous output if there is pause i.e. no output from whisper - segments = self.get_previous_output() - if len(segments): - self.send_transcription_to_client(segments) + if len(self.text): + if self.eos and last_segment is None: + self.send_transcription_to_client(' '.join([s.strip() for s in self.text])) + self.set_eos(False) + self.text = [] + elif not self.eos: + self.send_transcription_to_client(' '.join([s.strip() for s in self.text])) def speech_to_text(self): """ @@ -972,7 +968,12 @@ class ServeClientFasterWhisper(ServeClientBase): self.clip_audio_if_no_valid_segment() input_bytes, duration = self.get_audio_chunk_for_processing() - if duration < 1.0: + if duration < 0.6: + if len(self.text) and self.eos: + self.send_transcription_to_client(' '.join([s.strip() for s in self.text])) + self.set_eos(False) + self.text = [] + time.sleep(0.1) continue try: input_sample = input_bytes.copy() @@ -980,7 +981,7 @@ class ServeClientFasterWhisper(ServeClientBase): if result is None or self.language is None: self.timestamp_offset += duration - time.sleep(0.25) # wait for voice activity, result is None when no voice activity + time.sleep(0.1) # wait for voice activity, result is None when no voice activity continue self.handle_transcription_output(result, duration) @@ -1029,13 +1030,13 @@ class ServeClientFasterWhisper(ServeClientBase): dict or None: The last processed segment with its start time, end time, and transcribed text. Returns None if there are no valid segments to process. """ + last_segment = None offset = None self.current_out = '' # process complete segments if len(segments) > 1: for i, s in enumerate(segments[:-1]): text_ = s.text - self.text.append(text_) start, end = self.timestamp_offset + s.start, self.timestamp_offset + min(duration, s.end) if start >= end: @@ -1043,15 +1044,17 @@ class ServeClientFasterWhisper(ServeClientBase): if s.no_speech_prob > self.no_speech_thresh: continue + self.text.append(text_) self.transcript.append(self.format_segment(start, end, text_)) offset = min(duration, s.end) - self.current_out += segments[-1].text - last_segment = self.format_segment( - self.timestamp_offset + segments[-1].start, - self.timestamp_offset + min(duration, segments[-1].end), - self.current_out - ) + if segments[-1].no_speech_prob <= self.no_speech_thresh: + self.current_out += segments[-1].text + last_segment = self.format_segment( + self.timestamp_offset + segments[-1].start, + self.timestamp_offset + min(duration, segments[-1].end), + self.current_out + ) # if same incomplete segment is seen multiple times then update the offset # and append the segment to the list @@ -1060,7 +1063,7 @@ class ServeClientFasterWhisper(ServeClientBase): else: self.same_output_threshold = 0 - if self.same_output_threshold > 5: + if self.same_output_threshold > 2: if not len(self.text) or self.text[-1].strip().lower() != self.current_out.strip().lower(): self.text.append(self.current_out) self.transcript.append(self.format_segment(