1577-2 / backend /utils.py
jts-ai-team's picture
Upload 23 files
1354c32 verified
Raw History Blame
3.45 kB
import numpy as np
import librosa
import io
import os
import warnings
from pydub import AudioSegment
from dotenv import load_dotenv
from fastrtc import get_cloudflare_turn_credentials_async, get_cloudflare_turn_credentials
try:
import torch
except ModuleNotFoundError:
torch = None # type: ignore
warnings.filterwarnings("ignore")
# load_dotenv(override = True)
# --- Device Configuration ---
def get_device():
"""Gets the best available device for PyTorch."""
if torch is None:
return "cpu"
if torch.cuda.is_available():
return "cuda"
elif hasattr(torch.backends, "mps") and torch.backends.mps.is_available():
return "mps"
else:
return "cpu"
device = get_device()
print(f"Using device: {device}")
# --- Cloud Credentials ---
async def get_async_credentials():
"""Asynchronously fetches Cloudflare TURN credentials."""
return await get_cloudflare_turn_credentials_async(hf_token=os.getenv('HF_TOKEN'))
def get_sync_credentials(ttl=360_000):
"""Synchronously fetches Cloudflare TURN credentials."""
return get_cloudflare_turn_credentials(ttl=ttl)
def setup_gcp_credentials():
"""Sets up Google Cloud credentials from an environment variable."""
gcp_service_account_json_str = os.getenv("GCP_SERVICE_ACCOUNT_JSON")
if gcp_service_account_json_str:
print("GCP service account JSON loaded from environment variable.")
else:
print("Warning: GCP_SERVICE_ACCOUNT_JSON is not set; Google Cloud clients may fail.")
return gcp_service_account_json_str
# --- Audio Processing ---
def audiosegment_to_numpy(audio, target_sample_rate=16000):
samples = np.array(audio.get_array_of_samples(), dtype=np.float32)
if audio.channels > 1:
samples = samples.reshape((-1, audio.channels)).mean(axis=1)
samples /= np.iinfo(audio.array_type).max
if audio.frame_rate != target_sample_rate:
samples = librosa.resample(samples, orig_sr=audio.frame_rate, target_sr=target_sample_rate)
return samples
def preprocess_audio(audio, target_channels=1, target_frame_rate=16000):
"""
Preprocess the audio using pydub AudioSegment by setting the number of channels and frame rate.
Args:
audio (tuple): A tuple (sample_rate, audio_array) where audio_array is a NumPy array.
target_channels (int): Desired number of channels (default is 1 for mono).
target_frame_rate (int): Desired frame rate (default is 16000 Hz).
Returns:
np.ndarray: The processed audio as a NumPy array.
"""
sample_rate, audio_array = audio
target_frame_rate = sample_rate
audio_array_int16 = audio_array.astype(np.int16)
audio_bytes = audio_array_int16.tobytes()
audio_io = io.BytesIO(audio_bytes)
segment = AudioSegment.from_raw(audio_io, sample_width=2, frame_rate=sample_rate, channels=1)
segment = segment.set_channels(target_channels)
segment = segment.set_frame_rate(target_frame_rate)
return audiosegment_to_numpy(segment)
# --- Conversation Utilities ---
def is_valid_turn(turn):
"""Return True if turn is a valid dict with non-empty 'role' and 'content' strings."""
return (
isinstance(turn, dict)
and "role" in turn
and "content" in turn
and isinstance(turn["role"], str)
and isinstance(turn["content"], str)
and turn["role"].strip() != ""
and turn["content"].strip() != ""
)