Compare commits

...
9 Commits
Author SHA1 Message Date
shadowcz007 3bac87ee52 system prompt 2024-06-04 08:50:58 +08:00
shadowcz007 9a01701019 whisper+chat 2024-06-04 08:40:17 +08:00
shadowcz007 693954ee23 ing 2024-06-03 20:12:12 +08:00
shadowcz007 bf4ba91e7a update 2024-06-03 17:43:41 +08:00
shadowcz007 0828353253 Update main.py 2024-06-03 16:16:02 +08:00
shadowcz007 1997c7ad8f Update live.html 2024-06-03 16:12:06 +08:00
shadowcz007 b1e62440e4 test 2024-06-02 22:38:15 +08:00
shadowcz007 d549a5eb6a whisper 2024-06-02 19:53:07 +08:00
shadowcz007 77bfb08d76 web 2024-06-02 17:12:11 +08:00
20 changed files with 1327 additions and 3 deletions
+80 -3
View File
@@ -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):
View File
+12
View File
@@ -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}")
+9
View File
@@ -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
+27
View File
@@ -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
+26
View File
@@ -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.")
+54
View File
@@ -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)
+54
View File
@@ -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()
+7
View File
@@ -0,0 +1,7 @@
websockets
speechbrain
pyannote-audio
asyncio
sentence-transformers
transformers
faster-whisper
+88
View File
@@ -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)
View File
+50
View File
@@ -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
+23
View File
@@ -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}")
+16
View File
@@ -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
View File
@@ -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>