Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 0 additions & 4 deletions ml_v2/encoder/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,14 +17,10 @@ async def vectorize_audio(file: UploadFile = File(...)):

if not file:
raise HTTPException(status_code=400, detail="file not provided")


audio = preprocess_audio(file)

logger.info("Preprocessed audio successfully")

embedding = extract_embedding(audio)

logger.info("Generated embedding successfully")

return {
Expand Down
135 changes: 133 additions & 2 deletions ml_v2/orchastrater/main.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,137 @@
def main():
print("Hello from orchastrater!")
from fileinput import filename
import os
import shutil
import tempfile
from typing import final
from fastapi import FastAPI, UploadFile, File, Form, HTTPException
from fastapi.middleware.cors import CORSMiddleware
import uvicorn


from utils.logger import get_logger

from services import whisper_utils
from services import remote_encoder
from services import db_embeddings
from services import db_clusters
from services import clustering
from services import supabase

app = FastAPI(title="BhashaSuraksha Orchastrator")
logger = get_logger("main")


app.add_middleware(
CORSMiddleware,
allow_origins=["*"], # In production, change this to your Frontend URL
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)

@app.get("/")
def health_check():
return {"status" : "healthy", "service": "orchastrator"}

@app.post("/process-audio")
async def process_audio(
file: UploadFile = File(...),
region:str = Form("Unknown"),
lat:float = Form(None),
lng:float = Form(None)
):
tmp_path = None
try:
logger.info(f"Recieved request:{file.filename} from region:{region}")

filename = file.filename or "audio.wav"
file_ext = os.path.splitext(filename)[1] or ".wav"
with tempfile.NamedTemporaryFile(delete=False, suffix=file_ext) as tmp:
shutil.copyfileobj(file.file,tmp)
tmp_path = tmp.name
logger.info(f"saved temp file into {tmp_path}")

#whisper
transcription = whisper_utils.transcribe_audio(tmp_path)
transcript_text = transcription["text"]
detected_language = transcription["language"]
confidence = transcription["probability"]
logger.info(f"Whisper results: transcription:{transcript_text} with language:{detected_language} with confidence:{confidence}")

#encoder
embedding = await remote_encoder.get_audio_embedding(tmp_path)

if not embedding:
raise HTTPException(status_code=500,detail="Failed to generate embedding")

#clustering
existing_clusters = db_clusters.get_all_clusters()
best_cluster_id, distance = clustering.find_best_cluster(embedding,existing_clusters)
final_cluster_id = None

if best_cluster_id is not None:
logger.info(f"Joining Cluster {best_cluster_id} (Distance:{distance:.4f})")
final_cluster_id = best_cluster_id

match = next(c for c in existing_clusters if c['id']==best_cluster_id)

new_centroid = clustering.calculate_new_centroid(
match["centroid"],
match["sampleCount"],
embedding
)

db_clusters.update_cluster_centroid(
best_cluster_id,
new_centroid,
match["sampleCount"]+1
)

else:
logger.info("No matching cluster found, Creating new Cluster")
final_cluster_id = db_clusters.create_new_cluster(embedding)

public_url = supabase.upload_audio_file(tmp_path,filename)

sample_id = db_embeddings.create_unknown_sample(
file_url=public_url,
language_guess=detected_language,
confidence=confidence,
transcript=transcript_text,
region=region,
lat=lat,
lng=lng,
keywords="",
embedding=embedding,
cluster_id=final_cluster_id
)

return {
"status": "success",
"sample_id": sample_id,
"transcript": transcript_text,
"detected_language":detected_language,
"assigned_cluster_id": final_cluster_id,
"is_new_cluster": (final_cluster_id != best_cluster_id) if best_cluster_id is not None else True,
"file_url": public_url
}

except Exception as e:
logger.error(f"Processing failed: {e}")
raise HTTPException(status_code=500,detail=str(e))

finally:
if tmp_path and os.path.exists(tmp_path):
os.remove(tmp_path)
logger.info("cleaned temp path")

def main():
uvicorn.run(
"main:app",
host="0.0.0.0",
port=8000,
reload=True
)

if __name__ == "__main__":
main()
5 changes: 5 additions & 0 deletions ml_v2/orchastrater/pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -7,9 +7,14 @@ requires-python = ">=3.11"
dependencies = [
"fastapi>=0.128.0",
"faster-whisper>=1.2.1",
"httpx>=0.28.1",
"numpy>=2.4.1",
"psycopg2-binary>=2.9.11",
"pydantic>=2.12.5",
"python-dotenv>=1.2.1",
"python-multipart>=0.0.21",
"requests>=2.32.5",
"scikit-learn>=1.8.0",
"supabase>=2.27.2",
"uvicorn>=0.40.0",
]
60 changes: 60 additions & 0 deletions ml_v2/orchastrater/services/clustering.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,60 @@
import numpy as np
from sklearn.metrics.pairwise import cosine_distances
from utils.logger import get_logger

logger = get_logger("clustering")

# 0.0 = Identical and 1.0 = Completely Opposite
SIMILARITY_THRESHOLD = 0.2

def find_best_cluster(new_embedding: list, existing_clusters: list):
"""
args:
new_embedding: the embedding vector of the users voice
existing_clusters: a list of dicts of all the clusters

output:
tuple(best_cluster_id, distance) or (None, None) if None exists
"""
if not existing_clusters:
return (None, None)

centroids = np.array([c['centroid'] for c in existing_clusters])
user_vector = np.array(new_embedding).reshape(1,-1)

"""
The [0] here is used to flatten the list
example output without [0]:
[
[0.1,0.2,0.3]
]
with [0]:
[0.1,0.2,0.3]
"""
dists = cosine_distances(user_vector,centroids)[0]

min_dist_index = np.argmin(dists)
min_dist = dists[min_dist_index]

best_cluster = existing_clusters[min_dist_index]

logger.info(f"Closest cluster: ID {best_cluster['id']} and Distance {min_dist:.4f}")

if min_dist < SIMILARITY_THRESHOLD:
return best_cluster['id'], float(min_dist)
else:
logger.info(f"No match found. CLosest was {min_dist:.4f} > {SIMILARITY_THRESHOLD}")
return None, float(min_dist)

def calculate_new_centroid(current_centroid: list, current_count: int, new_embedding: list):
"""
Updates a specific clusters centroid using a wighted average with the formula:

New = ((Old * Count) + New_sample) / Count+1
"""
old_vec = np.array(current_centroid)
new_vec = np.array(new_embedding)
updated_vec = ((old_vec * current_count) + new_vec) / (current_count + 1)

return updated_vec.tolist()

39 changes: 39 additions & 0 deletions ml_v2/orchastrater/services/db_clusters.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,39 @@
import json
from utils.db import excecute_query
from utils.logger import get_logger

logger = get_logger("db_clusters")

def get_all_clusters():
query = 'SELECT "id", "centroid", "sampleCount" FROM "Cluster";'

return excecute_query(query, fetch_all=True)

def create_new_cluster(centroid: list):

query = """
INSERT INTO "Cluster" ("centroid", "sampleCount","createdAt") VALUES (%s::jsonb, 1, NOW())
RETURNING id;
"""

centroid_json = json.dumps(centroid)
result = excecute_query(query, (centroid_json,), fetch_one=True)
logger.info(f"Created a new cluster with id: {result['id']}")

return result['id']

def update_cluster_centroid(cluster_id: int, new_centroid: list, new_count: int):

query = """
UPDATE "Cluster"
SET "centroid" = %s::jsonb, "sampleCount" = %s
WHERE "id" = %s;
"""

centroid_json = json.dumps(new_centroid)
excecute_query(query, (centroid_json, new_count, cluster_id))

logger.info(f"Updated cluster: {cluster_id} centroid with the new count: {new_count}")



61 changes: 61 additions & 0 deletions ml_v2/orchastrater/services/db_embeddings.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,61 @@
import json
from utils.db import excecute_query
from utils.logger import get_logger

logger = get_logger("db_embeddings")

def create_unknown_sample(
file_url: str,
language_guess: str,
confidence: str,
transcript: str,
region: str,
lat: float,
lng: float,
keywords: str,
embedding: list,
cluster_id: int = None
):
"""To insert into the unknown samples table"""

query = """
INSERT INTO "UnknownSample"(
"fileUrl",
"languageGuess",
"confidence",
"transcript",
"region",
"lat",
"lng",
"keywords",
"embedding",
"clusterId"
)
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s::jsonb, %s)
RETURNING id;
"""

embedding_json = json.dumps(embedding)

params = (
file_url,
language_guess,
confidence,
transcript,
region,
lat,
lng,
keywords,
embedding_json,
cluster_id,
)

try:
result = excecute_query(query,params,fetch_one=True)
if result:
logger.info(f"Saved UnknownSample with id: {result['id']}")
return result['id']

except Exception as e:
logger.error(f"Failed to save UnknownSample: {e}")
raise
15 changes: 7 additions & 8 deletions ml_v2/orchastrater/services/remote_encoder.py
Original file line number Diff line number Diff line change
@@ -1,12 +1,12 @@
# orchestrator/services/remote_encoder.py
import httpx
import requests
import os
from utils.logger import get_logger
from utils.env import ENCODER_URL

logger = get_logger(__name__)
def get_audio_embedding(file_path: str):

async def get_audio_embedding(file_path: str):
"""
Sends an audio file to the Encoder service and returns the vector.
"""
Expand All @@ -19,13 +19,12 @@ def get_audio_embedding(file_path: str):
logger.info(f"Sending {file_path} to Encoder service at {url}...")

try:
with open(file_path,"rb") as f:
files = {"file": f}

response = requests.post(url, files=files)
async with httpx.AsyncClient(timeout=30.0) as client:
with open(file_path,"rb") as f:
files = {"file": f}
response = await client.post(url, files=files)

response.raise_for_status()

data = response.json()

return data["embedding"]
Expand Down
Loading
Loading