From 81768be08547d5bbd645323529b6af80015a2920 Mon Sep 17 00:00:00 2001 From: makaveli10 Date: Wed, 31 May 2023 19:00:12 +0530 Subject: [PATCH] update client to websocket --- client.py | 194 +++++++++++++--------------------------- requirements/client.txt | 3 +- 2 files changed, 65 insertions(+), 132 deletions(-) diff --git a/client.py b/client.py index 8c908bd..a4ae3d7 100644 --- a/client.py +++ b/client.py @@ -2,56 +2,84 @@ import io import os import argparse import wave -import uuid -import hashlib -import base64 -import time import numpy as np import scipy import ffmpeg import torch -import socket, pickle, pyaudio, struct +import pyaudio import threading import textwrap import json import torchaudio -from dataclasses import dataclass +import websocket + CHUNK = 1024 FORMAT = pyaudio.paInt16 CHANNELS = 1 RATE = 16000 RECORD_SECONDS = 60000 -all_segments = [] -@dataclass(frozen=True) -class Constants: - ACK = b"acknowledged" - RECORDING_OVER = b"audio_data_over" - RECEIVED_AUDIO_FILE = b"audio_file_sent" - RECEIVING_AUDIO_FILE = b"sending_audio_file" + +def on_message(ws, message): + message = json.loads(message) + text = [] + if len(message): + for seg in message: + if len(text): + if text[-1] != seg["text"]: + text.append(seg["text"]) + else: + text.append(seg["text"]) + if len(text) > 3: + text = text[-3:] + wrapper = textwrap.TextWrapper(width=60) + word_list = wrapper.wrap(text="".join(text)) + # Print each line. + os.system('clear') + for element in word_list: + print(element) + +def on_error(ws, error): + print(error) + +def on_close(ws, close_status_code, close_msg): + print("### websocket connection closed ###") + +def on_open(ws): + print("Opened connection") + class Client: - def __init__(self, topic=None, host=None, port=None): + def __init__(self, host=None, port=None): self.timestamp_offset = 0.0 self.audio_bytes = None self.p = pyaudio.PyAudio() - self.payload_size = struct.calcsize("Q") self.stream = self.p.open(format=FORMAT, channels=CHANNELS, rate=RATE, input=True, frames_per_buffer=CHUNK) print(self.p.get_sample_size(FORMAT)) - self.client_socket = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - host_ip = 'localhost' if host is None else host - port = 5901 if port is None else port - socket_address = (host_ip, port) - self.client_socket.connect(socket_address) - print("CLIENT CONNECTED TO", socket_address) + # create websocket connection + if host is not None and port is not None: + socket_url = f"ws://{host}:{port}" + self.client_socket = websocket.WebSocketApp(socket_url, + on_open=on_open, + on_message=on_message, + on_error=on_error, + on_close=on_close) + else: + print("No host or port specified.") + return + + # start websocket client in a thread + self.ws_thread = threading.Thread(target=self.client_socket.run_forever) + self.ws_thread.setDaemon(True) + self.ws_thread.start() # voice activity detection model self.vad_model, _ = torch.hub.load(repo_or_dir='snakers4/silero-vad', @@ -61,43 +89,14 @@ class Client: self.window_size = 1024 self.vad_threshold = 0.4 - # subscribing to the correct topic - if topic is not None: - self.topic = topic - else: - self.topic = self.get_mac_address().decode() - self.frames = b"" - data = b"" - while True: - while len(data) < self.payload_size: - packet = self.client_socket.recv(4*1024) #4K - if not packet: break - data+=packet - packed_msg_size = data[:self.payload_size] - data = data[self.payload_size:] - try: - msg_size = struct.unpack("Q",packed_msg_size)[0] - except struct.error: - break - while len(data) < msg_size: - data += self.client_socket.recv(4*1024) - frame_data = data[:msg_size] - frame_data = pickle.loads(frame_data) - if Constants.ACK in frame_data: - print("Server is ready. Sending audio ...") - break print("* recording") def send_packet_to_server(self, message): - a = pickle.dumps(message) - message = struct.pack("Q",len(a))+a - self.client_socket.sendall(message) - - def get_mac_address(self): - mac = hex(uuid.getnode()) - hasher = hashlib.sha1(mac.encode()) - return base64.urlsafe_b64encode(hasher.digest()[:5]) + try: + self.client_socket.send(message, websocket.ABNF.OPCODE_BINARY) + except Exception as e: + print(e) @staticmethod def bytes_to_audio_tensor(audio_bytes): @@ -107,7 +106,7 @@ class Client: ) scipy.io.wavfile.write(bytes_io, RATE, raw_data) audio, _ = torchaudio.load(bytes_io) - return audio.squeeze(0) + return audio.squeeze(0), raw_data.astype(np.float32) / 32768.0 def play_file(self, filename): # read audio and create pyaudio stream @@ -124,34 +123,20 @@ class Client: if data==b'': break # voice activity detection - chunk_tensor = Client.bytes_to_audio_tensor(data) + chunk_tensor, audio_array = Client.bytes_to_audio_tensor(data) try: speech_prob = self.vad_model(chunk_tensor, RATE).item() except ValueError: break # input audio chunk is too short if speech_prob > self.vad_threshold: - data_dict = { - "topic": self.topic, - "audio": data - } - self.send_packet_to_server(data_dict) + self.send_packet_to_server(audio_array.tobytes()) self.stream.write(data) self.wf.close() self.stream.close() - # let the server know that we're done - data = Constants.RECORDING_OVER - self.send_packet_to_server(data) - with open("results.json", "w") as f: - json_dict = json.dumps(all_segments, indent=2) - f.write(json_dict) - except KeyboardInterrupt: - # write all segments to a file - with open("results.json", "w") as f: - json_dict = json.dumps(all_segments, indent=2) - f.write(json_dict) + print("Keyboard interrupt.") def get_client_socket(self): @@ -176,15 +161,11 @@ class Client: self.frames += data # voice activity detection - chunk_tensor = Client.bytes_to_audio_tensor(data) + chunk_tensor , audio_array = Client.bytes_to_audio_tensor(data) speech_prob = self.vad_model(chunk_tensor, RATE).item() if speech_prob > self.vad_threshold: - data_dict = { - "topic": self.topic, - "audio": data - } - self.send_packet_to_server(data_dict) + self.send_packet_to_server(audio_array.tobytes()) # save frames if more than a minute if len(self.frames) > 60*RATE: @@ -205,16 +186,9 @@ class Client: self.stream.close() self.p.terminate() - # let the server know that we're done - data = Constants.RECORDING_OVER - self.send_packet_to_server(data) # combine all the audio files self.write_output_recording(n_audio_file, out_file) - # write all segments to a file - with open("results.json", "w") as f: - json_dict = json.dumps(all_segments, indent=2) - f.write(json_dict) def write_output_recording(self, n_audio_file, out_file): input_files = [f"chunks/{i}.wav" for i in range(n_audio_file) if os.path.exists(f"chunks/{i}.wav")] @@ -234,41 +208,6 @@ class Client: wf.close() -def recieve_response(client_socket): - data = b"" - payload_size = struct.calcsize("Q") - - while True: - while len(data) < payload_size: - packet = client_socket.recv(4*1024) # 4K - if not packet: break - data+=packet - packed_msg_size = data[:payload_size] - data = data[payload_size:] - try: - msg_size = struct.unpack("Q",packed_msg_size)[0] - except struct.error: - break - while len(data) < msg_size: - data += client_socket.recv(4*1024) - frame_data = data[:msg_size] - data = data[msg_size:] - response = pickle.loads(frame_data) - - if response is not None and isinstance(response, dict): - os.system('clear') - text = response['text'] - segments = response['segments'] - if len(segments): - for seg in segments: - all_segments.append(seg) - wrapper = textwrap.TextWrapper(width=50) - word_list = wrapper.wrap(text=text) - # Print each line. - for element in word_list: - print(element) - - def resample(file: str, sr: int = 16000): """ # https://github.com/openai/whisper/blob/7858aa9c08d98f75575035ecd6481f462d66ca27/whisper/audio.py#L22 @@ -301,20 +240,13 @@ def resample(file: str, sr: int = 16000): if __name__=="__main__": parser = argparse.ArgumentParser() parser.add_argument('--audio', type=str, help='audio file to transcribe') - parser.add_argument('--topic', default=None, type=str, help='topic to subscribe for results') - parser.add_argument('--host', default=None, type=str, help='server address to connect to') - parser.add_argument('--port', default=None, type=str, help='server port to connect to') + parser.add_argument('--host', default=None, type=str, help='websocket server address to connect to') + parser.add_argument('--port', default=None, type=str, help='websocket server port to connect to') opt = parser.parse_args() - c = Client(topic=opt.topic, host=opt.host, port=opt.port) - while True: - if c.get_client_socket() is not None: - break - client_socket = c.get_client_socket() - t2 = threading.Thread(target=recieve_response, args=(client_socket, )) - t2.start() + c = Client(host=opt.host, port=opt.port) + if opt.audio is not None: resampled_file = resample(opt.audio) c.play_file(resampled_file) else: c.record() - t2.join() diff --git a/requirements/client.txt b/requirements/client.txt index ce9121d..dec830f 100644 --- a/requirements/client.txt +++ b/requirements/client.txt @@ -2,4 +2,5 @@ PyAudio ffmpeg-python scipy torch==1.12.1 -torchaudio==0.12.1 \ No newline at end of file +torchaudio==0.12.1 +websocket-client \ No newline at end of file