Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
3bac87ee52 | ||
|
|
9a01701019 | ||
|
|
693954ee23 | ||
|
|
bf4ba91e7a | ||
|
|
0828353253 | ||
|
|
1997c7ad8f | ||
|
|
b1e62440e4 | ||
|
|
d549a5eb6a | ||
|
|
77bfb08d76 |
+80
-3
@@ -1,4 +1,3 @@
|
||||
#
|
||||
import os
|
||||
import subprocess
|
||||
import importlib.util
|
||||
@@ -13,11 +12,13 @@ from PIL import Image
|
||||
from comfy.cli_args import args
|
||||
python = sys.executable
|
||||
|
||||
# print("sys.path", sys.path)
|
||||
|
||||
#修复 sys.stdout.isatty() object has no attribute 'isatty'
|
||||
try:
|
||||
sys.stdout.isatty()
|
||||
except:
|
||||
print('#fix sys.stdout.isatty')
|
||||
# print('#fix sys.stdout.isatty')
|
||||
sys.stdout.isatty = lambda: False
|
||||
|
||||
llama_port=None
|
||||
@@ -575,7 +576,7 @@ async def mixlab_app_handler(request):
|
||||
return web.Response(text=html_data, content_type='text/html')
|
||||
else:
|
||||
return web.Response(text="HTML file not found", status=404)
|
||||
|
||||
|
||||
|
||||
@routes.post('/mixlab/workflow')
|
||||
async def mixlab_workflow_hander(request):
|
||||
@@ -718,6 +719,44 @@ async def post_prompt_result(request):
|
||||
return web.json_response({"result":res})
|
||||
|
||||
|
||||
|
||||
def start_local_live_thread(data):
|
||||
import asyncio
|
||||
from VoiceStreamAI.server import Server
|
||||
from VoiceStreamAI.asr.asr_factory import ASRFactory
|
||||
from VoiceStreamAI.vad.vad_factory import VADFactory
|
||||
|
||||
model="large-v3"
|
||||
if "model" in data:
|
||||
model=data['model']
|
||||
|
||||
vad_pipeline = VADFactory.create_vad_pipeline("pyannote")
|
||||
#device
|
||||
asr_pipeline = ASRFactory.create_asr_pipeline("faster_whisper", **{"model_size":model})
|
||||
|
||||
port=8765
|
||||
if 'port' in data:
|
||||
port=data['port']
|
||||
|
||||
llm_port=9000
|
||||
if 'llm_port' in data:
|
||||
llm_port=data['llm_port']
|
||||
|
||||
server = Server(vad_pipeline,
|
||||
asr_pipeline,
|
||||
host="127.0.0.1",
|
||||
port=port,
|
||||
sampling_rate=16000,
|
||||
samples_width=2,
|
||||
llm_port=llm_port
|
||||
)
|
||||
|
||||
loop = asyncio.new_event_loop()
|
||||
asyncio.set_event_loop(loop)
|
||||
loop.run_until_complete(server.start())
|
||||
loop.run_forever()
|
||||
|
||||
|
||||
async def start_local_llm(data):
|
||||
global llama_port,llama_model,llama_chat_format
|
||||
if llama_port and llama_model and llama_chat_format:
|
||||
@@ -747,6 +786,8 @@ async def start_local_llm(data):
|
||||
|
||||
|
||||
chat_format="chatml"
|
||||
if "model" in data and "function-calling" in data['model']:
|
||||
chat_format="functionary-v2"
|
||||
|
||||
model_alias=os.path.basename(model)
|
||||
|
||||
@@ -829,6 +870,42 @@ async def my_hander_method(request):
|
||||
|
||||
return web.json_response(result)
|
||||
|
||||
|
||||
@routes.post('/mixlab/start_live')
|
||||
async def mixlab_live_start_handler(request):
|
||||
import threading
|
||||
llm=await start_local_llm({
|
||||
"model":"Phi-3-mini-4k-instruct-Q5_K_S.gguf",
|
||||
"n_gpu_layers":2
|
||||
})
|
||||
|
||||
# {"port":llama_port,"model":llama_model,"chat_format":llama_chat_format}
|
||||
|
||||
os.environ['HF_ENDPOINT'] = 'https://hf-mirror.com' #hf_hub_download 里的下载地址修改
|
||||
os.environ['PYANNOTE_AUTH_TOKEN'] = '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'
|
||||
|
||||
# Create and start the thread
|
||||
data = {
|
||||
"llm_port":llm['port'],
|
||||
"port":8725,
|
||||
"model":"large-v3"
|
||||
} # Replace with your actual data if needed
|
||||
thread = threading.Thread(target=start_local_live_thread, args=(data,))
|
||||
thread.start()
|
||||
|
||||
return web.json_response(data)
|
||||
|
||||
|
||||
@routes.get('/mixlab/live')
|
||||
async def mixlab_live_handler(request):
|
||||
html_file = os.path.join(current_path, "web/live.html")
|
||||
if os.path.exists(html_file):
|
||||
with open(html_file, 'r', encoding='utf-8', errors='ignore') as f:
|
||||
html_data = f.read()
|
||||
return web.Response(text=html_data, content_type='text/html')
|
||||
else:
|
||||
return web.Response(text="HTML file not found", status=404)
|
||||
|
||||
# 重启服务
|
||||
@routes.post('/mixlab/re_start')
|
||||
def re_start(request):
|
||||
|
||||
@@ -0,0 +1,12 @@
|
||||
from VoiceStreamAI.asr.whisper_asr import WhisperASR
|
||||
from VoiceStreamAI.asr.faster_whisper_asr import FasterWhisperASR
|
||||
|
||||
class ASRFactory:
|
||||
@staticmethod
|
||||
def create_asr_pipeline(type, **kwargs):
|
||||
if type == "whisper":
|
||||
return WhisperASR(**kwargs)
|
||||
if type == "faster_whisper":
|
||||
return FasterWhisperASR(**kwargs)
|
||||
else:
|
||||
raise ValueError(f"Unknown ASR pipeline type: {type}")
|
||||
@@ -0,0 +1,9 @@
|
||||
class ASRInterface:
|
||||
async def transcribe(self, client):
|
||||
"""
|
||||
Transcribe the given audio data.
|
||||
|
||||
:param client: The client object with all the member variables including the buffer
|
||||
:return: The transcription structure, see for example the faster_whisper_asr.py file.
|
||||
"""
|
||||
raise NotImplementedError("This method should be implemented by subclasses.")
|
||||
@@ -0,0 +1,142 @@
|
||||
import os
|
||||
from faster_whisper import WhisperModel
|
||||
|
||||
from VoiceStreamAI.asr.asr_interface import ASRInterface
|
||||
from VoiceStreamAI.audio_utils import save_audio_to_file
|
||||
|
||||
import folder_paths
|
||||
|
||||
language_codes = {
|
||||
"afrikaans": "af",
|
||||
"amharic": "am",
|
||||
"arabic": "ar",
|
||||
"assamese": "as",
|
||||
"azerbaijani": "az",
|
||||
"bashkir": "ba",
|
||||
"belarusian": "be",
|
||||
"bulgarian": "bg",
|
||||
"bengali": "bn",
|
||||
"tibetan": "bo",
|
||||
"breton": "br",
|
||||
"bosnian": "bs",
|
||||
"catalan": "ca",
|
||||
"czech": "cs",
|
||||
"welsh": "cy",
|
||||
"danish": "da",
|
||||
"german": "de",
|
||||
"greek": "el",
|
||||
"english": "en",
|
||||
"spanish": "es",
|
||||
"estonian": "et",
|
||||
"basque": "eu",
|
||||
"persian": "fa",
|
||||
"finnish": "fi",
|
||||
"faroese": "fo",
|
||||
"french": "fr",
|
||||
"galician": "gl",
|
||||
"gujarati": "gu",
|
||||
"hausa": "ha",
|
||||
"hawaiian": "haw",
|
||||
"hebrew": "he",
|
||||
"hindi": "hi",
|
||||
"croatian": "hr",
|
||||
"haitian": "ht",
|
||||
"hungarian": "hu",
|
||||
"armenian": "hy",
|
||||
"indonesian": "id",
|
||||
"icelandic": "is",
|
||||
"italian": "it",
|
||||
"japanese": "ja",
|
||||
"javanese": "jw",
|
||||
"georgian": "ka",
|
||||
"kazakh": "kk",
|
||||
"khmer": "km",
|
||||
"kannada": "kn",
|
||||
"korean": "ko",
|
||||
"latin": "la",
|
||||
"luxembourgish": "lb",
|
||||
"lingala": "ln",
|
||||
"lao": "lo",
|
||||
"lithuanian": "lt",
|
||||
"latvian": "lv",
|
||||
"malagasy": "mg",
|
||||
"maori": "mi",
|
||||
"macedonian": "mk",
|
||||
"malayalam": "ml",
|
||||
"mongolian": "mn",
|
||||
"marathi": "mr",
|
||||
"malay": "ms",
|
||||
"maltese": "mt",
|
||||
"burmese": "my",
|
||||
"nepali": "ne",
|
||||
"dutch": "nl",
|
||||
"norwegian nynorsk": "nn",
|
||||
"norwegian": "no",
|
||||
"occitan": "oc",
|
||||
"punjabi": "pa",
|
||||
"polish": "pl",
|
||||
"pashto": "ps",
|
||||
"portuguese": "pt",
|
||||
"romanian": "ro",
|
||||
"russian": "ru",
|
||||
"sanskrit": "sa",
|
||||
"sindhi": "sd",
|
||||
"sinhalese": "si",
|
||||
"slovak": "sk",
|
||||
"slovenian": "sl",
|
||||
"shona": "sn",
|
||||
"somali": "so",
|
||||
"albanian": "sq",
|
||||
"serbian": "sr",
|
||||
"sundanese": "su",
|
||||
"swedish": "sv",
|
||||
"swahili": "sw",
|
||||
"tamil": "ta",
|
||||
"telugu": "te",
|
||||
"tajik": "tg",
|
||||
"thai": "th",
|
||||
"turkmen": "tk",
|
||||
"tagalog": "tl",
|
||||
"turkish": "tr",
|
||||
"tatar": "tt",
|
||||
"ukrainian": "uk",
|
||||
"urdu": "ur",
|
||||
"uzbek": "uz",
|
||||
"vietnamese": "vi",
|
||||
"yiddish": "yi",
|
||||
"yoruba": "yo",
|
||||
"chinese": "zh",
|
||||
"cantonese": "yue",
|
||||
}
|
||||
|
||||
|
||||
class FasterWhisperASR(ASRInterface):
|
||||
def __init__(self, **kwargs):
|
||||
model_size = kwargs.get('model_size', "large-v3")
|
||||
device = kwargs.get('device', "cuda")
|
||||
model_root = os.path.join(folder_paths.models_dir, "whisper")
|
||||
# Run on GPU with FP16
|
||||
self.asr_pipeline = WhisperModel(model_size, device=device, compute_type="float16",download_root=model_root)
|
||||
|
||||
async def transcribe(self, client):
|
||||
file_path = await save_audio_to_file(client.scratch_buffer, client.get_file_name())
|
||||
|
||||
language = None if client.config['language'] is None else language_codes.get(client.config['language'].lower())
|
||||
segments, info = self.asr_pipeline.transcribe(file_path, word_timestamps=True, language=language)
|
||||
|
||||
segments = list(segments) # The transcription will actually run here.
|
||||
os.remove(file_path)
|
||||
|
||||
flattened_words = [word for segment in segments for word in segment.words]
|
||||
|
||||
to_return = {
|
||||
"language": info.language,
|
||||
"language_probability": info.language_probability,
|
||||
"text": ' '.join([s.text.strip() for s in segments]),
|
||||
"words":
|
||||
[
|
||||
{"word": w.word, "start": w.start, "end": w.end, "probability":w.probability} for w in flattened_words
|
||||
]
|
||||
}
|
||||
return to_return
|
||||
|
||||
@@ -0,0 +1,27 @@
|
||||
from transformers import pipeline
|
||||
from VoiceStreamAI.asr.asr_interface import ASRInterface
|
||||
from VoiceStreamAI.audio_utils import save_audio_to_file
|
||||
import os
|
||||
|
||||
class WhisperASR(ASRInterface):
|
||||
def __init__(self, **kwargs):
|
||||
model_name = kwargs.get('model_name', "openai/whisper-large-v3")
|
||||
self.asr_pipeline = pipeline("automatic-speech-recognition", model=model_name)
|
||||
|
||||
async def transcribe(self, client):
|
||||
file_path = await save_audio_to_file(client.scratch_buffer, client.get_file_name())
|
||||
|
||||
if client.config['language'] is not None:
|
||||
to_return = self.asr_pipeline(file_path, generate_kwargs={"language": client.config['language']})['text']
|
||||
else:
|
||||
to_return = self.asr_pipeline(file_path)['text']
|
||||
|
||||
os.remove(file_path)
|
||||
|
||||
to_return = {
|
||||
"language": "UNSUPPORTED_BY_HUGGINGFACE_WHISPER",
|
||||
"language_probability": None,
|
||||
"text": to_return.strip(),
|
||||
"words": "UNSUPPORTED_BY_HUGGINGFACE_WHISPER"
|
||||
}
|
||||
return to_return
|
||||
@@ -0,0 +1,26 @@
|
||||
import wave
|
||||
import os
|
||||
|
||||
async def save_audio_to_file(audio_data, file_name, audio_dir="audio_files", audio_format="wav"):
|
||||
"""
|
||||
Saves the audio data to a file.
|
||||
|
||||
:param client_id: Unique identifier for the client.
|
||||
:param audio_data: The audio data to save.
|
||||
:param file_counters: Dictionary to keep track of file counts for each client.
|
||||
:param audio_dir: Directory where audio files will be saved.
|
||||
:param audio_format: Format of the audio file.
|
||||
:return: Path to the saved audio file.
|
||||
"""
|
||||
|
||||
os.makedirs(audio_dir, exist_ok=True)
|
||||
|
||||
file_path = os.path.join(audio_dir, file_name)
|
||||
|
||||
with wave.open(file_path, 'wb') as wav_file:
|
||||
wav_file.setnchannels(1) # Assuming mono audio
|
||||
wav_file.setsampwidth(2)
|
||||
wav_file.setframerate(16000)
|
||||
wav_file.writeframes(audio_data)
|
||||
|
||||
return file_path
|
||||
@@ -0,0 +1,142 @@
|
||||
import os
|
||||
import asyncio
|
||||
import json
|
||||
import time
|
||||
|
||||
from VoiceStreamAI.buffering_strategy.buffering_strategy_interface import BufferingStrategyInterface
|
||||
from openai import OpenAI
|
||||
|
||||
class SilenceAtEndOfChunk(BufferingStrategyInterface):
|
||||
"""
|
||||
A buffering strategy that processes audio at the end of each chunk with silence detection.
|
||||
|
||||
This class is responsible for handling audio chunks, detecting silence at the end of each chunk,
|
||||
and initiating the transcription process for the chunk.
|
||||
|
||||
Attributes:
|
||||
client (Client): The client instance associated with this buffering strategy.
|
||||
chunk_length_seconds (float): Length of each audio chunk in seconds.
|
||||
chunk_offset_seconds (float): Offset time in seconds to be considered for processing audio chunks.
|
||||
"""
|
||||
|
||||
def __init__(self, client, **kwargs):
|
||||
"""
|
||||
Initialize the SilenceAtEndOfChunk buffering strategy.
|
||||
|
||||
Args:
|
||||
client (Client): The client instance associated with this buffering strategy.
|
||||
**kwargs: Additional keyword arguments, including 'chunk_length_seconds' and 'chunk_offset_seconds'.
|
||||
"""
|
||||
self.client = client
|
||||
|
||||
self.chunk_length_seconds = os.environ.get('BUFFERING_CHUNK_LENGTH_SECONDS')
|
||||
if not self.chunk_length_seconds:
|
||||
self.chunk_length_seconds = kwargs.get('chunk_length_seconds')
|
||||
self.chunk_length_seconds = float(self.chunk_length_seconds)
|
||||
|
||||
self.chunk_offset_seconds = os.environ.get('BUFFERING_CHUNK_OFFSET_SECONDS')
|
||||
if not self.chunk_offset_seconds:
|
||||
self.chunk_offset_seconds = kwargs.get('chunk_offset_seconds')
|
||||
self.chunk_offset_seconds = float(self.chunk_offset_seconds)
|
||||
|
||||
self.error_if_not_realtime = os.environ.get('ERROR_IF_NOT_REALTIME')
|
||||
if not self.error_if_not_realtime:
|
||||
self.error_if_not_realtime = kwargs.get('error_if_not_realtime', False)
|
||||
|
||||
self.processing_flag = False
|
||||
|
||||
self.messages=[]
|
||||
|
||||
def process_audio(self, websocket, vad_pipeline, asr_pipeline,llm_port):
|
||||
"""
|
||||
Process audio chunks by checking their length and scheduling asynchronous processing.
|
||||
|
||||
This method checks if the length of the audio buffer exceeds the chunk length and, if so,
|
||||
it schedules asynchronous processing of the audio.
|
||||
|
||||
Args:
|
||||
websocket (Websocket): The WebSocket connection for sending transcriptions.
|
||||
vad_pipeline: The voice activity detection pipeline.
|
||||
asr_pipeline: The automatic speech recognition pipeline.
|
||||
"""
|
||||
chunk_length_in_bytes = self.chunk_length_seconds * self.client.sampling_rate * self.client.samples_width
|
||||
if len(self.client.buffer) > chunk_length_in_bytes:
|
||||
if self.processing_flag:
|
||||
exit("Error in realtime processing: tried processing a new chunk while the previous one was still being processed")
|
||||
|
||||
self.client.scratch_buffer += self.client.buffer
|
||||
self.client.buffer.clear()
|
||||
self.processing_flag = True
|
||||
# Schedule the processing in a separate task
|
||||
asyncio.create_task(self.process_audio_async(websocket, vad_pipeline, asr_pipeline,llm_port))
|
||||
|
||||
async def process_audio_async(self, websocket, vad_pipeline, asr_pipeline,llm_port):
|
||||
"""
|
||||
Asynchronously process audio for activity detection and transcription.
|
||||
|
||||
This method performs heavy processing, including voice activity detection and transcription of
|
||||
the audio data. It sends the transcription results through the WebSocket connection.
|
||||
|
||||
Args:
|
||||
websocket (Websocket): The WebSocket connection for sending transcriptions.
|
||||
vad_pipeline: The voice activity detection pipeline.
|
||||
asr_pipeline: The automatic speech recognition pipeline.
|
||||
"""
|
||||
start = time.time()
|
||||
vad_results = await vad_pipeline.detect_activity(self.client)
|
||||
|
||||
if len(vad_results) == 0:
|
||||
self.client.scratch_buffer.clear()
|
||||
self.client.buffer.clear()
|
||||
self.processing_flag = False
|
||||
return
|
||||
|
||||
last_segment_should_end_before = ((len(self.client.scratch_buffer) / (self.client.sampling_rate * self.client.samples_width)) - self.chunk_offset_seconds)
|
||||
if vad_results[-1]['end'] < last_segment_should_end_before:
|
||||
transcription = await asr_pipeline.transcribe(self.client)
|
||||
if transcription['text'] != '':
|
||||
end = time.time()
|
||||
transcription['processing_time'] = end - start
|
||||
|
||||
transcription['status']="chat_start"
|
||||
|
||||
json_transcription = json.dumps(transcription)
|
||||
|
||||
await websocket.send(json_transcription)
|
||||
|
||||
# Point to the local server
|
||||
client = OpenAI(base_url=f"http://localhost:{llm_port}/v1", api_key="lm-studio")
|
||||
|
||||
|
||||
messages=[
|
||||
{"role": "system", "content": "You are a friendly and engaging AI designed to interact with users in a conversational manner. Your personality is that of a sophisticated and polite young professional who is both a designer and a programmer. You are well-mannered, articulate, and possess a good sense of humor. Your goal is to provide helpful and insightful responses while maintaining a pleasant and enjoyable conversation. Be sure to use your knowledge in design and programming to enrich the dialogue and offer relevant advice or information when appropriate. Always be respectful and considerate of the user's feelings and perspectives. Additionally, you are fluent in both English and Chinese, and can seamlessly switch between the two languages to best assist users."},
|
||||
]+self.messages[-10:0]+[{"role": "user", "content":transcription['text']}]
|
||||
# print('#messages',messages)
|
||||
|
||||
completion = client.chat.completions.create(
|
||||
model="model-identifier",
|
||||
messages=messages,
|
||||
temperature=0.7,
|
||||
)
|
||||
|
||||
transcription['asistant'] = completion.choices[0].message.content
|
||||
|
||||
transcription['status']="chat_end"
|
||||
|
||||
json_transcription = json.dumps(transcription)
|
||||
|
||||
self.messages.append({
|
||||
"role": "user",
|
||||
"content":transcription['text']})
|
||||
|
||||
self.messages.append({
|
||||
"role": "asistant",
|
||||
"content": transcription['asistant']
|
||||
})
|
||||
# print('#messages',completion.choices[0].message.content)
|
||||
|
||||
await websocket.send(json_transcription)
|
||||
self.client.scratch_buffer.clear()
|
||||
self.client.increment_file_counter()
|
||||
|
||||
self.processing_flag = False
|
||||
@@ -0,0 +1,41 @@
|
||||
from VoiceStreamAI.buffering_strategy.buffering_strategies import SilenceAtEndOfChunk
|
||||
|
||||
class BufferingStrategyFactory:
|
||||
"""
|
||||
A factory class for creating instances of different buffering strategies.
|
||||
|
||||
This factory provides a centralized way to instantiate various buffering strategies
|
||||
based on the type specified. It abstracts the creation logic, making it easier to
|
||||
manage and extend with new buffering strategy types.
|
||||
|
||||
Methods:
|
||||
create_buffering_strategy: Creates and returns an instance of a specified buffering strategy.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def create_buffering_strategy(type, client, **kwargs):
|
||||
"""
|
||||
Creates an instance of a buffering strategy based on the specified type.
|
||||
|
||||
This method acts as a factory for creating buffering strategy objects. It returns
|
||||
an instance of the strategy corresponding to the given type. If the type is not
|
||||
recognized, it raises a ValueError.
|
||||
|
||||
Args:
|
||||
type (str): The type of buffering strategy to create. Currently supports 'silence_at_end_of_chunk'.
|
||||
client (Client): The client instance to be associated with the buffering strategy.
|
||||
**kwargs: Additional keyword arguments specific to the buffering strategy being created.
|
||||
|
||||
Returns:
|
||||
An instance of the specified buffering strategy.
|
||||
|
||||
Raises:
|
||||
ValueError: If the specified type is not recognized or supported.
|
||||
|
||||
Example:
|
||||
strategy = BufferingStrategyFactory.create_buffering_strategy("silence_at_end_of_chunk", client)
|
||||
"""
|
||||
if type == "silence_at_end_of_chunk":
|
||||
return SilenceAtEndOfChunk(client, **kwargs)
|
||||
else:
|
||||
raise ValueError(f"Unknown buffering strategy type: {type}")
|
||||
@@ -0,0 +1,31 @@
|
||||
class BufferingStrategyInterface:
|
||||
"""
|
||||
An interface class for buffering strategies in audio processing systems.
|
||||
|
||||
This class defines the structure for buffering strategies used in handling
|
||||
and processing audio data. It serves as a template for creating custom buffering
|
||||
strategies that fit specific requirements of an audio processing pipeline.
|
||||
|
||||
Subclasses should implement the methods defined in this interface to ensure
|
||||
consistency and compatibility with the system's audio processing framework.
|
||||
|
||||
Methods:
|
||||
process_audio: Process audio data. This method should be implemented by subclasses.
|
||||
"""
|
||||
|
||||
def process_audio(self, websocket, vad_pipeline, asr_pipeline):
|
||||
"""
|
||||
Process audio data using the given WebSocket connection, VAD pipeline, and ASR pipeline.
|
||||
|
||||
This method is intended to be overridden in subclasses to provide specific logic
|
||||
for handling and processing audio data in different buffering strategies.
|
||||
|
||||
Args:
|
||||
websocket (Websocket): The WebSocket connection for communication with clients.
|
||||
vad_pipeline: The Voice Activity Detection (VAD) pipeline used for detecting speech in the audio.
|
||||
asr_pipeline: The Automatic Speech Recognition (ASR) pipeline used for transcribing speech in the audio.
|
||||
|
||||
Raises:
|
||||
NotImplementedError: If the method is not implemented in the subclass.
|
||||
"""
|
||||
raise NotImplementedError("This method should be implemented by subclasses.")
|
||||
@@ -0,0 +1,54 @@
|
||||
from VoiceStreamAI.buffering_strategy.buffering_strategy_factory import BufferingStrategyFactory
|
||||
|
||||
class Client:
|
||||
"""
|
||||
Represents a client connected to the VoiceStreamAI server.
|
||||
|
||||
This class maintains the state for each connected client, including their
|
||||
unique identifier, audio buffer, configuration, and a counter for processed audio files.
|
||||
|
||||
Attributes:
|
||||
client_id (str): A unique identifier for the client.
|
||||
buffer (bytearray): A buffer to store incoming audio data.
|
||||
config (dict): Configuration settings for the client, like chunk length and offset.
|
||||
file_counter (int): Counter for the number of audio files processed.
|
||||
total_samples (int): Total number of audio samples received from this client.
|
||||
sampling_rate (int): The sampling rate of the audio data in Hz.
|
||||
samples_width (int): The width of each audio sample in bits.
|
||||
"""
|
||||
def __init__(self, client_id, sampling_rate, samples_width):
|
||||
self.client_id = client_id
|
||||
self.buffer = bytearray()
|
||||
self.scratch_buffer = bytearray()
|
||||
self.config = {"language": None,
|
||||
"processing_strategy": "silence_at_end_of_chunk",
|
||||
"processing_args": {
|
||||
"chunk_length_seconds": 5,
|
||||
"chunk_offset_seconds": 0.1
|
||||
}
|
||||
}
|
||||
self.file_counter = 0
|
||||
self.total_samples = 0
|
||||
self.sampling_rate = sampling_rate
|
||||
self.samples_width = samples_width
|
||||
self.buffering_strategy = BufferingStrategyFactory.create_buffering_strategy(self.config['processing_strategy'], self, **self.config['processing_args'])
|
||||
|
||||
def update_config(self, config_data):
|
||||
self.config.update(config_data)
|
||||
self.buffering_strategy = BufferingStrategyFactory.create_buffering_strategy(self.config['processing_strategy'], self, **self.config['processing_args'])
|
||||
|
||||
def append_audio_data(self, audio_data):
|
||||
self.buffer.extend(audio_data)
|
||||
self.total_samples += len(audio_data) / self.samples_width
|
||||
|
||||
def clear_buffer(self):
|
||||
self.buffer.clear()
|
||||
|
||||
def increment_file_counter(self):
|
||||
self.file_counter += 1
|
||||
|
||||
def get_file_name(self):
|
||||
return f"{self.client_id}_{self.file_counter}.wav"
|
||||
|
||||
def process_audio(self, websocket, vad_pipeline, asr_pipeline,llm_port):
|
||||
self.buffering_strategy.process_audio(websocket, vad_pipeline, asr_pipeline,llm_port)
|
||||
@@ -0,0 +1,54 @@
|
||||
import argparse
|
||||
import asyncio
|
||||
import json
|
||||
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
# 获取当前文件的绝对路径
|
||||
current_file_path = os.path.abspath(__file__)
|
||||
|
||||
# 获取当前文件的目录
|
||||
current_directory = os.path.dirname(current_file_path)
|
||||
sys.path.append(str(Path(current_directory).parent))
|
||||
# print("sys.path", current_directory)
|
||||
|
||||
|
||||
from VoiceStreamAI.server import Server
|
||||
from VoiceStreamAI.asr.asr_factory import ASRFactory
|
||||
from VoiceStreamAI.vad.vad_factory import VADFactory
|
||||
|
||||
def parse_args():
|
||||
parser = argparse.ArgumentParser(description="VoiceStreamAI Server: Real-time audio transcription using self-hosted Whisper and WebSocket")
|
||||
parser.add_argument("--vad-type", type=str, default="pyannote", help="Type of VAD pipeline to use (e.g., 'pyannote')")
|
||||
parser.add_argument("--vad-args", type=str, default='{"auth_token": "huggingface_token"}', help="JSON string of additional arguments for VAD pipeline")
|
||||
parser.add_argument("--asr-type", type=str, default="faster_whisper", help="Type of ASR pipeline to use (e.g., 'whisper')")
|
||||
parser.add_argument("--asr-args", type=str, default='{"model_size": "large-v3"}', help="JSON string of additional arguments for ASR pipeline")
|
||||
parser.add_argument("--host", type=str, default="127.0.0.1", help="Host for the WebSocket server")
|
||||
parser.add_argument("--port", type=int, default=8765, help="Port for the WebSocket server")
|
||||
parser.add_argument("--certfile", type=str, default=None, help="The path to the SSL certificate (cert file) if using secure websockets")
|
||||
parser.add_argument("--keyfile", type=str, default=None, help="The path to the SSL key file if using secure websockets")
|
||||
return parser.parse_args()
|
||||
|
||||
def main():
|
||||
args = parse_args()
|
||||
|
||||
try:
|
||||
vad_args = json.loads(args.vad_args)
|
||||
asr_args = json.loads(args.asr_args)
|
||||
except json.JSONDecodeError as e:
|
||||
print(f"Error parsing JSON arguments: {e}")
|
||||
return
|
||||
|
||||
vad_pipeline = VADFactory.create_vad_pipeline(args.vad_type, **vad_args)
|
||||
asr_pipeline = ASRFactory.create_asr_pipeline(args.asr_type, **asr_args)
|
||||
|
||||
server = Server(vad_pipeline, asr_pipeline, host=args.host, port=args.port, sampling_rate=16000, samples_width=2, certfile=args.certfile, keyfile=args.keyfile)
|
||||
|
||||
asyncio.get_event_loop().run_until_complete(server.start())
|
||||
asyncio.get_event_loop().run_forever()
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,7 @@
|
||||
websockets
|
||||
speechbrain
|
||||
pyannote-audio
|
||||
asyncio
|
||||
sentence-transformers
|
||||
transformers
|
||||
faster-whisper
|
||||
@@ -0,0 +1,88 @@
|
||||
import websockets
|
||||
import uuid
|
||||
import json
|
||||
import asyncio
|
||||
import ssl
|
||||
|
||||
from VoiceStreamAI.audio_utils import save_audio_to_file
|
||||
from VoiceStreamAI.client import Client
|
||||
|
||||
class Server:
|
||||
"""
|
||||
Represents the WebSocket server for handling real-time audio transcription.
|
||||
|
||||
This class manages WebSocket connections, processes incoming audio data,
|
||||
and interacts with VAD and ASR pipelines for voice activity detection and
|
||||
speech recognition.
|
||||
|
||||
Attributes:
|
||||
vad_pipeline: An instance of a voice activity detection pipeline.
|
||||
asr_pipeline: An instance of an automatic speech recognition pipeline.
|
||||
host (str): Host address of the server.
|
||||
port (int): Port on which the server listens.
|
||||
sampling_rate (int): The sampling rate of audio data in Hz.
|
||||
samples_width (int): The width of each audio sample in bits.
|
||||
connected_clients (dict): A dictionary mapping client IDs to Client objects.
|
||||
"""
|
||||
def __init__(self, vad_pipeline, asr_pipeline, host='localhost', port=8765, sampling_rate=16000, samples_width=2, certfile = None, keyfile = None,llm_port=9000):
|
||||
self.vad_pipeline = vad_pipeline
|
||||
self.asr_pipeline = asr_pipeline
|
||||
self.host = host
|
||||
self.port = port
|
||||
self.sampling_rate = sampling_rate
|
||||
self.samples_width = samples_width
|
||||
self.certfile = certfile
|
||||
self.keyfile = keyfile
|
||||
self.connected_clients = {}
|
||||
|
||||
self.llm_port=llm_port
|
||||
|
||||
async def handle_audio(self, client, websocket):
|
||||
while True:
|
||||
message = await websocket.recv()
|
||||
|
||||
if isinstance(message, bytes):
|
||||
client.append_audio_data(message)
|
||||
elif isinstance(message, str):
|
||||
config = json.loads(message)
|
||||
if config.get('type') == 'config':
|
||||
client.update_config(config['data'])
|
||||
continue
|
||||
else:
|
||||
print(f"Unexpected message type from {client.client_id}")
|
||||
|
||||
# this is synchronous, any async operation is in BufferingStrategy
|
||||
client.process_audio(websocket, self.vad_pipeline, self.asr_pipeline,self.llm_port)
|
||||
|
||||
|
||||
async def handle_websocket(self, websocket, path):
|
||||
client_id = str(uuid.uuid4())
|
||||
client = Client(client_id, self.sampling_rate, self.samples_width)
|
||||
self.connected_clients[client_id] = client
|
||||
|
||||
print(f"Client {client_id} connected")
|
||||
|
||||
try:
|
||||
await self.handle_audio(client, websocket)
|
||||
except websockets.ConnectionClosed as e:
|
||||
print(f"Connection with {client_id} closed: {e}")
|
||||
finally:
|
||||
del self.connected_clients[client_id]
|
||||
|
||||
def start(self):
|
||||
if self.certfile:
|
||||
# Create an SSL context to enforce encrypted connections
|
||||
ssl_context = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER)
|
||||
|
||||
# Load your server's certificate and private key
|
||||
# Replace 'your_cert_path.pem' and 'your_key_path.pem' with the actual paths to your files
|
||||
ssl_context.load_cert_chain(certfile=self.certfile, keyfile=self.keyfile)
|
||||
|
||||
print(f"WebSocket server ready to accept secure connections on {self.host}:{self.port}")
|
||||
|
||||
# Pass the SSL context to the serve function along with the host and port
|
||||
# Ensure the secure flag is set to True if using a secure WebSocket protocol (wss://)
|
||||
return websockets.serve(self.handle_websocket, self.host, self.port, ssl=ssl_context)
|
||||
else:
|
||||
print(f"WebSocket server ready to accept secure connections on {self.host}:{self.port}")
|
||||
return websockets.serve(self.handle_websocket, self.host, self.port)
|
||||
@@ -0,0 +1,50 @@
|
||||
from os import remove
|
||||
import os
|
||||
|
||||
from pyannote.core import Segment
|
||||
from pyannote.audio import Model
|
||||
from pyannote.audio.pipelines import VoiceActivityDetection
|
||||
|
||||
from VoiceStreamAI.vad.vad_interface import VADInterface
|
||||
from VoiceStreamAI.audio_utils import save_audio_to_file
|
||||
|
||||
|
||||
class PyannoteVAD(VADInterface):
|
||||
"""
|
||||
Pyannote-based implementation of the VADInterface.
|
||||
"""
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
"""
|
||||
Initializes Pyannote's VAD pipeline.
|
||||
|
||||
Args:
|
||||
model_name (str): The model name for Pyannote.
|
||||
auth_token (str, optional): Authentication token for Hugging Face.
|
||||
"""
|
||||
|
||||
model_name = kwargs.get('model_name', "pyannote/segmentation")
|
||||
|
||||
auth_token = os.environ.get('PYANNOTE_AUTH_TOKEN')
|
||||
if not auth_token:
|
||||
auth_token = kwargs.get('auth_token')
|
||||
|
||||
if auth_token is None:
|
||||
raise ValueError("Missing required env var in PYANNOTE_AUTH_TOKEN or argument in --vad-args: 'auth_token'")
|
||||
|
||||
pyannote_args = kwargs.get('pyannote_args', {"onset": 0.5, "offset": 0.5, "min_duration_on": 0.3, "min_duration_off": 0.3})
|
||||
self.model = Model.from_pretrained(model_name, use_auth_token=auth_token)
|
||||
self.vad_pipeline = VoiceActivityDetection(segmentation=self.model)
|
||||
self.vad_pipeline.instantiate(pyannote_args)
|
||||
|
||||
async def detect_activity(self, client):
|
||||
audio_file_path = await save_audio_to_file(client.scratch_buffer, client.get_file_name())
|
||||
vad_results = self.vad_pipeline(audio_file_path)
|
||||
remove(audio_file_path)
|
||||
vad_segments = []
|
||||
if len(vad_results) > 0:
|
||||
vad_segments = [
|
||||
{"start": segment.start, "end": segment.end, "confidence": 1.0}
|
||||
for segment in vad_results.itersegments()
|
||||
]
|
||||
return vad_segments
|
||||
@@ -0,0 +1,23 @@
|
||||
from VoiceStreamAI.vad.pyannote_vad import PyannoteVAD
|
||||
|
||||
class VADFactory:
|
||||
"""
|
||||
Factory for creating instances of VAD systems.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def create_vad_pipeline(type, **kwargs):
|
||||
"""
|
||||
Creates a VAD pipeline based on the specified type.
|
||||
|
||||
Args:
|
||||
type (str): The type of VAD pipeline to create (e.g., 'pyannote').
|
||||
kwargs: Additional arguments for the VAD pipeline creation.
|
||||
|
||||
Returns:
|
||||
VADInterface: An instance of a class that implements VADInterface.
|
||||
"""
|
||||
if type == "pyannote":
|
||||
return PyannoteVAD(**kwargs)
|
||||
else:
|
||||
raise ValueError(f"Unknown VAD pipeline type: {type}")
|
||||
@@ -0,0 +1,16 @@
|
||||
class VADInterface:
|
||||
"""
|
||||
Interface for voice activity detection (VAD) systems.
|
||||
"""
|
||||
|
||||
async def detect_activity(self, client):
|
||||
"""
|
||||
Detects voice activity in the given audio data.
|
||||
|
||||
Args:
|
||||
client (src.Client): The client to detect on
|
||||
|
||||
Returns:
|
||||
List: VAD result, a list of objects containing "start", "end", "confidence"
|
||||
"""
|
||||
raise NotImplementedError("This method should be implemented by subclasses.")
|
||||
+525
@@ -0,0 +1,525 @@
|
||||
<!DOCTYPE html>
|
||||
<!--
|
||||
VoiceStreamAI Client Interface
|
||||
Real-time audio transcription using self-hosted Whisper and WebSocket
|
||||
|
||||
Contributor:
|
||||
- Alessandro Saccoia - alessandro.saccoia@gmail.com
|
||||
-->
|
||||
<html lang="en">
|
||||
|
||||
<head>
|
||||
<meta charset="UTF-8">
|
||||
<title>Audio Stream to WebSocket Server</title>
|
||||
<style>
|
||||
body {
|
||||
font-family: Arial, sans-serif;
|
||||
margin: 0;
|
||||
padding: 0;
|
||||
background: #f4f4f4;
|
||||
text-align: center;
|
||||
}
|
||||
|
||||
h1 {
|
||||
color: #333;
|
||||
}
|
||||
|
||||
.controls {
|
||||
margin: 20px auto;
|
||||
padding: 10px;
|
||||
width: 80%;
|
||||
display: flex;
|
||||
justify-content: space-around;
|
||||
align-items: center;
|
||||
}
|
||||
|
||||
.control-group {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
align-items: center;
|
||||
}
|
||||
|
||||
.controls input,
|
||||
.controls button,
|
||||
.controls select {
|
||||
padding: 8px;
|
||||
margin: 5px;
|
||||
border: 1px solid #ddd;
|
||||
border-radius: 5px;
|
||||
font-size: 0.9em;
|
||||
}
|
||||
|
||||
#transcription {
|
||||
margin: 20px auto;
|
||||
border: 1px solid #ddd;
|
||||
padding: 10px;
|
||||
width: 80%;
|
||||
height: 150px;
|
||||
overflow-y: auto;
|
||||
background: white;
|
||||
}
|
||||
|
||||
.label {
|
||||
font-size: 0.9em;
|
||||
color: #555;
|
||||
margin-bottom: 5px;
|
||||
}
|
||||
|
||||
button {
|
||||
cursor: pointer;
|
||||
}
|
||||
|
||||
.buffering-strategy-panel {
|
||||
margin-top: 10px;
|
||||
}
|
||||
|
||||
/* ... existing styles ... */
|
||||
.hidden {
|
||||
display: none;
|
||||
}
|
||||
</style>
|
||||
|
||||
<style>
|
||||
body {
|
||||
margin: 0;
|
||||
padding: 0;
|
||||
/* background-color: #333;
|
||||
color:white */
|
||||
}
|
||||
|
||||
#mic_container {
|
||||
display: flex;
|
||||
width: 100%;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
}
|
||||
|
||||
#mic {
|
||||
border: 1px solid #ddd;
|
||||
border-radius: 4px;
|
||||
margin-top: 1rem;
|
||||
|
||||
width: 300px
|
||||
}
|
||||
|
||||
#asistant {
|
||||
width: 300px;
|
||||
padding: 12px;
|
||||
font-size: 12px;
|
||||
overflow-y: scroll;
|
||||
height: 300px;
|
||||
}
|
||||
</style>
|
||||
</head>
|
||||
|
||||
<body>
|
||||
|
||||
<div id="mic_container">
|
||||
<div id="mic"></div>
|
||||
<div id="asistant"></div>
|
||||
</div>
|
||||
|
||||
<h1>VAD + Whisper + </h1>
|
||||
|
||||
<button id="init_server">Server</button>
|
||||
|
||||
<div class="controls">
|
||||
<div class="control-group">
|
||||
<label class="label" for="websocketAddress">WebSocket Address:</label>
|
||||
<input type="text" id="websocketAddress" value="ws://localhost:8725">
|
||||
</div>
|
||||
<div class="control-group">
|
||||
<label class="label" for="bufferingStrategySelect" onchange="toggleBufferingStrategyPanel()">Buffering
|
||||
Strategy:</label>
|
||||
<select id="bufferingStrategySelect">
|
||||
<option value="silence_at_end_of_chunk" selected>Silence at End of Chunk</option>
|
||||
</select>
|
||||
</div>
|
||||
<div class="silence_at_end_of_chunk_options_panel">
|
||||
<div class="control-group">
|
||||
<label class="label" for="chunk_length_seconds">Chunk Length (s):</label>
|
||||
<input type="number" id="chunk_length_seconds" value="3" min="1">
|
||||
</div>
|
||||
<div class="control-group">
|
||||
<label class="label" for="chunk_offset_seconds">Silence at the End of Chunk (s):</label>
|
||||
<input type="number" id="chunk_offset_seconds" value="0.1" min="0">
|
||||
</div>
|
||||
</div>
|
||||
<div class="control-group">
|
||||
<label class="label" for="languageSelect">Language:</label>
|
||||
<select id="languageSelect">
|
||||
<option value="multilingual">Multilingual</option>
|
||||
<option value="english">English</option>
|
||||
<option value="italian">Italian</option>
|
||||
<option value="spanish">Spanish</option>
|
||||
<option value="french">French</option>
|
||||
<option value="german">German</option>
|
||||
<option value="chinese">Chinese</option>
|
||||
<option value="arabic">Arabic</option>
|
||||
<option value="portuguese">Portuguese</option>
|
||||
<option value="russian">Russian</option>
|
||||
<option value="japanese">Japanese</option>
|
||||
<option value="dutch">Dutch</option>
|
||||
<option value="korean">Korean</option>
|
||||
<option value="hindi">Hindi</option>
|
||||
<option value="turkish">Turkish</option>
|
||||
<option value="swedish">Swedish</option>
|
||||
<option value="norwegian">Norwegian</option>
|
||||
<option value="danish">Danish</option>
|
||||
<option value="polish">Polish</option>
|
||||
<option value="finnish">Finnish</option>
|
||||
<option value="thai">Thai</option>
|
||||
<option value="czech">Czech</option>
|
||||
<option value="hungarian">Hungarian</option>
|
||||
<option value="greek">Greek</option>
|
||||
</select>
|
||||
</div>
|
||||
<button id="connectButton">Connect</button>
|
||||
</div>
|
||||
<button id="startButton" disabled>Start Streaming</button>
|
||||
<button id="stopButton" disabled>Stop Streaming</button>
|
||||
<div id="transcription"></div>
|
||||
<br />
|
||||
<div>WebSocket: <span id="webSocketStatus">Not Connected</span></div>
|
||||
<div>Detected Language: <span id="detected_language">Undefined</span></div>
|
||||
<div>Last Processing Time: <span id="processing_time">Undefined</span></div>
|
||||
|
||||
<script type="module">
|
||||
|
||||
// Record plugin
|
||||
|
||||
import WaveSurfer from 'https://cdn.jsdelivr.net/npm/wavesurfer.js@7/dist/wavesurfer.esm.js'
|
||||
import RecordPlugin from 'https://cdn.jsdelivr.net/npm/wavesurfer.js/dist/plugins/record.esm.js'
|
||||
|
||||
// 可视化
|
||||
let wavesurfer, record;
|
||||
|
||||
// ws服务
|
||||
let websocket;
|
||||
let context;
|
||||
let processor;
|
||||
let globalStream;
|
||||
|
||||
const websocket_uri = 'ws://localhost:8765';
|
||||
const bufferSize = 4096;
|
||||
let isRecording = false;
|
||||
let chunk_length_seconds, chunk_offset_seconds, language;
|
||||
|
||||
// Record button
|
||||
const startButton = document.getElementById('startButton'),
|
||||
stopButton = document.getElementById('stopButton'),
|
||||
connectButton = document.getElementById('connectButton'),
|
||||
initServerButton = document.getElementById('init_server')
|
||||
|
||||
const createWaveSurfer = () => {
|
||||
// Create an instance of WaveSurfer
|
||||
if (wavesurfer) {
|
||||
wavesurfer.destroy()
|
||||
}
|
||||
wavesurfer = WaveSurfer.create({
|
||||
container: '#mic',
|
||||
waveColor: 'rgb(200, 0, 200)',
|
||||
progressColor: 'rgb(100, 0, 100)',
|
||||
renderFunction: (channels, ctx) => {
|
||||
const { width, height } = ctx.canvas
|
||||
// console.log(width, height)
|
||||
const scale = channels[0].length / width
|
||||
const step = 20
|
||||
|
||||
ctx.translate(0, height / 2)
|
||||
ctx.strokeStyle = ctx.fillStyle
|
||||
ctx.beginPath()
|
||||
|
||||
for (let i = 0; i < width; i += step * 2) {
|
||||
const index = Math.floor(i * scale)
|
||||
const value = Math.abs(channels[0][index])
|
||||
let x = i
|
||||
let y = value * height * 1.2
|
||||
|
||||
ctx.moveTo(x, 0)
|
||||
ctx.lineTo(x, y)
|
||||
ctx.arc(x + step / 2, y, step / 2, Math.PI, 0, true)
|
||||
ctx.lineTo(x + step, 0)
|
||||
|
||||
x = x + step
|
||||
y = -y
|
||||
ctx.moveTo(x, 0)
|
||||
ctx.lineTo(x, y)
|
||||
ctx.arc(x + step / 2, y, step / 2, Math.PI, 0, false)
|
||||
ctx.lineTo(x + step, 0)
|
||||
}
|
||||
|
||||
ctx.stroke()
|
||||
ctx.closePath()
|
||||
},
|
||||
// // Set a bar width
|
||||
// barWidth: 20,
|
||||
// // Optionally, specify the spacing between bars
|
||||
// barGap: 4,
|
||||
// // And the bar radius
|
||||
// barRadius: 2,
|
||||
})
|
||||
|
||||
// Initialize the Record plugin
|
||||
record = wavesurfer.registerPlugin(RecordPlugin.create({ scrollingWaveform: false, renderRecordedAudio: false }))
|
||||
// Render recorded audio
|
||||
// recButton.textContent = 'Record'
|
||||
|
||||
}
|
||||
|
||||
startButton.addEventListener('click', e => {
|
||||
record.startRecording()
|
||||
|
||||
if (isRecording) return;
|
||||
isRecording = true;
|
||||
|
||||
const AudioContext = window.AudioContext || window.webkitAudioContext;
|
||||
context = new AudioContext();
|
||||
|
||||
navigator.mediaDevices.getUserMedia({ audio: true }).then(stream => {
|
||||
globalStream = stream;
|
||||
const input = context.createMediaStreamSource(stream);
|
||||
processor = context.createScriptProcessor(bufferSize, 1, 1);
|
||||
processor.onaudioprocess = e => processAudio(e);
|
||||
input.connect(processor);
|
||||
processor.connect(context.destination);
|
||||
|
||||
sendAudioConfig();
|
||||
}).catch(error => console.error('Error accessing microphone', error));
|
||||
|
||||
// Disable start button and enable stop button
|
||||
startButton.disabled = true;
|
||||
stopButton.disabled = false;
|
||||
|
||||
});
|
||||
|
||||
stopButton.addEventListener('click', e => {
|
||||
if (!isRecording) return;
|
||||
isRecording = false;
|
||||
|
||||
if (globalStream) {
|
||||
globalStream.getTracks().forEach(track => track.stop());
|
||||
}
|
||||
if (processor) {
|
||||
processor.disconnect();
|
||||
processor = null;
|
||||
}
|
||||
if (context) {
|
||||
context.close().then(() => context = null);
|
||||
}
|
||||
startButton.disabled = false;
|
||||
stopButton.disabled = true;
|
||||
|
||||
// 可视化部分
|
||||
if (record.isRecording() || record.isPaused()) {
|
||||
record.stopRecording()
|
||||
return
|
||||
}
|
||||
})
|
||||
|
||||
connectButton.addEventListener('click', e => {
|
||||
initWebSocket()
|
||||
})
|
||||
|
||||
initServerButton.addEventListener('click', async e => {
|
||||
const response = await fetch('/mixlab/start_live', {
|
||||
method: 'POST',
|
||||
headers: {
|
||||
'Content-Type': 'application/json'
|
||||
},
|
||||
body: JSON.stringify({
|
||||
port: 8323,
|
||||
model: ''
|
||||
})
|
||||
})
|
||||
console.log(await response.json())
|
||||
|
||||
})
|
||||
|
||||
// Initialize on page load
|
||||
window.onload = () => {
|
||||
initWebSocket()
|
||||
createWaveSurfer();
|
||||
};
|
||||
|
||||
function initWebSocket() {
|
||||
const websocketAddress = document.getElementById('websocketAddress').value;
|
||||
chunk_length_seconds = document.getElementById('chunk_length_seconds').value;
|
||||
chunk_offset_seconds = document.getElementById('chunk_offset_seconds').value;
|
||||
const selectedLanguage = document.getElementById('languageSelect').value;
|
||||
language = selectedLanguage !== 'multilingual' ? selectedLanguage : null;
|
||||
|
||||
if (!websocketAddress) {
|
||||
console.log("WebSocket address is required.");
|
||||
return;
|
||||
}
|
||||
|
||||
if (websocket) websocket.close()
|
||||
|
||||
websocket = new WebSocket(websocketAddress);
|
||||
websocket.onopen = () => {
|
||||
console.log("WebSocket connection established");
|
||||
document.getElementById("webSocketStatus").textContent = 'Connected';
|
||||
startButton.disabled = false;
|
||||
};
|
||||
websocket.onclose = event => {
|
||||
console.log("WebSocket connection closed", event);
|
||||
document.getElementById("webSocketStatus").textContent = 'Not Connected';
|
||||
stopButton.click()
|
||||
startButton.disabled = true;
|
||||
stopButton.disabled = true;
|
||||
|
||||
// setTimeout(()=>initWebSocket(),1000)
|
||||
};
|
||||
websocket.onmessage = event => {
|
||||
console.log("Message from server:", event.data);
|
||||
const transcript_data = JSON.parse(event.data);
|
||||
|
||||
if (transcript_data.status === 'chat_start') {
|
||||
updateTranscription(transcript_data);
|
||||
stopButton.click();
|
||||
} else if (transcript_data.status === 'chat_end') {
|
||||
let asistant = decodeURIComponent(transcript_data.asistant)
|
||||
document.getElementById('asistant').innerText = asistant;
|
||||
startButton.click();
|
||||
}
|
||||
|
||||
};
|
||||
websocket.onerror = () => {
|
||||
// setTimeout(()=>initWebSocket(),1000)
|
||||
}
|
||||
}
|
||||
|
||||
function updateTranscription(transcript_data) {
|
||||
const transcriptionDiv = document.getElementById('transcription');
|
||||
const languageDiv = document.getElementById('detected_language');
|
||||
|
||||
if (transcript_data['words'] && transcript_data['words'].length > 0) {
|
||||
// Append words with color based on their probability
|
||||
transcript_data['words'].forEach(wordData => {
|
||||
const span = document.createElement('span');
|
||||
const probability = wordData['probability'];
|
||||
span.textContent = wordData['word'] + ' ';
|
||||
|
||||
// Set the color based on the probability
|
||||
if (probability > 0.9) {
|
||||
span.style.color = 'green';
|
||||
} else if (probability > 0.6) {
|
||||
span.style.color = 'orange';
|
||||
} else {
|
||||
span.style.color = 'red';
|
||||
}
|
||||
|
||||
transcriptionDiv.appendChild(span);
|
||||
});
|
||||
|
||||
// Add a new line at the end
|
||||
transcriptionDiv.appendChild(document.createElement('br'));
|
||||
} else {
|
||||
// Fallback to plain text
|
||||
transcriptionDiv.textContent += transcript_data['text'] + '\n';
|
||||
}
|
||||
|
||||
// Update the language information
|
||||
if (transcript_data['language'] && transcript_data['language_probability']) {
|
||||
languageDiv.textContent = transcript_data['language'] + ' (' + transcript_data['language_probability'].toFixed(2) + ')';
|
||||
}
|
||||
|
||||
// Update the processing time, if available
|
||||
const processingTimeDiv = document.getElementById('processing_time');
|
||||
if (transcript_data['processing_time']) {
|
||||
processingTimeDiv.textContent = 'Processing time: ' + transcript_data['processing_time'].toFixed(2) + ' seconds';
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
|
||||
|
||||
function sendAudioConfig() {
|
||||
let selectedStrategy = document.getElementById('bufferingStrategySelect').value;
|
||||
let processingArgs = {};
|
||||
|
||||
if (selectedStrategy === 'silence_at_end_of_chunk') {
|
||||
processingArgs = {
|
||||
chunk_length_seconds: parseFloat(document.getElementById('chunk_length_seconds').value),
|
||||
chunk_offset_seconds: parseFloat(document.getElementById('chunk_offset_seconds').value)
|
||||
};
|
||||
}
|
||||
|
||||
const audioConfig = {
|
||||
type: 'config',
|
||||
data: {
|
||||
sampleRate: context.sampleRate,
|
||||
bufferSize: bufferSize,
|
||||
channels: 1, // Assuming mono channel
|
||||
language: language,
|
||||
processing_strategy: selectedStrategy,
|
||||
processing_args: processingArgs
|
||||
}
|
||||
};
|
||||
|
||||
websocket.send(JSON.stringify(audioConfig));
|
||||
}
|
||||
|
||||
function downsampleBuffer(buffer, inputSampleRate, outputSampleRate) {
|
||||
if (inputSampleRate === outputSampleRate) {
|
||||
return buffer;
|
||||
}
|
||||
var sampleRateRatio = inputSampleRate / outputSampleRate;
|
||||
var newLength = Math.round(buffer.length / sampleRateRatio);
|
||||
var result = new Float32Array(newLength);
|
||||
var offsetResult = 0;
|
||||
var offsetBuffer = 0;
|
||||
while (offsetResult < result.length) {
|
||||
var nextOffsetBuffer = Math.round((offsetResult + 1) * sampleRateRatio);
|
||||
var accum = 0, count = 0;
|
||||
for (var i = offsetBuffer; i < nextOffsetBuffer && i < buffer.length; i++) {
|
||||
accum += buffer[i];
|
||||
count++;
|
||||
}
|
||||
result[offsetResult] = accum / count;
|
||||
offsetResult++;
|
||||
offsetBuffer = nextOffsetBuffer;
|
||||
}
|
||||
return result;
|
||||
}
|
||||
|
||||
function processAudio(e) {
|
||||
const inputSampleRate = context.sampleRate;
|
||||
const outputSampleRate = 16000; // Target sample rate
|
||||
|
||||
const left = e.inputBuffer.getChannelData(0);
|
||||
const downsampledBuffer = downsampleBuffer(left, inputSampleRate, outputSampleRate);
|
||||
const audioData = convertFloat32ToInt16(downsampledBuffer);
|
||||
|
||||
if (websocket && websocket.readyState === WebSocket.OPEN) {
|
||||
websocket.send(audioData);
|
||||
}
|
||||
}
|
||||
|
||||
function convertFloat32ToInt16(buffer) {
|
||||
let l = buffer.length;
|
||||
const buf = new Int16Array(l);
|
||||
while (l--) {
|
||||
buf[l] = Math.min(1, buffer[l]) * 0x7FFF;
|
||||
}
|
||||
return buf.buffer;
|
||||
}
|
||||
|
||||
|
||||
function toggleBufferingStrategyPanel() {
|
||||
var selectedStrategy = document.getElementById('bufferingStrategySelect').value;
|
||||
if (selectedStrategy === 'silence_at_end_of_chunk') {
|
||||
var panel = document.getElementById('silence_at_end_of_chunk_options_panel');
|
||||
panel.classList.remove('hidden');
|
||||
} else {
|
||||
var panel = document.getElementById('silence_at_end_of_chunk_options_panel');
|
||||
panel.classList.add('hidden');
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
</script>
|
||||
</body>
|
||||
|
||||
</html>
|
||||
Reference in New Issue
Block a user