Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
664d505446 | ||
|
|
cd5b471892 | ||
|
|
826778e9fd | ||
|
|
cb6d8c6d94 |
+23
-4
@@ -16,8 +16,14 @@ __all__ = ["EXTRA_MODELS", "normalize", "disputed", "build_draft"]
|
||||
|
||||
log = logging.getLogger("talkscore-asr")
|
||||
|
||||
# Модели подобраны разными по устройству: RNN-T молчит на трудных участках,
|
||||
# CTC говорит хоть что-то. Их ошибки не совпадают, поэтому согласие значимо.
|
||||
# RNN-T молчит на трудных участках, CTC говорит хоть что-то - ошибаются они
|
||||
# в разных местах, поэтому согласие значимо.
|
||||
#
|
||||
# Сторонние модели проверены на эталоне и не подошли: Vosk 61 %, NeMo ru
|
||||
# 67 и 72 %, T-one 74 %, Whisper 145 % при 33 % у GigaAM. Разница слишком
|
||||
# велика - их несогласие было бы шумом, а не сигналом. Варианты GigaAM
|
||||
# обучены одной командой, и это ограничение метода: места, где ошибаются
|
||||
# все три, черновик не подсветит.
|
||||
EXTRA_MODELS = ("gigaam-v3-e2e-ctc", "gigaam-v2-rnnt")
|
||||
WORD_RE = re.compile(r"[\w-]+", re.UNICODE)
|
||||
GAP_NOTICE_SEC = 3.0
|
||||
@@ -57,7 +63,8 @@ def disputed(base: list[str], others: list[list[str]]) -> list[set[str]]:
|
||||
return variants
|
||||
|
||||
|
||||
def build_draft(pipeline, samples, duration: float, progress=None) -> dict:
|
||||
def build_draft(pipeline, samples, duration: float, progress=None,
|
||||
threads: int = 0) -> dict:
|
||||
"""Прогоняет запись несколькими моделями и размечает расхождения.
|
||||
|
||||
pipeline уже держит основную модель и детектор речи - второй копии
|
||||
@@ -72,8 +79,11 @@ def build_draft(pipeline, samples, duration: float, progress=None) -> dict:
|
||||
for name in EXTRA_MODELS:
|
||||
if progress:
|
||||
progress(f"загружаю модель {name}")
|
||||
# Черновик собирается, когда очередь обычно пуста: отдаём моделям
|
||||
# все ядра, иначе они делят те же потоки, что и обычная задача.
|
||||
options = {"sess_options": _session_options(threads)} if threads else {}
|
||||
try:
|
||||
extra[name] = onnx_asr.load_model(name, quantization="int8")
|
||||
extra[name] = onnx_asr.load_model(name, quantization="int8", **options)
|
||||
except Exception as exc: # noqa: BLE001 - без модели просто меньше сверок
|
||||
log.warning("модель %s недоступна: %s", name, exc)
|
||||
|
||||
@@ -121,6 +131,15 @@ def build_draft(pipeline, samples, duration: float, progress=None) -> dict:
|
||||
}
|
||||
|
||||
|
||||
def _session_options(threads: int):
|
||||
import onnxruntime
|
||||
|
||||
options = onnxruntime.SessionOptions()
|
||||
options.intra_op_num_threads = threads
|
||||
options.inter_op_num_threads = 1
|
||||
return options
|
||||
|
||||
|
||||
def _gaps(turns: list[dict], duration: float) -> list[dict]:
|
||||
"""Промежутки, где не распознала ни одна модель.
|
||||
|
||||
|
||||
+13
-12
@@ -348,7 +348,9 @@ def _build_draft(job_id: str, audio: Path) -> None:
|
||||
wav = Path(tmp) / "audio.wav"
|
||||
to_wav16k(audio, wav, pipeline.ffmpeg, settings.normalize)
|
||||
samples = read_wav(wav)
|
||||
result = build_draft(pipeline, samples, len(samples) / 16000, progress)
|
||||
# Пока считается черновик, очередь обычно пуста - берём все ядра.
|
||||
result = build_draft(pipeline, samples, len(samples) / 16000, progress,
|
||||
threads=os.cpu_count() or 4)
|
||||
with _draft_lock:
|
||||
_drafts[job_id] = {"status": "ready", "draft": result}
|
||||
log.info("черновик по задаче %s готов: слов %d, согласны %d",
|
||||
@@ -360,29 +362,28 @@ def _build_draft(job_id: str, audio: Path) -> None:
|
||||
|
||||
|
||||
@app.get("/v1/jobs/{job_id}/draft", include_in_schema=False)
|
||||
def job_draft(job_id: str, request: Request, build: bool = False):
|
||||
def job_draft(job_id: str, request: Request):
|
||||
"""Черновик для правки на слух. Первый запрос запускает сборку."""
|
||||
token = page_guard(request)
|
||||
with _draft_lock:
|
||||
state = dict(_drafts.get(job_id) or {})
|
||||
|
||||
# Известное состояние отвечает само за себя: запись могли удалить по
|
||||
# сроку, пока черновик собирался, и отвечать тогда 404 неправильно.
|
||||
if state.get("status") == "ready":
|
||||
return HTMLResponse(render_draft(state["draft"], job_id, token))
|
||||
|
||||
audio = record_path(job_id)
|
||||
if audio is None:
|
||||
raise HTTPException(status_code=404,
|
||||
detail="запись не сохранена, черновик собрать не из чего")
|
||||
|
||||
if state.get("status") != "building" and (build or not state):
|
||||
if state.get("status") == "failed":
|
||||
raise HTTPException(status_code=500, detail=state.get("error", "не собрался"))
|
||||
if state.get("status") != "building":
|
||||
audio = record_path(job_id)
|
||||
if audio is None:
|
||||
raise HTTPException(
|
||||
status_code=404, detail="запись не сохранена, черновик собрать не из чего")
|
||||
with _draft_lock:
|
||||
_drafts[job_id] = {"status": "building", "stage": "запускаю"}
|
||||
threading.Thread(target=_build_draft, args=(job_id, audio),
|
||||
name=f"draft-{job_id[:8]}", daemon=True).start()
|
||||
state = {"status": "building", "stage": "запускаю"}
|
||||
|
||||
if state.get("status") == "failed":
|
||||
raise HTTPException(status_code=500, detail=state.get("error", "не собрался"))
|
||||
# Страница сама перезапросит себя: сборка занимает минуты.
|
||||
return HTMLResponse(_draft_waiting(job_id, state.get("stage", ""), token),
|
||||
status_code=202)
|
||||
|
||||
+1
-1
@@ -1 +1 @@
|
||||
__version__ = "0.24.0"
|
||||
__version__ = "0.25.0"
|
||||
|
||||
@@ -469,3 +469,55 @@ class TestRecords:
|
||||
|
||||
def test_missing_record_gives_404(self, client):
|
||||
assert client.get("/v1/jobs/" + "0" * 32 + "/audio").status_code == 404
|
||||
|
||||
|
||||
class TestDraftEndpoint:
|
||||
"""Сборка идёт минуты: страница не должна ни висеть, ни врать про готовность."""
|
||||
|
||||
def test_without_record_gives_404(self, client):
|
||||
"""Собирать не из чего: запись не сохранена."""
|
||||
import app.main as m
|
||||
|
||||
m._drafts.clear()
|
||||
assert client.get("/v1/jobs/" + "0" * 32 + "/draft").status_code == 404
|
||||
|
||||
def test_building_answers_202_and_refreshes(self, client, tmp_path, monkeypatch):
|
||||
import app.main as m
|
||||
|
||||
m._drafts.clear()
|
||||
m._drafts["abc"] = {"status": "building", "stage": "распознаю"}
|
||||
response = client.get("/v1/jobs/abc/draft")
|
||||
assert response.status_code == 202
|
||||
assert "распознаю" in response.text
|
||||
assert "http-equiv=\"refresh\"" in response.text
|
||||
|
||||
def test_ready_draft_is_shown(self, client):
|
||||
import app.main as m
|
||||
|
||||
m._drafts.clear()
|
||||
m._drafts["abc"] = {"status": "ready", "draft": {
|
||||
"turns": [{"start": 0.0, "end": 2.0,
|
||||
"words": [{"text": "Привет", "variants": []},
|
||||
{"text": "мир", "variants": ["миру"]}]}],
|
||||
"duration_sec": 2.0, "models": 3, "total_words": 2,
|
||||
"agreed_words": 1, "gaps": []}}
|
||||
response = client.get("/v1/jobs/abc/draft")
|
||||
assert response.status_code == 200
|
||||
assert "<mark>" in response.text and "миру" in response.text
|
||||
|
||||
def test_failure_is_reported(self, client):
|
||||
"""Молчаливое зависание хуже честной ошибки."""
|
||||
import app.main as m
|
||||
|
||||
m._drafts.clear()
|
||||
m._drafts["abc"] = {"status": "failed", "error": "модель не загрузилась"}
|
||||
response = client.get("/v1/jobs/abc/draft")
|
||||
assert response.status_code == 500
|
||||
assert "модель не загрузилась" in response.text
|
||||
|
||||
def test_draft_needs_token(self, client):
|
||||
import app.main as m
|
||||
|
||||
m._drafts.clear()
|
||||
assert client.get("/v1/jobs/abc/draft",
|
||||
headers={"Authorization": ""}).status_code == 401
|
||||
|
||||
Reference in New Issue
Block a user