Spaces:
Running
Running
Download server.py from mocomoco-inc/AudioDecisionModel: direct link, hf CLI and curl.
- Browser
- Download file 15.5 kB
-
https://huggingface.co/spaces/mocomoco-inc/AudioDecisionModel/resolve/main/server.py
- Command line
-
hf download hf://spaces/mocomoco-inc/AudioDecisionModel/server.py
-
curl -L -o server.py https://huggingface.co/spaces/mocomoco-inc/AudioDecisionModel/resolve/main/server.py
15.5 kB
| """CPU-only API for the Mocovoice Audio Jev demo. | |
| The model files are local to the Space. Requests are never written to disk. | |
| """ | |
| from __future__ import annotations | |
| import asyncio | |
| import json | |
| import math | |
| import os | |
| import threading | |
| import time | |
| from pathlib import Path | |
| from typing import Any | |
| import numpy as np | |
| from fastapi import FastAPI, HTTPException, Request | |
| from fastapi.concurrency import run_in_threadpool | |
| from fastapi.responses import FileResponse, PlainTextResponse | |
| from pydantic import BaseModel, ValidationError | |
| import torch | |
| from server_runtime.typed_runtime import TypedAudioPredictor | |
| from server_runtime.asr_semantic import ( | |
| ASR_ADAPTATION_SHA256, | |
| ASR_VARIANT, | |
| QwenASRSemanticPredictor, | |
| fixed_semantic_task, | |
| ) | |
| from server_runtime.transcript_keyword import ( | |
| answer_keyword_questions, | |
| classify_keyword_question, | |
| ) | |
| from server_runtime.language_id import ( | |
| answer_language_questions, | |
| classify_language_question, | |
| ) | |
| from server_runtime.clap_sound_router import ( | |
| ClapSoundRouter, | |
| sound_answer, | |
| ) | |
| from server_runtime.learned_modality_router import ( | |
| LearnedModalityRouter, | |
| learned_fuse_answers, | |
| ) | |
| ROOT = Path(__file__).resolve().parent | |
| CHECKPOINT = ROOT / "server_runtime" / "checkpoint.pt" | |
| # A full 30-second JSON array of 480,000 float samples is around 10 MiB. | |
| MAX_BODY_BYTES = 12 * 1024 * 1024 | |
| class AnalyzeRequest(BaseModel): | |
| samples: list[float] | |
| context: str | |
| questions: dict[str, Any] | |
| class ModelBundle: | |
| def __init__(self) -> None: | |
| # This Space is deliberately CPU-only and serializes requests below. | |
| # Two threads match the target Space allocation and keep Qwen decoding | |
| # below real time on the measured short-audio CPU benchmark. | |
| torch.set_num_threads(max(1, min(8, int(os.environ.get("JEV_TORCH_THREADS", "2"))))) | |
| self.semantic = QwenASRSemanticPredictor(ROOT / "server_runtime" / "semantic_v8", device="cpu") | |
| self._predictor: TypedAudioPredictor | None = None | |
| self._predictor_lock = threading.Lock() | |
| # Checkpoint heads are local and small; the 615 MB pinned CLAP base is | |
| # loaded only after a supported non-verbal question reaches the server. | |
| self.sound = ClapSoundRouter(ROOT / "server_runtime") | |
| self.router = LearnedModalityRouter(ROOT / "models") | |
| def typed_predictor(self) -> TypedAudioPredictor: | |
| # The generic option scorer remains available for turn state and user | |
| # schemas. Avoid its second Qwen audio-tower copy until it is needed. | |
| if self._predictor is None: | |
| with self._predictor_lock: | |
| if self._predictor is None: | |
| self._predictor = TypedAudioPredictor(CHECKPOINT, device="cpu") | |
| return self._predictor | |
| _model: ModelBundle | None = None | |
| _load_lock = threading.Lock() | |
| _request_lock = asyncio.Lock() | |
| def get_model() -> ModelBundle: | |
| global _model | |
| if _model is None: | |
| with _load_lock: | |
| if _model is None: | |
| _model = ModelBundle() | |
| return _model | |
| def status() -> dict[str, Any]: | |
| bundle = _model | |
| return { | |
| "ready": bundle is not None, | |
| "backend": "server-cpu", | |
| "precision": "float32", | |
| "model_revision": "cloud-unified-qwen06-ja-adapt-semantic-v8-clap-learned-router-v7", | |
| "model": "Adapted Qwen3-ASR-0.6B and CLAP with learned sparse routing ensemble and option-weighted fusion", | |
| "asr_adaptation": {"candidate": ASR_VARIANT, "sha256": ASR_ADAPTATION_SHA256}, | |
| "on_device": False, | |
| "audio_sent_to_server": True, | |
| } | |
| def fail(message: str) -> None: | |
| raise ValueError(message) | |
| def options_for(question: Any) -> list[tuple[str, str]]: | |
| if not isinstance(question, dict) or not isinstance(question.get("instructions"), str) or not question["instructions"].strip(): | |
| fail("各質問にinstructions(質問文)を指定してください。") | |
| if len(question["instructions"]) > 2000: | |
| fail("質問文は2000文字以内にしてください。") | |
| kind = question.get("type") | |
| if kind == "noul": | |
| criteria = question.get("criteria", {"false": "いいえ。その条件を満たさない。", "true": "はい。その条件を満たす。"}) | |
| if not isinstance(criteria, dict) or set(criteria) != {"false", "true"}: | |
| fail("Noulのcriteriaにはfalseとtrueを指定してください。") | |
| pairs = [("false", criteria["false"]), ("true", criteria["true"])] | |
| elif kind == "choice": | |
| criteria = question.get("criteria") | |
| if not isinstance(criteria, dict): | |
| fail("Choiceのcriteriaは候補IDと説明のオブジェクトです。") | |
| pairs = list(criteria.items()) | |
| if not 2 <= len(pairs) <= 32: | |
| fail("Choiceの候補数は2〜32個です。") | |
| elif kind == "score": | |
| criteria = question.get("criteria") | |
| if not isinstance(criteria, list) or not 2 <= len(criteria) <= 10: | |
| fail("Scoreのcriteriaは2〜10段階の説明を順に並べた配列です。") | |
| pairs = [(str(i), value) for i, value in enumerate(criteria)] | |
| else: | |
| fail("typeはnoul・choice・scoreのいずれかです。") | |
| if any(not key or not isinstance(value, str) or not value.strip() or len(value) > 1000 for key, value in pairs): | |
| fail("各候補には空でない説明を1000文字以内で指定してください。") | |
| return pairs | |
| def waveform_windows(samples: np.ndarray) -> list[np.ndarray]: | |
| last = max(0, len(samples) - 80000) | |
| starts = list(range(0, last + 1, 40000)) | |
| if starts[-1] != last: | |
| starts.append(last) | |
| return [samples[start : min(start + 80000, len(samples))] for start in starts] | |
| def normalize(wave: np.ndarray) -> np.ndarray: | |
| # Match JS accumulation semantics closely; both use observed-population variance. | |
| mean = sum(float(value) for value in wave) / len(wave) | |
| variance = sum((float(value) - mean) ** 2 for value in wave) / len(wave) | |
| return np.asarray([(float(value) - mean) / math.sqrt(variance + 1e-7) for value in wave], dtype=np.float32) | |
| def analyze(payload: AnalyzeRequest) -> dict[str, Any]: | |
| began = time.perf_counter() | |
| samples = np.asarray(payload.samples, dtype=np.float32) | |
| if samples.ndim != 1 or len(samples) < 4000 or len(samples) > 480000 or not np.isfinite(samples).all(): | |
| fail("音声は0.25秒〜30秒の範囲で入力してください。") | |
| if len(payload.context) > 2000: | |
| fail("文脈は2000文字以内にしてください。") | |
| entries = list(payload.questions.items()) | |
| if not 1 <= len(entries) <= 16: | |
| fail("質問数は1〜16個にしてください。") | |
| # Validate with the browser-visible Japanese error contract before model work. | |
| for _, question in entries: | |
| options_for(question) | |
| bundle = get_model() | |
| # Question names and lexical aliases never enter this gate. ModernBERT | |
| # embeddings produce continuous speech/non-verbal/joint weights, followed | |
| # by sparse top-1 or two-expert execution. | |
| routes = bundle.router.route_all(payload.context, payload.questions) | |
| language_questions = {name: question for name, question in payload.questions.items() | |
| if routes[name].mode in ("language", "joint")} | |
| language_id_questions = {name: question for name, question in language_questions.items() | |
| if classify_language_question(question) is not None} | |
| fixed_questions = {name: question for name, question in language_questions.items() | |
| if name not in language_id_questions and fixed_semantic_task(question) is not None} | |
| keyword_questions = {name: question for name, question in language_questions.items() | |
| if name not in language_id_questions and fixed_semantic_task(question) is None | |
| and classify_keyword_question(question) is not None} | |
| generic_questions = {name: question for name, question in language_questions.items() | |
| if name not in language_id_questions and name not in fixed_questions | |
| and name not in keyword_questions} | |
| needs_sound = any(selected.mode in ("sound", "joint") for selected in routes.values()) | |
| sound_scores = bundle.sound.score(samples) if needs_sound else None | |
| result = {"transcript": None, "answers": {}, "action_executable": False, | |
| "evidence": {"audio_duration_ms": len(samples) / 16, "sample_rate_hz": 16_000, | |
| "language_experts": {}}, "latency_ms": {}} | |
| auto_transcription = None | |
| if language_id_questions: | |
| asr_began = time.perf_counter() | |
| auto_transcription = bundle.semantic.transcribe_result(samples, detect_language=True) | |
| result["latency_ms"]["language_asr"] = (time.perf_counter() - asr_began) * 1000 | |
| result["transcript"] = auto_transcription.text | |
| result["answers"].update(answer_language_questions( | |
| auto_transcription.language, language_id_questions | |
| )) | |
| result["evidence"]["language_experts"]["native_language_id"] = { | |
| "model": f"Qwen/Qwen3-ASR-0.6B + {ASR_VARIANT}", | |
| "detected_language": auto_transcription.language or "unknown", | |
| "decoder": "one automatic-language ASR decode shared with transcript tasks", | |
| "adaptation_sha256": ASR_ADAPTATION_SHA256, | |
| "tasks": {name: classify_language_question(question).task | |
| for name, question in language_id_questions.items()}, | |
| } | |
| if fixed_questions: | |
| fixed_result = bundle.semantic.predict( | |
| samples, payload.context, fixed_questions, transcription=auto_transcription | |
| ) | |
| result["transcript"] = fixed_result["transcript"] | |
| result["answers"].update(fixed_result["answers"]) | |
| result["evidence"]["language_experts"]["fixed_semantics"] = fixed_result["evidence"] | |
| result["latency_ms"].update({f"fixed_{key}": value for key, value in fixed_result["latency_ms"].items()}) | |
| if keyword_questions: | |
| if result["transcript"] is None: | |
| asr_began = time.perf_counter() | |
| result["transcript"] = bundle.semantic.transcribe(samples) | |
| result["latency_ms"]["keyword_asr"] = (time.perf_counter() - asr_began) * 1000 | |
| result["answers"].update(answer_keyword_questions(result["transcript"], keyword_questions)) | |
| result["evidence"]["language_experts"]["transcript_keyword"] = { | |
| "model": f"Qwen/Qwen3-ASR-0.6B + {ASR_VARIANT}", | |
| "selection": "clean-v3 validation; frozen surface-reading-lemma matcher, threshold 0.60, minimum fuzzy length 2", | |
| "adaptation_sha256": ASR_ADAPTATION_SHA256, | |
| "tasks": {name: classify_keyword_question(question).task | |
| for name, question in keyword_questions.items()}, | |
| } | |
| if generic_questions: | |
| generic_result = bundle.typed_predictor().predict(samples, payload.context, generic_questions) | |
| result["answers"].update(generic_result["answers"]) | |
| result["evidence"]["language_experts"]["generic_typed"] = generic_result["evidence"] | |
| result["latency_ms"].update({f"generic_{key}": value for key, value in generic_result["latency_ms"].items()}) | |
| for name, question in payload.questions.items(): | |
| selected = routes[name] | |
| if selected.mode == "sound": | |
| result["answers"][name] = sound_answer(question, sound_scores) | |
| elif selected.mode == "joint": | |
| acoustic_answer = sound_answer(question, sound_scores) | |
| result["answers"][name] = learned_fuse_answers( | |
| question, result["answers"][name], acoustic_answer, selected | |
| ) | |
| result["answers"] = {name: result["answers"][name] for name in payload.questions} | |
| result.setdefault("evidence", {})["question_routing"] = {name: selected.public() for name, selected in routes.items()} | |
| result.setdefault("latency_ms", {})["total"] = (time.perf_counter() - began) * 1000 | |
| result["model"] = "mocovoice-audio-jev-cloud-unified-qwen06-ja-adapt-semantic-v8-clap-learned-router-v7" | |
| result["execution"] = status() | |
| result["warnings"] = [] | |
| result["scope"] = "研究用試作。未知タスクへの正確さと確率校正は未検証。" | |
| return result | |
| app = FastAPI(docs_url=None, redoc_url=None, openapi_url=None) | |
| async def load_models_at_startup() -> None: | |
| # The status endpoint is usable as a readiness probe, so load before serving. | |
| await run_in_threadpool(get_model) | |
| async def read_json_limited(request: Request) -> dict[str, Any]: | |
| content_length = request.headers.get("content-length") | |
| if content_length is not None: | |
| try: | |
| if int(content_length) > MAX_BODY_BYTES: | |
| raise HTTPException(status_code=413, detail="Request body must be 12 MiB or smaller.") | |
| except ValueError as error: | |
| raise HTTPException(status_code=400, detail="Invalid Content-Length.") from error | |
| chunks: list[bytes] = [] | |
| received = 0 | |
| async for chunk in request.stream(): | |
| received += len(chunk) | |
| if received > MAX_BODY_BYTES: | |
| raise HTTPException(status_code=413, detail="Request body must be 12 MiB or smaller.") | |
| chunks.append(chunk) | |
| try: | |
| body = json.loads(b"".join(chunks)) | |
| except (UnicodeDecodeError, json.JSONDecodeError) as error: | |
| raise HTTPException(status_code=400, detail="Request body must be valid JSON.") from error | |
| if not isinstance(body, dict): | |
| raise HTTPException(status_code=400, detail="Request body must be a JSON object.") | |
| return body | |
| async def api_status() -> dict[str, Any]: | |
| return status() | |
| async def api_analyze(request: Request) -> dict[str, Any]: | |
| try: | |
| payload = AnalyzeRequest.model_validate(await read_json_limited(request)) | |
| except ValidationError as error: | |
| raise HTTPException(status_code=400, detail="Invalid analyze request.") from error | |
| async with _request_lock: | |
| try: | |
| return await run_in_threadpool(analyze, payload) | |
| except ValueError as error: | |
| raise HTTPException(status_code=400, detail=str(error)) from error | |
| _EXACT_STATIC = { | |
| "index.html", "styles.css", "ui.js", "inference.js", "inference-worker.js", | |
| "audio-recorder.js", "audio-player.js", "README.md", "mocomoco-inc-logo.svg", "favicon.svg", | |
| } | |
| _STATIC_PREFIXES = ("models/", "vendor/", "fonts/", "examples/", "licenses/") | |
| async def home() -> FileResponse: | |
| return FileResponse(ROOT / "index.html", headers={"Cache-Control": "private, no-store"}) | |
| async def static_asset(asset_path: str) -> FileResponse: | |
| """Serve only browser assets; source and server checkpoints are never files.""" | |
| candidate = (ROOT / asset_path).resolve() | |
| try: | |
| relative = candidate.relative_to(ROOT).as_posix() | |
| except ValueError: | |
| return PlainTextResponse("Not found.", status_code=404) | |
| # Authorize the resolved path, never the untrusted URL string. This makes | |
| # ``models/../server.py`` and encoded dot-segment variants unavailable. | |
| allowed = relative in _EXACT_STATIC or relative.startswith(_STATIC_PREFIXES) | |
| if not allowed or not candidate.is_file(): | |
| return PlainTextResponse("Not found.", status_code=404) | |
| return FileResponse(candidate) | |