speaker-id/ah.py
yumoqing 41b362f0b6 feat: speaker-id service skeleton with ahserver + longtasks
- ahserver app on port 9095
- endpoints: status, enroll, identify, verify, compare
- ECAPA-TDNN model placeholder
- longtasks async task runner (enroll/identify/verify)
- direct cosine comparison endpoint
2026-07-21 17:33:07 +08:00

166 lines
5.8 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

# -*- coding:utf-8 -*-
"""
Speaker ID Service — 声纹识别服务
基于 SpeechBrain ECAPA-TDNNahserver + longtasks 异步模式。
"""
from ahserver.webapp import webapp
from ahserver.serverenv import ServerEnv
from ahserver.configuredServer import add_startup
from longtasks.longtasks import LongTasks, schedule_once
from appPublic.registerfunction import RegisterFunction
from appPublic.log import debug, exception
from traceback import format_exc
import json
# ── Task runner ────────────────────────────────────────────
class SpeakerTasks(LongTasks):
async def process_task(self, payload, workid=None):
if isinstance(payload, str):
payload = json.loads(payload)
task_type = payload.get('task_type', '')
if task_type == 'enroll':
return await self._enroll(payload)
elif task_type == 'identify':
return await self._identify(payload)
elif task_type == 'verify':
return await self._verify(payload)
raise ValueError(f'Unknown task_type: {task_type}')
async def _enroll(self, payload):
audio = payload.get('audio', '')
speaker_id = payload.get('speaker_id', '')
# TODO: load model, compute embedding, store
return {'speaker_id': speaker_id, 'status': 'enrolled'}
async def _identify(self, payload):
audio = payload.get('audio', '')
# TODO: compute embedding, search DB
return {'matches': []}
async def _verify(self, payload):
audio1 = payload.get('audio1', '')
audio2 = payload.get('audio2', '')
# TODO: compute embeddings, compare
return {'match': False, 'score': 0.0}
# ── API handlers ───────────────────────────────────────────
async def status_handler(request, params_kw, *args, **kwargs):
return json.dumps({
"service": "speaker-id",
"model": "ECAPA-TDNN (spkrec-ecapa-voxceleb)",
"model_loaded": False,
"embedding_dim": 192,
"endpoints": [
"/api/status", "/api/enroll", "/api/identify",
"/api/verify", "/api/compare"
]
}, indent=2, ensure_ascii=False)
async def enroll_handler(request, params_kw, *args, **kwargs):
try:
env = request._run_ns
longtasks = env.longtasks
audio = params_kw.get('audio', '')
speaker_id = params_kw.get('speaker_id', '')
speaker_name = params_kw.get('speaker_name', '')
if not audio or not speaker_id:
return json.dumps({"error": "audio and speaker_id required"})
task_id = await longtasks.submit({
'task_type': 'enroll',
'audio': audio,
'speaker_id': speaker_id,
'speaker_name': speaker_name
})
return json.dumps({"status": "ACCEPTED", "task_id": task_id})
except Exception as e:
exception(f"enroll: {e}, {format_exc()}")
return json.dumps({"error": str(e)})
async def identify_handler(request, params_kw, *args, **kwargs):
try:
env = request._run_ns
longtasks = env.longtasks
audio = params_kw.get('audio', '')
if not audio:
return json.dumps({"error": "audio required"})
task_id = await longtasks.submit({
'task_type': 'identify',
'audio': audio,
'threshold': float(params_kw.get('threshold', 0.7)),
'top_k': int(params_kw.get('top_k', 3))
})
return json.dumps({"status": "ACCEPTED", "task_id": task_id})
except Exception as e:
exception(f"identify: {e}, {format_exc()}")
return json.dumps({"error": str(e)})
async def verify_handler(request, params_kw, *args, **kwargs):
try:
env = request._run_ns
longtasks = env.longtasks
audio1 = params_kw.get('audio1', '')
audio2 = params_kw.get('audio2', '')
if not audio1 or not audio2:
return json.dumps({"error": "audio1 and audio2 required"})
task_id = await longtasks.submit({
'task_type': 'verify',
'audio1': audio1,
'audio2': audio2
})
return json.dumps({"status": "ACCEPTED", "task_id": task_id})
except Exception as e:
exception(f"verify: {e}, {format_exc()}")
return json.dumps({"error": str(e)})
async def compare_handler(request, params_kw, *args, **kwargs):
"""直接比对两个声纹向量 (实时, 不走 longtasks)"""
emb1 = params_kw.get('embedding1', [])
emb2 = params_kw.get('embedding2', [])
if not emb1 or not emb2:
return json.dumps({"error": "embedding1 and embedding2 required"})
# Cosine similarity
import math
dot = sum(a * b for a, b in zip(emb1, emb2))
norm1 = math.sqrt(sum(a * a for a in emb1))
norm2 = math.sqrt(sum(b * b for b in emb2))
score = dot / (norm1 * norm2) if norm1 and norm2 else 0.0
return json.dumps({"status": "SUCCEEDED", "similarity": round(score, 6)})
async def on_app_built(app):
env = ServerEnv()
longtasks = env.longtasks
if longtasks:
schedule_once(0.1, longtasks.run)
debug('speaker-id longtasks worker started')
def init():
env = ServerEnv()
longtasks = SpeakerTasks(
'redis://127.0.0.1:6379',
'speaker_id',
worker_cnt=2,
stuck_seconds=1800,
max_age_hours=24
)
env.longtasks = longtasks
add_startup(on_app_built)
rf = RegisterFunction()
rf.register("status", status_handler)
rf.register("enroll", enroll_handler)
rf.register("identify", identify_handler)
rf.register("verify", verify_handler)
rf.register("compare", compare_handler)
if __name__ == '__main__':
webapp(init)