mirror of
https://github.com/mudler/LocalAI.git
synced 2026-09-09 04:29:35 -04:00
Compare commits
231
Commits
No files matched your search
@@ -16,8 +16,7 @@ side (`pkg/oci/cosignverify` plus the gallery YAML).
|
||||
per-arch manifest before checking signatures.
|
||||
- **Storage:** Signatures are written as OCI 1.1 referrers
|
||||
(`--registry-referrers-mode=oci-1-1`) in the new Sigstore bundle format
|
||||
(current cosign releases do this by default; no `--new-bundle-format`
|
||||
flag). No `:sha256-<hex>.sig` tag clutter.
|
||||
(`--new-bundle-format`). No `:sha256-<hex>.sig` tag clutter.
|
||||
- **Consumer:** `pkg/oci/cosignverify` discovers the bundle via the
|
||||
referrers API, hands it to `sigstore-go`, and verifies it against the
|
||||
policy declared in the gallery YAML (`Gallery.Verification`).
|
||||
@@ -34,14 +33,15 @@ to sign. The job needs:
|
||||
|
||||
- `permissions: { id-token: write, contents: read }` at the job level so
|
||||
the runner can exchange its GitHub OIDC token for a Fulcio cert.
|
||||
- `sigstore/cosign-installer@v3` step (current cosign releases already
|
||||
default to the new bundle format).
|
||||
- `sigstore/cosign-installer@v3` step (the pinned cosign v2 release needs
|
||||
`--new-bundle-format` explicitly).
|
||||
- After each `docker buildx imagetools create`, resolve the resulting
|
||||
list digest with `docker buildx imagetools inspect <tag> --format
|
||||
'{{.Manifest.Digest}}'` and sign:
|
||||
|
||||
```sh
|
||||
cosign sign --yes --recursive \
|
||||
--new-bundle-format \
|
||||
--registry-referrers-mode=oci-1-1 \
|
||||
"${REGISTRY_REPO}@${DIGEST}"
|
||||
```
|
||||
@@ -67,10 +67,12 @@ entry (`backend/index.yaml`):
|
||||
|
||||
```yaml
|
||||
- name: localai
|
||||
url: github:mudler/LocalAI/backend/index.yaml@master
|
||||
url: https://index.localai.io/backends
|
||||
mirrors:
|
||||
- github:mudler/LocalAI/backend/index.yaml@master
|
||||
verification:
|
||||
issuer: "https://token.actions.githubusercontent.com"
|
||||
identity_regex: "^https://github\\.com/mudler/LocalAI/\\.github/workflows/backend_merge\\.yml@refs/heads/master$"
|
||||
identity_regex: "^https://github\\.com/mudler/LocalAI/\\.github/workflows/backend_merge\\.yml@refs/(heads/master|tags/.+)$"
|
||||
# Optional revocation cutoff; advance during incident response.
|
||||
# not_before: "2026-06-01T00:00:00Z"
|
||||
```
|
||||
|
||||
@@ -198,7 +198,7 @@ Two properties this relies on:
|
||||
|
||||
The same reasoning applies to master pushes, and the volume is larger there: on 2026-07-30, **12 of the 23 queued `image.yml` runs** were commits like "add 1 new model to gallery" or a docs fix, each rebuilding all 18 container images.
|
||||
|
||||
`image.yml` now has a `changes` job that decides once whether the push can affect any image; the other 11 jobs carry `needs: changes` plus an `if:` on its output. Verified against the shipped `Dockerfile`: the final stage copies only `entrypoint.sh`, `healthcheck.sh` and the `local-ai` binary, there is no `go:embed` of `gallery/` or `docs/`, and the gallery is fetched at runtime from `github:mudler/LocalAI/gallery/index.yaml@master`. A gallery-only commit therefore produces byte-identical images, and the gallery change reaches users through GitHub immediately whether or not an image is rebuilt.
|
||||
`image.yml` now has a `changes` job that decides once whether the push can affect any image; the other 11 jobs carry `needs: changes` plus an `if:` on its output. Verified against the shipped `Dockerfile`: the final stage copies only `entrypoint.sh`, `healthcheck.sh` and the `local-ai` binary, there is no `go:embed` of `gallery/` or `docs/`, and the gallery is fetched at runtime from `https://index.localai.io/models` (a caching mirror of `gallery/index.yaml` on master, with `github:mudler/LocalAI/gallery/index.yaml@master` as the fallback mirror). A gallery-only commit therefore produces byte-identical images, and the gallery change reaches users over the network immediately whether or not an image is rebuilt.
|
||||
|
||||
Two properties to preserve if you touch it:
|
||||
|
||||
|
||||
@@ -8,8 +8,15 @@ build_type=${2-}
|
||||
# ggml-cpu/arch/x86/repack.cpp at -march=sapphirerapids: the job sits on that one
|
||||
# translation unit until GitHub kills it at 6h. gcc builds the same file in
|
||||
# seconds, so only the SYCL images have to give up the CPU variant matrix.
|
||||
#
|
||||
# ROCm runs out of the same 6h budget for a different reason: volume, not a
|
||||
# stall. hipcc compiles ggml's HIP kernels once per entry in AMDGPU_TARGETS,
|
||||
# which is eleven architectures (gfx908 through gfx1201), and the CPU variant
|
||||
# matrix lands on top of that. The job built in 2h27m before it was added and
|
||||
# has been killed at exactly 6h00m on every run since, so no ROCm llama-cpp
|
||||
# image has been published since 2026-08-01.
|
||||
case "$build_type" in
|
||||
sycl*)
|
||||
sycl*|hipblas*)
|
||||
echo llama-cpp-fallback
|
||||
exit 0
|
||||
;;
|
||||
|
||||
@@ -59,7 +59,9 @@ backend/rust/*/target
|
||||
backend-images
|
||||
local-backends
|
||||
local-ai
|
||||
.claude
|
||||
.crush
|
||||
.tools
|
||||
protoc
|
||||
tests
|
||||
|
||||
|
||||
+146
-4
@@ -860,6 +860,19 @@ include:
|
||||
dockerfile: "./backend/Dockerfile.golang"
|
||||
context: "./"
|
||||
ubuntu-version: '2404'
|
||||
- build-type: 'cublas'
|
||||
cuda-major-version: "12"
|
||||
cuda-minor-version: "8"
|
||||
platforms: 'linux/amd64'
|
||||
tag-latest: 'auto'
|
||||
tag-suffix: '-gpu-nvidia-cuda-12-nemo-speech-cpp'
|
||||
runs-on: 'ubuntu-latest'
|
||||
base-image: "ubuntu:24.04"
|
||||
skip-drivers: 'false'
|
||||
backend: "nemo-speech-cpp"
|
||||
dockerfile: "./backend/Dockerfile.golang"
|
||||
context: "./"
|
||||
ubuntu-version: '2404'
|
||||
- build-type: 'cublas'
|
||||
cuda-major-version: "12"
|
||||
cuda-minor-version: "8"
|
||||
@@ -1911,6 +1924,19 @@ include:
|
||||
dockerfile: "./backend/Dockerfile.golang"
|
||||
context: "./"
|
||||
ubuntu-version: '2404'
|
||||
- build-type: 'cublas'
|
||||
cuda-major-version: "13"
|
||||
cuda-minor-version: "0"
|
||||
platforms: 'linux/amd64'
|
||||
tag-latest: 'auto'
|
||||
tag-suffix: '-gpu-nvidia-cuda-13-nemo-speech-cpp'
|
||||
runs-on: 'ubuntu-latest'
|
||||
base-image: "ubuntu:24.04"
|
||||
skip-drivers: 'false'
|
||||
backend: "nemo-speech-cpp"
|
||||
dockerfile: "./backend/Dockerfile.golang"
|
||||
context: "./"
|
||||
ubuntu-version: '2404'
|
||||
- build-type: 'cublas'
|
||||
cuda-major-version: "13"
|
||||
cuda-minor-version: "0"
|
||||
@@ -1963,6 +1989,24 @@ include:
|
||||
backend: "parakeet-cpp"
|
||||
dockerfile: "./backend/Dockerfile.golang"
|
||||
context: "./"
|
||||
# The CUDA-13 counterpart to the JetPack r36.4.0 row in the nemo-speech-cpp
|
||||
# block below. A Jetson whose CUDA 13 runtime is present reports the
|
||||
# nvidia-l4t-cuda-13 capability, and pointing that key at the JetPack image
|
||||
# would hand it a ggml linked against CUDA 12 whose libcudart.so.12 is not
|
||||
# there to dlopen. Same base and runner as the parakeet-cpp row above.
|
||||
- build-type: 'cublas'
|
||||
cuda-major-version: "13"
|
||||
cuda-minor-version: "0"
|
||||
platforms: 'linux/arm64'
|
||||
skip-drivers: 'false'
|
||||
tag-latest: 'auto'
|
||||
tag-suffix: '-nvidia-l4t-cuda-13-arm64-nemo-speech-cpp'
|
||||
base-image: "ubuntu:24.04"
|
||||
ubuntu-version: '2404'
|
||||
runs-on: 'ubuntu-24.04-arm'
|
||||
backend: "nemo-speech-cpp"
|
||||
dockerfile: "./backend/Dockerfile.golang"
|
||||
context: "./"
|
||||
- build-type: 'cublas'
|
||||
cuda-major-version: "13"
|
||||
cuda-minor-version: "0"
|
||||
@@ -2540,7 +2584,7 @@ include:
|
||||
tag-latest: 'auto'
|
||||
tag-suffix: '-gpu-intel-vllm'
|
||||
runs-on: 'ubuntu-latest'
|
||||
base-image: "intel/oneapi-basekit:2025.3.0-0-devel-ubuntu24.04"
|
||||
base-image: "intel/oneapi-basekit:2025.3.2-0-devel-ubuntu24.04"
|
||||
skip-drivers: 'false'
|
||||
backend: "vllm"
|
||||
dockerfile: "./backend/Dockerfile.python"
|
||||
@@ -3164,9 +3208,10 @@ include:
|
||||
# consumed. Same reason CUDA needs its toolkit in base-image rather than in a
|
||||
# builder image: this is the ds4 shape, not the llama-cpp one.
|
||||
#
|
||||
# No ROCm entry: upstream has no HIP configuration. No CUDA arm64 or L4T
|
||||
# entry: upstream documents and validates CUDA on x86 only. Darwin/Metal is in
|
||||
# the includeDarwin matrix below, built by scripts/build/audio-cpp-darwin.sh.
|
||||
# ROCm uses upstream's HIP backend and the project-wide ROCm 7.2.1 base. No
|
||||
# CUDA arm64 or L4T entry: upstream documents and validates CUDA on x86 only.
|
||||
# Darwin/Metal is in the includeDarwin matrix below, built by
|
||||
# scripts/build/audio-cpp-darwin.sh.
|
||||
#
|
||||
# No vulkan entry either, though Dockerfile.audio-cpp and the backend Makefile
|
||||
# both handle BUILD_TYPE=vulkan for local builds. Every other vulkan backend
|
||||
@@ -3234,6 +3279,19 @@ include:
|
||||
dockerfile: "./backend/Dockerfile.audio-cpp"
|
||||
context: "./"
|
||||
ubuntu-version: '2404'
|
||||
- build-type: 'hipblas'
|
||||
cuda-major-version: ""
|
||||
cuda-minor-version: ""
|
||||
platforms: 'linux/amd64'
|
||||
tag-latest: 'auto'
|
||||
tag-suffix: '-gpu-rocm-hipblas-audio-cpp'
|
||||
runs-on: 'ubuntu-latest'
|
||||
base-image: "rocm/dev-ubuntu-24.04:7.2.1"
|
||||
skip-drivers: 'false'
|
||||
backend: "audio-cpp"
|
||||
dockerfile: "./backend/Dockerfile.audio-cpp"
|
||||
context: "./"
|
||||
ubuntu-version: '2404'
|
||||
- build-type: ''
|
||||
cuda-major-version: ""
|
||||
cuda-minor-version: ""
|
||||
@@ -4183,6 +4241,86 @@ include:
|
||||
dockerfile: "./backend/Dockerfile.golang"
|
||||
context: "./"
|
||||
ubuntu-version: '2404'
|
||||
# nemo-speech-cpp
|
||||
#
|
||||
# No hipblas and no sycl rows, unlike the parakeet-cpp block above: upstream
|
||||
# NeMo-Speech.cpp builds ggml with CUDA, Vulkan or Metal only, so a ROCm or
|
||||
# SYCL image would be a CPU build wearing a GPU tag.
|
||||
#
|
||||
# cpu and vulkan are per-arch pairs sharing one tag-suffix, so
|
||||
# backend-merge-jobs assembles a multi-arch manifest from the two digests.
|
||||
# The arm64 legs are not redundant with the Jetson image below: an ARM server
|
||||
# with no NVIDIA GPU reports the "default" capability and would otherwise pull
|
||||
# an amd64-only manifest.
|
||||
- build-type: ''
|
||||
cuda-major-version: ""
|
||||
cuda-minor-version: ""
|
||||
platforms: 'linux/amd64'
|
||||
platform-tag: 'amd64'
|
||||
tag-latest: 'auto'
|
||||
tag-suffix: '-cpu-nemo-speech-cpp'
|
||||
runs-on: 'ubuntu-latest'
|
||||
base-image: "ubuntu:24.04"
|
||||
skip-drivers: 'false'
|
||||
backend: "nemo-speech-cpp"
|
||||
dockerfile: "./backend/Dockerfile.golang"
|
||||
context: "./"
|
||||
ubuntu-version: '2404'
|
||||
- build-type: ''
|
||||
cuda-major-version: ""
|
||||
cuda-minor-version: ""
|
||||
platforms: 'linux/arm64'
|
||||
platform-tag: 'arm64'
|
||||
tag-latest: 'auto'
|
||||
tag-suffix: '-cpu-nemo-speech-cpp'
|
||||
runs-on: 'ubuntu-24.04-arm'
|
||||
base-image: "ubuntu:24.04"
|
||||
skip-drivers: 'false'
|
||||
backend: "nemo-speech-cpp"
|
||||
dockerfile: "./backend/Dockerfile.golang"
|
||||
context: "./"
|
||||
ubuntu-version: '2404'
|
||||
- build-type: 'vulkan'
|
||||
cuda-major-version: ""
|
||||
cuda-minor-version: ""
|
||||
platforms: 'linux/amd64'
|
||||
platform-tag: 'amd64'
|
||||
tag-latest: 'auto'
|
||||
tag-suffix: '-gpu-vulkan-nemo-speech-cpp'
|
||||
runs-on: 'ubuntu-latest'
|
||||
base-image: "ubuntu:24.04"
|
||||
skip-drivers: 'false'
|
||||
backend: "nemo-speech-cpp"
|
||||
dockerfile: "./backend/Dockerfile.golang"
|
||||
context: "./"
|
||||
ubuntu-version: '2404'
|
||||
- build-type: 'vulkan'
|
||||
cuda-major-version: ""
|
||||
cuda-minor-version: ""
|
||||
platforms: 'linux/arm64'
|
||||
platform-tag: 'arm64'
|
||||
tag-latest: 'auto'
|
||||
tag-suffix: '-gpu-vulkan-nemo-speech-cpp'
|
||||
runs-on: 'ubuntu-24.04-arm'
|
||||
base-image: "ubuntu:24.04"
|
||||
skip-drivers: 'false'
|
||||
backend: "nemo-speech-cpp"
|
||||
dockerfile: "./backend/Dockerfile.golang"
|
||||
context: "./"
|
||||
ubuntu-version: '2404'
|
||||
- build-type: 'cublas'
|
||||
cuda-major-version: "12"
|
||||
cuda-minor-version: "0"
|
||||
platforms: 'linux/arm64'
|
||||
skip-drivers: 'false'
|
||||
tag-latest: 'auto'
|
||||
tag-suffix: '-nvidia-l4t-arm64-nemo-speech-cpp'
|
||||
base-image: "nvcr.io/nvidia/l4t-jetpack:r36.4.0"
|
||||
runs-on: 'ubuntu-24.04-arm'
|
||||
backend: "nemo-speech-cpp"
|
||||
dockerfile: "./backend/Dockerfile.golang"
|
||||
context: "./"
|
||||
ubuntu-version: '2204'
|
||||
# moss-transcribe-cpp
|
||||
- build-type: ''
|
||||
cuda-major-version: ""
|
||||
@@ -6226,6 +6364,10 @@ includeDarwin:
|
||||
tag-suffix: "-metal-darwin-arm64-moss-transcribe-cpp"
|
||||
build-type: "metal"
|
||||
lang: "go"
|
||||
- backend: "nemo-speech-cpp"
|
||||
tag-suffix: "-metal-darwin-arm64-nemo-speech-cpp"
|
||||
build-type: "metal"
|
||||
lang: "go"
|
||||
- backend: "ced"
|
||||
tag-suffix: "-metal-darwin-arm64-ced"
|
||||
build-type: "metal"
|
||||
|
||||
@@ -3,7 +3,7 @@
|
||||
# darwin (Apple Silicon) install path. The macOS/Metal build
|
||||
# (backend/python/vllm/install.sh, Darwin branch) installs vllm-metal, which is
|
||||
# version-locked to a specific vLLM source release. install.sh derives that vLLM
|
||||
# version at build time from vllm-metal's own installer (`vllm_v=`) at the pinned
|
||||
# version at build time from vllm-metal's own installer at the pinned
|
||||
# tag, so there is only ONE value to bump here -- mirroring bump_vllm_wheel.sh,
|
||||
# which bumps the Linux cu130 wheel pin.
|
||||
#
|
||||
@@ -32,10 +32,10 @@ LATEST_TAG=$(gh_curl -H "Accept: application/vnd.github+json" \
|
||||
# The coupled vLLM source version lives in vllm-metal's installer at that tag.
|
||||
NEW_VLLM_VERSION=$(gh_curl \
|
||||
"https://raw.githubusercontent.com/$REPO/$LATEST_TAG/install.sh" \
|
||||
| grep -oE 'vllm_v="[0-9]+\.[0-9]+\.[0-9]+"' | head -1 | cut -d'"' -f2)
|
||||
| "$(dirname "${BASH_SOURCE[0]}")/../scripts/lib/extract-vllm-metal-version.sh")
|
||||
|
||||
if [ -z "$LATEST_TAG" ] || [ -z "$NEW_VLLM_VERSION" ]; then
|
||||
echo "Could not resolve vllm-metal tag ($LATEST_TAG) or its vllm_v ($NEW_VLLM_VERSION)." >&2
|
||||
echo "Could not resolve vllm-metal tag ($LATEST_TAG) or its vLLM version ($NEW_VLLM_VERSION)." >&2
|
||||
exit 1
|
||||
fi
|
||||
|
||||
|
||||
Executable
+44
@@ -0,0 +1,44 @@
|
||||
#!/usr/bin/env bash
|
||||
set -euo pipefail
|
||||
|
||||
python3 - <<'PY'
|
||||
from pathlib import Path
|
||||
|
||||
home = Path("website/layouts/index.html").read_text()
|
||||
css = Path("website/static/css/site.css").read_text()
|
||||
install = Path("docs/content/getting-started/install.md").read_text()
|
||||
containers = Path("docs/content/getting-started/containers.md").read_text()
|
||||
|
||||
def require(condition, message):
|
||||
if not condition:
|
||||
raise SystemExit(f"FAIL: {message}")
|
||||
|
||||
require("Drop-in replacement for most upstream APIs." in home,
|
||||
"homepage must use the requested drop-in API heading")
|
||||
require("Everything else plugs into LocalAI." not in home,
|
||||
"old runtime heading must be removed")
|
||||
require("When the engine we need" not in home,
|
||||
"hero must describe user outcomes instead of team implementation")
|
||||
require('href="mailto:contact@localai.io"' in home and "business" in home.lower(),
|
||||
"homepage must provide a direct business contact action")
|
||||
require(home.index('id="localai"') < home.index('id="proof-quotes"') < home.index('id="mission"'),
|
||||
"headline testimonials must directly follow the runtime section")
|
||||
require(home.count('id="proof-quotes"') == 1,
|
||||
"headline testimonials must appear exactly once")
|
||||
require('id="engines"' not in home and "Engines we build" not in home,
|
||||
"homepage engine showcase must be removed")
|
||||
require('href="/docs/installation/index.html"' in home,
|
||||
"installation guide action must use the direct installation URL")
|
||||
require('<iframe' in install and "youtube.com/embed/cMVNnlqwfw4" in install,
|
||||
"installation page must embed the walkthrough video")
|
||||
require("## Quick Start" not in install,
|
||||
"installation landing page must not duplicate Quick Start")
|
||||
for text in ("CUDA 12", "CUDA 13", "ROCm", "Intel", "Jetson", "Vulkan", "fallback"):
|
||||
require(text.lower() in containers.lower(), f"GPU chooser must explain {text}")
|
||||
require('class="sn__e"><a href="https://github.com/mudler/parakeet.cpp">parakeet.cpp</a>' in home,
|
||||
"capability engine names must link to their repositories")
|
||||
require(".pane{min-height:" in css.replace(" ", ""),
|
||||
"all installation panes must have a fixed minimum height")
|
||||
|
||||
print("website review 143 source checks passed")
|
||||
PY
|
||||
@@ -71,8 +71,8 @@ jobs:
|
||||
|
||||
# cosign signs each pushed manifest list with --recursive so the
|
||||
# index and every per-arch entry get an attached Sigstore bundle.
|
||||
# Recent cosign releases always emit the new bundle format, so
|
||||
# there's no extra CLI flag to opt into it.
|
||||
# Cosign v2.4.1 emits the current bundle format by default; the
|
||||
# verifier discovers those bundles through OCI 1.1 referrers.
|
||||
- name: Install cosign
|
||||
if: github.event_name != 'pull_request'
|
||||
uses: sigstore/cosign-installer@v3
|
||||
|
||||
@@ -62,6 +62,10 @@ jobs:
|
||||
variable: "MOSS_VERSION"
|
||||
branch: "master"
|
||||
file: "backend/go/moss-transcribe-cpp/Makefile"
|
||||
- repository: "NVIDIA/NeMo-Speech.cpp"
|
||||
variable: "NEMO_SPEECH_VERSION"
|
||||
branch: "main"
|
||||
file: "backend/go/nemo-speech-cpp/Makefile"
|
||||
- repository: "localai-org/ced.cpp"
|
||||
variable: "CED_VERSION"
|
||||
branch: "main"
|
||||
|
||||
@@ -14,6 +14,7 @@ on:
|
||||
|
||||
permissions:
|
||||
contents: write
|
||||
pull-requests: write
|
||||
|
||||
concurrency:
|
||||
group: refresh-site-counters
|
||||
@@ -23,22 +24,32 @@ jobs:
|
||||
refresh:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: actions/checkout@v7
|
||||
|
||||
- name: Read the counts off the GitHub API
|
||||
env:
|
||||
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
run: ./.github/ci/refresh-site-counters.sh
|
||||
|
||||
- name: Commit only if something moved
|
||||
- name: Show changes
|
||||
run: |
|
||||
if git diff --quiet -- website/data/stats.yaml; then
|
||||
echo "counters unchanged, nothing to commit"
|
||||
exit 0
|
||||
echo "counters unchanged"
|
||||
else
|
||||
git diff --unified=0 -- website/data/stats.yaml
|
||||
fi
|
||||
git diff --unified=0 -- website/data/stats.yaml
|
||||
git config user.name "github-actions[bot]"
|
||||
git config user.email "41898282+github-actions[bot]@users.noreply.github.com"
|
||||
git add website/data/stats.yaml
|
||||
git commit -m "chore(website): refresh the counters"
|
||||
git push
|
||||
|
||||
- name: Create pull request when counters moved
|
||||
uses: peter-evans/create-pull-request@v8
|
||||
with:
|
||||
token: ${{ secrets.UPDATE_BOT_TOKEN }}
|
||||
push-to-fork: ci-forks/LocalAI
|
||||
commit-message: "chore(website): refresh the counters"
|
||||
title: "chore(website): refresh the counters"
|
||||
body: |
|
||||
Weekly refresh of the landing-page counters from the GitHub API.
|
||||
|
||||
This PR was created automatically by the `refresh-site-counters` workflow.
|
||||
branch: update/site-counters
|
||||
delete-branch: true
|
||||
labels: automated
|
||||
@@ -11,7 +11,7 @@ jobs:
|
||||
if: github.repository == 'mudler/LocalAI'
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/stale@1e223db275d687790206a7acac4d1a11bd6fe629 # v9
|
||||
- uses: actions/stale@4391f3da665fdf50b6810c1a66712fb9ba21aa93 # v9
|
||||
with:
|
||||
stale-issue-message: 'This issue is stale because it has been open 90 days with no activity. Remove stale label or comment or this will be closed in 5 days.'
|
||||
stale-pr-message: 'This PR is stale because it has been open 90 days with no activity. Remove stale label or comment or this will be closed in 10 days.'
|
||||
|
||||
@@ -50,6 +50,7 @@ jobs:
|
||||
sherpa-onnx: ${{ steps.detect.outputs.sherpa-onnx }}
|
||||
whisper: ${{ steps.detect.outputs.whisper }}
|
||||
parakeet-cpp: ${{ steps.detect.outputs.parakeet-cpp }}
|
||||
nemo-speech-cpp: ${{ steps.detect.outputs.nemo-speech-cpp }}
|
||||
steps:
|
||||
- name: Checkout repository
|
||||
uses: actions/checkout@v7
|
||||
@@ -525,6 +526,7 @@ jobs:
|
||||
- name: Build llama-cpp backend image and run gRPC e2e tests
|
||||
run: |
|
||||
make test-extra-backend-llama-cpp
|
||||
make test-extra-backend-llama-cpp-embeddings
|
||||
tests-llama-cpp-grpc-transcription:
|
||||
needs: detect-changes
|
||||
if: needs.detect-changes.outputs.llama-cpp == 'true' || needs.detect-changes.outputs.run-all == 'true'
|
||||
@@ -900,6 +902,57 @@ jobs:
|
||||
- name: Test magpie-tts-cpp
|
||||
run: |
|
||||
make --jobs=5 --output-sync=target -C backend/go/magpie-tts-cpp test
|
||||
# Per-backend unit suite for nemo-speech-cpp. This job exists for one reason
|
||||
# above all: abi_test.go asserts the size and field offsets of every Go mirror
|
||||
# struct against the C ABI it is dlopened into. Those assertions are the only
|
||||
# thing standing between a purego symbol rename or an upstream header change
|
||||
# and silent memory corruption at run time, and they are worthless unless
|
||||
# something executes them. `make -C backend/go/nemo-speech-cpp test` sets
|
||||
# NEMO_SPEECH_REQUIRE_LIBS=1, which turns "library missing" from a skip into a
|
||||
# failure, so this job cannot report green having checked nothing.
|
||||
#
|
||||
# The backend Makefile's `test` target depends on `stage-libs`, so it clones
|
||||
# upstream at the pinned SHA and builds the native runtime itself. There is no
|
||||
# separate build step for that reason, and no model download: the specs are
|
||||
# ABI and pure-Go only.
|
||||
#
|
||||
# WITH_NORM=OFF skips the Sparrowhawk/OpenFST inverse-text-normalization
|
||||
# stack, which is the single most expensive leg of the build and needs a gcc-12
|
||||
# pin because OpenFST's templates ICE on gcc-13/14. It costs no coverage here:
|
||||
# nothing in include/nemo_speech/{asr,tts,diar,nmt}.h is conditional on it (the
|
||||
# only preprocessor conditionals in those headers are include guards,
|
||||
# __cplusplus and the _WIN32 export macros), so every struct layout this suite
|
||||
# checks is identical either way. The shipped images still build WITH_NORM=ON;
|
||||
# that path is covered by the backend image build in backend_pr.yml.
|
||||
tests-nemo-speech-cpp:
|
||||
needs: detect-changes
|
||||
if: needs.detect-changes.outputs.nemo-speech-cpp == 'true' || needs.detect-changes.outputs.run-all == 'true'
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 90
|
||||
steps:
|
||||
- name: Clone
|
||||
uses: actions/checkout@v7
|
||||
with:
|
||||
submodules: true
|
||||
- name: Dependencies
|
||||
run: |
|
||||
sudo apt-get update
|
||||
sudo apt-get install -y build-essential cmake ninja-build curl libopenblas-dev ffmpeg
|
||||
- name: Setup Go
|
||||
uses: actions/setup-go@v5
|
||||
- name: Display Go version
|
||||
run: go version
|
||||
- name: Proto Dependencies
|
||||
run: |
|
||||
curl -L -s https://github.com/protocolbuffers/protobuf/releases/download/v26.1/protoc-26.1-linux-x86_64.zip -o protoc.zip && \
|
||||
unzip -j -d /usr/local/bin protoc.zip bin/protoc && \
|
||||
rm protoc.zip
|
||||
go install google.golang.org/protobuf/cmd/protoc-gen-go@v1.34.2
|
||||
go install google.golang.org/grpc/cmd/protoc-gen-go-grpc@1958fcbe2ca8bd93af633f11e97d44e567e945af
|
||||
PATH="$PATH:$HOME/go/bin" make protogen-go
|
||||
- name: Test nemo-speech-cpp
|
||||
run: |
|
||||
make --jobs=5 --output-sync=target -C backend/go/nemo-speech-cpp WITH_NORM=OFF test
|
||||
# Per-backend smoke for rfdetr-cpp: builds the .so + Go binary and runs
|
||||
# `make -C backend/go/rfdetr-cpp test`. test.sh fetches the small (~20 MB)
|
||||
# rfdetr-nano-q8_0 GGUF from the published mudler/rfdetr-cpp-nano HF repo
|
||||
|
||||
@@ -65,6 +65,12 @@ jobs:
|
||||
- name: Test (with coverage gate)
|
||||
run: |
|
||||
PATH="$PATH:/root/go/bin" make --jobs 5 --output-sync=target test-coverage-check
|
||||
# tests/integration is outside the coverage roots because its store specs
|
||||
# need a live backend. test-stores builds and installs local-store before
|
||||
# running the complete suite, so new local-store specs are collected
|
||||
# automatically without adding another workflow entry.
|
||||
- name: Test local-store integration
|
||||
run: PATH="$PATH:$HOME/go/bin" make test-stores
|
||||
- name: Upload coverage report
|
||||
if: ${{ always() }}
|
||||
uses: actions/upload-artifact@v4
|
||||
|
||||
@@ -52,6 +52,8 @@ jobs:
|
||||
run: |
|
||||
sudo apt-get update
|
||||
sudo apt-get install -y build-essential libopus-dev
|
||||
- name: Run stale chunk recovery tests
|
||||
run: PATH="$PATH:$HOME/go/bin" make test-ui-stale-chunk
|
||||
# Builds an instrumented UI bundle, runs the Playwright specs, and fails
|
||||
# if line coverage regressed beyond the jitter tolerance (the gate is
|
||||
# in `make test-ui-coverage-check`). PLAYWRIGHT_CHROMIUM_PATH is unset
|
||||
|
||||
@@ -0,0 +1,21 @@
|
||||
## Design Context
|
||||
|
||||
### Users
|
||||
|
||||
LocalAI serves both single-host users who want to install and try models quickly and experienced developers, ML engineers, system administrators, and DevOps operators who manage production hosts or distributed clusters. The interface must support first-time discovery without hiding the runtime state, configuration, and control that returning operators need.
|
||||
|
||||
### Brand Personality
|
||||
|
||||
Capable, easy to use, and trustworthy. The interface should make sophisticated local-AI infrastructure feel understandable and under control. It should be direct and calm rather than playful, ornamental, or intimidating.
|
||||
|
||||
### Aesthetic Direction
|
||||
|
||||
Use LocalAI's established technical, editorial design language: Geist typography, compact information density, sharp geometry, deep blue-black surfaces, action blue, mint for healthy/local/live state, and amber only for decisions requiring attention. Support both dark and light themes. Avoid generic card dashboards, decorative gradients, glass effects, and visual noise.
|
||||
|
||||
### Design Principles
|
||||
|
||||
1. Use progressive disclosure to serve newcomers and operators in the same workflow: make the common path obvious, then reveal operational depth in context.
|
||||
2. Organize navigation around user intent and lifecycle state, not implementation concepts or nested containers.
|
||||
3. Give each resource one canonical home; expose discovery, installed state, and runtime state as clear views of that resource instead of duplicating management surfaces.
|
||||
4. Keep operational status visible and trustworthy through precise labels, explicit scope, and actionable state—not decoration.
|
||||
5. Preserve information density for expert use while flattening navigation and reducing repeated summaries, tabs, rails, and panels.
|
||||
@@ -27,6 +27,7 @@ To be removed, open a pull request deleting your row, or email
|
||||
|
||||
| Organisation | What they use it for | Status |
|
||||
|---|---|---|
|
||||
| [walcz.de](https://walcz.de) | Self-hosted appliance for a German B2B consultancy: local-only inference on AMD Strix Halo (gfx1151/ROCm), agents with MCP tools, RAG over an internal knowledge base, and a document/bookkeeping pipeline. | Production |
|
||||
| _Your organisation here_ | | |
|
||||
|
||||
## What this list is not
|
||||
|
||||
@@ -33,6 +33,7 @@ LocalAI follows the Linux kernel project's [guidelines for AI coding assistants]
|
||||
| [.agents/localai-assistant-mcp.md](.agents/localai-assistant-mcp.md) | LocalAI Assistant chat modality — adding admin tools to the in-process MCP server, editing skill prompts, keeping REST + MCP + skills in sync |
|
||||
| [.agents/backend-signing.md](.agents/backend-signing.md) | Backend OCI image signing (keyless cosign + sigstore-go) — producer-side CI setup, consumer-side gallery `verification:` block, strict mode (`LOCALAI_REQUIRE_BACKEND_INTEGRITY`), revocation via `not_before` |
|
||||
| [.agents/preparing-a-release.md](.agents/preparing-a-release.md) | Cutting a release: PR labels, `RELEASE_NOTES_vX.Y.Z.md`, the blog post under `website/content/blog/`, and the demo clips under `website/static/media/` |
|
||||
| [.impeccable.md](.impeccable.md) | Design context for UI/UX work — users, brand personality, aesthetic direction, and design principles |
|
||||
|
||||
## Quick Reference
|
||||
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
# Disable parallel execution for backend builds
|
||||
.NOTPARALLEL: backends/diffusers backends/llama-cpp backends/turboquant backends/bonsai backends/outetts backends/piper backends/stablediffusion-ggml backends/trellis2cpp backends/trellis2cpp-darwin backends/whisper backends/crispasr backends/parakeet-cpp backends/moss-transcribe-cpp backends/faster-whisper backends/silero-vad backends/local-store backends/valkey-store backends/cloud-proxy backends/huggingface backends/rfdetr backends/rfdetr-cpp backends/insightface backends/speaker-recognition backends/kitten-tts backends/kokoro backends/chatterbox backends/llama-cpp-darwin backends/neutts build-darwin-python-backend build-darwin-go-backend backends/mlx backends/diffuser-darwin backends/mlx-vlm backends/mlx-audio backends/mlx-distributed backends/stablediffusion-ggml-darwin backends/vllm backends/vllm-omni backends/longcat-video backends/sglang backends/moonshine backends/pocket-tts backends/qwen-tts backends/faster-qwen3-tts backends/qwen-asr backends/nemo backends/voxcpm backends/whisperx backends/ace-step backends/acestep-cpp backends/fish-speech backends/voxtral backends/opus backends/trl backends/llama-cpp-quantization backends/kokoros backends/sam3-cpp backends/qwen3-tts-cpp backends/moss-tts-cpp backends/magpie-tts-cpp backends/vllm-cpp backends/omnivoice-cpp backends/vibevoice-cpp backends/localvqe backends/tinygrad backends/sherpa-onnx backends/ds4 backends/ds4-darwin backends/liquid-audio backends/supertonic backends/depth-anything-cpp backends/privacy-filter backends/privacy-filter-darwin backends/audio-cpp backends/audio-cpp-darwin
|
||||
.NOTPARALLEL: backends/diffusers backends/llama-cpp backends/turboquant backends/bonsai backends/outetts backends/piper backends/stablediffusion-ggml backends/trellis2cpp backends/trellis2cpp-darwin backends/whisper backends/crispasr backends/parakeet-cpp backends/moss-transcribe-cpp backends/nemo-speech-cpp backends/faster-whisper backends/silero-vad backends/local-store backends/valkey-store backends/cloud-proxy backends/huggingface backends/rfdetr backends/rfdetr-cpp backends/insightface backends/speaker-recognition backends/kitten-tts backends/kokoro backends/chatterbox backends/llama-cpp-darwin backends/neutts build-darwin-python-backend build-darwin-go-backend backends/mlx backends/diffuser-darwin backends/mlx-vlm backends/mlx-audio backends/mlx-distributed backends/stablediffusion-ggml-darwin backends/vllm backends/vllm-omni backends/longcat-video backends/sglang backends/moonshine backends/pocket-tts backends/qwen-tts backends/faster-qwen3-tts backends/qwen-asr backends/nemo backends/voxcpm backends/whisperx backends/ace-step backends/acestep-cpp backends/fish-speech backends/voxtral backends/opus backends/trl backends/llama-cpp-quantization backends/kokoros backends/sam3-cpp backends/qwen3-tts-cpp backends/moss-tts-cpp backends/magpie-tts-cpp backends/vllm-cpp backends/omnivoice-cpp backends/vibevoice-cpp backends/localvqe backends/tinygrad backends/sherpa-onnx backends/ds4 backends/ds4-darwin backends/liquid-audio backends/supertonic backends/depth-anything-cpp backends/privacy-filter backends/privacy-filter-darwin backends/audio-cpp backends/audio-cpp-darwin
|
||||
|
||||
GOCMD=go
|
||||
GOTEST=$(GOCMD) test
|
||||
@@ -103,7 +103,7 @@ COVERAGE_E2E_LABELS?=!real-models
|
||||
COVERAGE_EXCLUDE_RE?=grpc/proto/.*[.]pb[.]go
|
||||
|
||||
|
||||
.PHONY: all test test-coverage test-coverage-baseline test-coverage-check test-backend-cpp test-build-scripts test-ui test-ui-coverage-baseline test-ui-coverage-check build vendor lint lint-all
|
||||
.PHONY: all test test-coverage test-coverage-baseline test-coverage-check test-backend-cpp test-build-scripts test-ui test-ui-stale-chunk test-ui-coverage-baseline test-ui-coverage-check build vendor lint lint-all
|
||||
|
||||
all: help
|
||||
|
||||
@@ -654,6 +654,7 @@ test-extra: prepare-test-extra
|
||||
$(MAKE) -C backend/go/depth-anything-cpp test
|
||||
$(MAKE) -C backend/go/supertonic test
|
||||
$(MAKE) -C backend/go/vllm-cpp test
|
||||
$(MAKE) -C backend/go/nemo-speech-cpp test
|
||||
$(MAKE) -C backend/go/trellis2cpp test
|
||||
$(MAKE) -C backend/go/valkey-store test
|
||||
|
||||
@@ -675,6 +676,7 @@ test-extra: prepare-test-extra
|
||||
## BACKEND_TEST_PROMPT Override the prompt used in predict/stream specs.
|
||||
## BACKEND_TEST_OPTIONS Comma-separated Options[] entries forwarded to LoadModel,
|
||||
## e.g. "tool_parser:hermes,reasoning_parser:qwen3".
|
||||
## BACKEND_TEST_EMBEDDING_LAYOUT Expected EmbeddingResult layout: "final" or "per_token".
|
||||
##
|
||||
## Direct usage (image already built, no docker-build-* dependency):
|
||||
##
|
||||
@@ -704,6 +706,7 @@ test-extra-backend: protogen-go
|
||||
BACKEND_TEST_CAPS="$$BACKEND_TEST_CAPS" \
|
||||
BACKEND_TEST_PROMPT="$$BACKEND_TEST_PROMPT" \
|
||||
BACKEND_TEST_OPTIONS="$$BACKEND_TEST_OPTIONS" \
|
||||
BACKEND_TEST_EMBEDDING_LAYOUT="$$BACKEND_TEST_EMBEDDING_LAYOUT" \
|
||||
BACKEND_TEST_TOOL_PROMPT="$$BACKEND_TEST_TOOL_PROMPT" \
|
||||
BACKEND_TEST_TOOL_NAME="$$BACKEND_TEST_TOOL_NAME" \
|
||||
BACKEND_TEST_CACHE_TYPE_K="$$BACKEND_TEST_CACHE_TYPE_K" \
|
||||
@@ -723,6 +726,15 @@ test-extra-backend-llama-cpp: docker-build-llama-cpp
|
||||
BACKEND_TEST_CAPS=health,load,predict,stream,logprobs,logit_bias \
|
||||
$(MAKE) test-extra-backend
|
||||
|
||||
## Raw llama.cpp embeddings are required by Go-side pooling. This exercises the
|
||||
## real C++ backend and verifies that it marks the flattened matrix per-token.
|
||||
test-extra-backend-llama-cpp-embeddings: docker-build-llama-cpp
|
||||
BACKEND_IMAGE=local-ai-backend:llama-cpp \
|
||||
BACKEND_TEST_CAPS=health,load,embeddings \
|
||||
BACKEND_TEST_OPTIONS=pooling:none \
|
||||
BACKEND_TEST_EMBEDDING_LAYOUT=per_token \
|
||||
$(MAKE) test-extra-backend
|
||||
|
||||
test-extra-backend-ik-llama-cpp: docker-build-ik-llama-cpp
|
||||
BACKEND_IMAGE=local-ai-backend:ik-llama-cpp $(MAKE) test-extra-backend
|
||||
|
||||
@@ -812,6 +824,7 @@ test-extra-backend-tinygrad-embeddings: docker-build-tinygrad
|
||||
BACKEND_IMAGE=local-ai-backend:tinygrad \
|
||||
BACKEND_TEST_MODEL_NAME=Qwen/Qwen3-0.6B \
|
||||
BACKEND_TEST_CAPS=health,load,embeddings \
|
||||
BACKEND_TEST_EMBEDDING_LAYOUT=final \
|
||||
$(MAKE) test-extra-backend
|
||||
|
||||
## tinygrad — Stable Diffusion 1.5. The original CompVis/runwayml repos have
|
||||
@@ -1298,6 +1311,7 @@ BACKEND_WHISPER = whisper|golang|.|false|true
|
||||
BACKEND_CRISPASR = crispasr|golang|.|false|true
|
||||
BACKEND_PARAKEET_CPP = parakeet-cpp|golang|.|false|true
|
||||
BACKEND_MOSS_TRANSCRIBE_CPP = moss-transcribe-cpp|golang|.|false|true
|
||||
BACKEND_NEMO_SPEECH_CPP = nemo-speech-cpp|golang|.|false|true
|
||||
BACKEND_DEPTH_ANYTHING_CPP = depth-anything-cpp|golang|.|false|true
|
||||
BACKEND_VOXTRAL = voxtral|golang|.|false|true
|
||||
BACKEND_ACESTEP_CPP = acestep-cpp|golang|.|false|true
|
||||
@@ -1400,6 +1414,7 @@ $(eval $(call generate-docker-build-target,$(BACKEND_WHISPER)))
|
||||
$(eval $(call generate-docker-build-target,$(BACKEND_CRISPASR)))
|
||||
$(eval $(call generate-docker-build-target,$(BACKEND_PARAKEET_CPP)))
|
||||
$(eval $(call generate-docker-build-target,$(BACKEND_MOSS_TRANSCRIBE_CPP)))
|
||||
$(eval $(call generate-docker-build-target,$(BACKEND_NEMO_SPEECH_CPP)))
|
||||
$(eval $(call generate-docker-build-target,$(BACKEND_DEPTH_ANYTHING_CPP)))
|
||||
$(eval $(call generate-docker-build-target,$(BACKEND_VOXTRAL)))
|
||||
$(eval $(call generate-docker-build-target,$(BACKEND_OPUS)))
|
||||
@@ -1456,7 +1471,7 @@ $(eval $(call generate-docker-build-target,$(BACKEND_SUPERTONIC)))
|
||||
docker-save-%: backend-images
|
||||
docker save local-ai-backend:$* -o backend-images/$*.tar
|
||||
|
||||
docker-build-backends: docker-build-llama-cpp docker-build-ik-llama-cpp docker-build-turboquant docker-build-bonsai docker-build-ds4 docker-build-rerankers docker-build-vllm docker-build-vllm-omni docker-build-longcat-video docker-build-sglang docker-build-transformers docker-build-outetts docker-build-diffusers docker-build-kokoro docker-build-faster-whisper docker-build-crispasr docker-build-coqui docker-build-chatterbox docker-build-vibevoice docker-build-liquid-audio docker-build-moonshine docker-build-pocket-tts docker-build-qwen-tts docker-build-fish-speech docker-build-faster-qwen3-tts docker-build-qwen-asr docker-build-nemo docker-build-voxcpm docker-build-whisperx docker-build-ace-step docker-build-acestep-cpp docker-build-voxtral docker-build-mlx-distributed docker-build-trl docker-build-llama-cpp-quantization docker-build-tinygrad docker-build-kokoros docker-build-sam3-cpp docker-build-rfdetr-cpp docker-build-qwen3-tts-cpp docker-build-moss-tts-cpp docker-build-magpie-tts-cpp docker-build-vllm-cpp docker-build-omnivoice-cpp docker-build-vibevoice-cpp docker-build-localvqe docker-build-insightface docker-build-speaker-recognition docker-build-sherpa-onnx docker-build-cloud-proxy docker-build-supertonic docker-build-depth-anything-cpp docker-build-moss-transcribe-cpp docker-build-privacy-filter docker-build-trellis2cpp docker-build-valkey-store docker-build-audio-cpp
|
||||
docker-build-backends: docker-build-llama-cpp docker-build-ik-llama-cpp docker-build-turboquant docker-build-bonsai docker-build-ds4 docker-build-rerankers docker-build-vllm docker-build-vllm-omni docker-build-longcat-video docker-build-sglang docker-build-transformers docker-build-outetts docker-build-diffusers docker-build-kokoro docker-build-faster-whisper docker-build-crispasr docker-build-coqui docker-build-chatterbox docker-build-vibevoice docker-build-liquid-audio docker-build-moonshine docker-build-pocket-tts docker-build-qwen-tts docker-build-fish-speech docker-build-faster-qwen3-tts docker-build-qwen-asr docker-build-nemo docker-build-voxcpm docker-build-whisperx docker-build-ace-step docker-build-acestep-cpp docker-build-voxtral docker-build-mlx-distributed docker-build-trl docker-build-llama-cpp-quantization docker-build-tinygrad docker-build-kokoros docker-build-sam3-cpp docker-build-rfdetr-cpp docker-build-qwen3-tts-cpp docker-build-moss-tts-cpp docker-build-magpie-tts-cpp docker-build-vllm-cpp docker-build-omnivoice-cpp docker-build-vibevoice-cpp docker-build-localvqe docker-build-insightface docker-build-speaker-recognition docker-build-sherpa-onnx docker-build-cloud-proxy docker-build-supertonic docker-build-depth-anything-cpp docker-build-moss-transcribe-cpp docker-build-nemo-speech-cpp docker-build-privacy-filter docker-build-trellis2cpp docker-build-valkey-store docker-build-audio-cpp
|
||||
|
||||
########################################################
|
||||
### Mock Backend for E2E Tests
|
||||
@@ -1502,6 +1517,13 @@ test-ui: build-mock-backend protogen-go
|
||||
$(GOCMD) build -o tests/e2e-ui/ui-test-server ./tests/e2e-ui
|
||||
cd core/http/react-ui && sh $(CURDIR)/scripts/ensure-playwright-browser.sh && bunx playwright test $(PLAYWRIGHT_WORKERS_FLAG)
|
||||
|
||||
## The stale-chunk specs need the production code-split bundle. The V8 coverage
|
||||
## bundle below inlines dynamic imports to keep every page in its denominator.
|
||||
test-ui-stale-chunk: build-mock-backend protogen-go
|
||||
cd core/http/react-ui && bun install && bun run build
|
||||
$(GOCMD) build -o tests/e2e-ui/ui-test-server ./tests/e2e-ui
|
||||
cd core/http/react-ui && sh $(CURDIR)/scripts/ensure-playwright-browser.sh && bunx playwright test --grep @production-chunks --workers=1
|
||||
|
||||
## React UI code coverage from the Playwright e2e suite. Builds a
|
||||
## NON-instrumented bundle with source maps (COVERAGE_V8=true), re-embeds it
|
||||
## into the ui-test-server (the dist is //go:embed'ed at compile time), runs the
|
||||
@@ -1517,7 +1539,7 @@ test-ui-coverage: build-mock-backend protogen-go
|
||||
$(GOCMD) build -o tests/e2e-ui/ui-test-server ./tests/e2e-ui && \
|
||||
( cd core/http/react-ui && rm -rf .nyc_output coverage && \
|
||||
sh $(CURDIR)/scripts/ensure-playwright-browser.sh && \
|
||||
PW_V8_COVERAGE=1 bunx playwright test $(PLAYWRIGHT_WORKERS_FLAG) && bun run coverage:report )
|
||||
PW_V8_COVERAGE=1 bunx playwright test --grep-invert @production-chunks $(PLAYWRIGHT_WORKERS_FLAG) && bun run coverage:report )
|
||||
|
||||
## UI coverage baseline (committed) and the strict gate that compares against
|
||||
## it — the React mirror of test-coverage-baseline / test-coverage-check.
|
||||
|
||||
@@ -5,9 +5,6 @@
|
||||
</h1>
|
||||
|
||||
<p align="center">
|
||||
<a href="https://github.com/go-skynet/LocalAI/stargazers" target="blank">
|
||||
<img src="https://img.shields.io/github/stars/go-skynet/LocalAI?style=for-the-badge" alt="LocalAI stars"/>
|
||||
</a>
|
||||
<a href='https://github.com/go-skynet/LocalAI/releases'>
|
||||
<img src='https://img.shields.io/github/release/go-skynet/LocalAI?&label=Latest&style=for-the-badge'>
|
||||
</a>
|
||||
@@ -231,7 +228,7 @@ Most backends wrap a best-in-class upstream engine. A handful of them are native
|
||||
|
||||
| Backend | What it does |
|
||||
|---------|-------------|
|
||||
| [vllm.cpp](https://github.com/mudler/vllm.cpp) | From-scratch C++20 port of vLLM for text generation: paged KV cache, continuous batching, prefix caching, safetensors + GGUF loading, engine-enforced structured output, on CPU, CUDA, Metal and Vulkan |
|
||||
| [vllm.cpp](https://github.com/mudler/vllm.cpp) | From-scratch C++20 port of vLLM for text generation: paged KV cache, continuous batching, prefix caching, safetensors + GGUF loading, engine-enforced structured output, on CPU, CUDA, Metal and Vulkan. Also serves MiniMax-H3 joint video+audio generation |
|
||||
| [parakeet.cpp](https://github.com/mudler/parakeet.cpp) | C++/GGML port of NVIDIA NeMo Parakeet ASR (tdt/ctc/rnnt/hybrid), with cache-aware streaming transcription |
|
||||
| [moss-transcribe.cpp](https://github.com/localai-org/moss-transcribe.cpp) | C++/GGML port of OpenMOSS MOSS-Transcribe-Diarize: joint long-form transcription, speaker diarization and timestamping in a single pass |
|
||||
| [moss-tts.cpp](https://github.com/mudler/moss-tts.cpp) | C++/GGML port of the OpenMOSS MOSS-TTS family: text-to-speech (MOSS-TTS-Local v1.5, 48 kHz stereo) with reference-audio voice cloning, through the MOSS-Audio-Tokenizer neural codec |
|
||||
@@ -318,10 +315,6 @@ Past sponsors
|
||||
|
||||
A special thanks to individual sponsors, a full list is on [GitHub](https://github.com/sponsors/mudler) and [buymeacoffee](https://buymeacoffee.com/mudler). Special shout out to [drikster80](https://github.com/drikster80) for being generous. Thank you everyone!
|
||||
|
||||
## Star history
|
||||
|
||||
[](https://star-history.com/#go-skynet/LocalAI&Date)
|
||||
|
||||
## License
|
||||
|
||||
LocalAI is a community-driven project created by [Ettore Di Giacinto](https://github.com/mudler/) and maintained by the [LocalAI team](#team).
|
||||
|
||||
@@ -6,13 +6,14 @@ ARG APT_PORTS_MIRROR=""
|
||||
# ASR, VAD, diarization, source separation and music generation, wrapped as a
|
||||
# LocalAI gRPC backend.
|
||||
#
|
||||
# BASE_IMAGE is ubuntu:24.04 for cpu and vulkan builds, or
|
||||
# nvidia/cuda:<ver>-devel-ubuntu24.04 for cublas builds; both ship apt and
|
||||
# Ubuntu Noble packages, and the CUDA base additionally provides
|
||||
# /usr/local/cuda. BUILD_TYPE selects the engine backend in the Makefile:
|
||||
# "" = portable CPU with all ggml CPU variants, "cublas" ->
|
||||
# -DENGINE_ENABLE_CUDA=ON, "vulkan" -> -DENGINE_ENABLE_VULKAN=ON. Darwin
|
||||
# (Metal) builds bypass this Dockerfile entirely.
|
||||
# BASE_IMAGE is ubuntu:24.04 for cpu and vulkan builds,
|
||||
# nvidia/cuda:<ver>-devel-ubuntu24.04 for cublas builds, or
|
||||
# rocm/dev-ubuntu-24.04:<ver> for hipblas builds. All ship apt and Ubuntu Noble
|
||||
# packages; the GPU bases also provide their toolkits. BUILD_TYPE selects the
|
||||
# engine backend in the Makefile: "" = portable CPU with all ggml CPU variants,
|
||||
# "cublas" -> -DENGINE_ENABLE_CUDA=ON, "hipblas" -> -DENGINE_ENABLE_HIP=ON,
|
||||
# and "vulkan" -> -DENGINE_ENABLE_VULKAN=ON. Darwin (Metal) builds bypass this
|
||||
# Dockerfile entirely.
|
||||
#
|
||||
# Upstream needs GCC 13 or newer, which ubuntu:24.04 and the CUDA 12/13
|
||||
# devel-ubuntu24.04 images all provide.
|
||||
@@ -62,7 +63,7 @@ ENV BUILD_TYPE=${BUILD_TYPE} \
|
||||
APT_MIRROR=${APT_MIRROR} \
|
||||
APT_PORTS_MIRROR=${APT_PORTS_MIRROR} \
|
||||
DEBIAN_FRONTEND=noninteractive \
|
||||
PATH=/usr/local/cuda/bin:${PATH}
|
||||
PATH=/opt/rocm/bin:/usr/local/cuda/bin:${PATH}
|
||||
|
||||
WORKDIR /build
|
||||
|
||||
@@ -73,7 +74,7 @@ WORKDIR /build
|
||||
# fallback of its own.
|
||||
#
|
||||
# BUILD_TYPE=vulkan additionally needs the loader headers and glslc; both are in
|
||||
# Noble. The CUDA toolkit for BUILD_TYPE=cublas comes from BASE_IMAGE.
|
||||
# Noble. The CUDA and ROCm toolkits come from their matching BASE_IMAGE.
|
||||
RUN --mount=type=bind,source=.docker/apt-mirror.sh,target=/usr/local/sbin/apt-mirror \
|
||||
sh /usr/local/sbin/apt-mirror && \
|
||||
apt-get update && \
|
||||
@@ -83,6 +84,9 @@ RUN --mount=type=bind,source=.docker/apt-mirror.sh,target=/usr/local/sbin/apt-mi
|
||||
if [ "${BUILD_TYPE}" = "vulkan" ]; then \
|
||||
apt-get install -y --no-install-recommends libvulkan-dev glslc; \
|
||||
fi && \
|
||||
if [ "${BUILD_TYPE}" = "hipblas" ]; then \
|
||||
apt-get install -y --no-install-recommends hipblas-dev rocblas-dev; \
|
||||
fi && \
|
||||
if [ "${TARGETARCH}" = "arm64" ]; then \
|
||||
apt-get install -y --no-install-recommends gcc-14 g++-14; \
|
||||
fi && \
|
||||
|
||||
@@ -248,6 +248,134 @@ RUN <<EOT bash
|
||||
fi
|
||||
EOT
|
||||
|
||||
# nemo-speech-cpp builds NVIDIA NeMo-Speech.cpp with text normalization enabled,
|
||||
# which compiles the Sparrowhawk/OpenFST WFST stack from source via
|
||||
# scripts/build_itn_deps.sh. That step needs gcc-12 specifically: OpenFST's
|
||||
# template-heavy translation units ICE on gcc-13 and gcc-14 at -O2, so upstream
|
||||
# pins gcc-12 for it while the runtime itself builds with the image default.
|
||||
# No update-alternatives here, so the default compiler is untouched; the backend
|
||||
# Makefile reaches gcc-12 by name for that one step.
|
||||
#
|
||||
# The rest is what build_itn_deps.sh and the WITH_NORM cmake block expect:
|
||||
# protobuf (headers plus protoc, which must come from the same apt set so the
|
||||
# generated stubs match the headers they compile against) and re2 for
|
||||
# Sparrowhawk, and autotools because OpenFST and Sparrowhawk ship autoconf
|
||||
# builds. ninja is not in the common apt list because this is the only Go
|
||||
# backend that configures with -G Ninja, and that list is a layer shared by
|
||||
# every backend image in the matrix.
|
||||
#
|
||||
# No libabsl-dev, despite upstream's Dockerfile installing it: upstream builds
|
||||
# against protobuf 25, which splits its runtime across libabsl_*, whereas every
|
||||
# base image in this matrix carries protobuf 3.21 (noble) or 3.12 (jammy), which
|
||||
# has no absl dependency. The cmake block's file(GLOB ... /usr/lib/libabsl_*.so)
|
||||
# would not match on Ubuntu anyway, since multiarch puts those under
|
||||
# /usr/lib/<triplet>/.
|
||||
#
|
||||
# Placed down here with the other per-backend gates rather than next to the
|
||||
# shared apt layer: Docker re-keys every layer below an inserted one, so adding
|
||||
# a step above the Vulkan SDK, CUDA, Go and protoc layers would force all of
|
||||
# them to re-execute once for every Go backend image, not just this one.
|
||||
# Nothing between there and here needs any of these packages (the Vulkan and
|
||||
# opus blocks install their own ninja and pkg-config, and the protoc download is
|
||||
# a release binary that needs neither libprotobuf-dev nor protoc from apt), and
|
||||
# nothing here needs anything those layers provide.
|
||||
#
|
||||
# The second half of this block backfills cmake. NeMo-Speech.cpp opens with
|
||||
# cmake_minimum_required(VERSION 3.26), which every noble base in the matrix
|
||||
# satisfies (24.04 ships 3.28) but the JetPack r36.4.0 row does not: that image
|
||||
# is jammy, whose apt cmake is 3.22, so configure aborts before it reads a
|
||||
# single one of our -D flags. This is the only Go backend that needs more than
|
||||
# jammy's cmake; parakeet-cpp and moss-transcribe-cpp share the same JetPack
|
||||
# base and both declare cmake_minimum_required(VERSION 3.18).
|
||||
#
|
||||
# Taken from Kitware's own release tarball rather than from their APT repo or
|
||||
# from pip. The tarball is a pinned URL with a published checksum, so the build
|
||||
# is reproducible and an upstream release cannot change what lands here; the
|
||||
# APT repo serves a moving 'latest', which today would be CMake 4.x, and 4.x
|
||||
# drops compatibility with cmake_minimum_required below 3.5 and so would break
|
||||
# vendored third_party subprojects that still declare one. pip would drag a
|
||||
# Python toolchain into a backend that otherwise has none. The binaries need
|
||||
# only glibc 2.17 and carry no libstdc++ DT_NEEDED, so jammy's 2.35 is far
|
||||
# above the floor. doc/, man/, ccmake and cmake-gui are left in the tarball;
|
||||
# this is a builder stage and the final image is FROM scratch, but there is no
|
||||
# reason to page 50 MB of Qt GUI and docs through the CI cache.
|
||||
#
|
||||
# Conditional on the installed cmake being too old rather than unconditional,
|
||||
# so the rows that already build green (noble cpu, vulkan, cublas and hipblas)
|
||||
# keep configuring with exactly the cmake they configure with today.
|
||||
#
|
||||
# The version test compares through two temp files and a grep on the exit
|
||||
# status rather than the obvious "$(sort -V ... | head -n1)". BuildKit delivers
|
||||
# a RUN heredoc through an outer shell with an unquoted delimiter, so the outer
|
||||
# shell expands the body before bash ever sees it: a $(...) here runs once, too
|
||||
# early, in a container where the files it reads do not exist yet, and its empty
|
||||
# output is then pasted into the script. Same reason there are no shell
|
||||
# variables below. ${BACKEND} and ${TARGETARCH} are fine because they are build
|
||||
# args, which BuildKit exports into that outer shell's environment.
|
||||
#
|
||||
# The symlink goes in /usr/local/bin, which precedes /usr/bin on PATH, so it
|
||||
# shadows apt's cmake. That is deliberate and, unlike the protoc shadowing that
|
||||
# broke Sparrowhawk earlier in this PR, it is inert: protoc has to agree with
|
||||
# the libprotobuf headers it generates against, whereas cmake is a standalone
|
||||
# build driver with no ABI relationship to anything in the image, and it locates
|
||||
# its own Modules/ tree by resolving the symlink back to /opt, so a 3.31 binary
|
||||
# can never read 3.22's modules. Scope is the ${BACKEND} gate: no other Go
|
||||
# backend image gets /opt/cmake or the symlink. Inside this image the only
|
||||
# other cmake consumers, the base apt layer and the Vulkan SDK build, both run
|
||||
# in layers above this one and have already finished.
|
||||
RUN <<EOT bash
|
||||
if [ "${BACKEND}" = "nemo-speech-cpp" ]; then
|
||||
set -e
|
||||
apt-get update
|
||||
apt-get install -y --no-install-recommends \
|
||||
gcc-12 g++-12 \
|
||||
ninja-build \
|
||||
libprotobuf-dev protobuf-compiler \
|
||||
libre2-dev \
|
||||
autoconf automake libtool pkg-config
|
||||
apt-get clean
|
||||
rm -rf /var/lib/apt/lists/*
|
||||
|
||||
echo 3.26.0 > /tmp/cmake-required
|
||||
cmake --version 2>/dev/null | head -n1 | cut -d' ' -f3 > /tmp/cmake-present
|
||||
if [ ! -s /tmp/cmake-present ]; then
|
||||
echo 0.0.0 > /tmp/cmake-present
|
||||
fi
|
||||
if sort -V /tmp/cmake-required /tmp/cmake-present | head -n1 | grep -qxF 3.26.0; then
|
||||
echo "==> cmake is new enough for NeMo-Speech.cpp:"
|
||||
cmake --version | head -n1
|
||||
else
|
||||
echo "==> cmake is below the 3.26 NeMo-Speech.cpp requires; installing 3.31.12. Found:"
|
||||
cat /tmp/cmake-present
|
||||
mkdir -p /opt/cmake
|
||||
if [ "${TARGETARCH}" = "arm64" ]; then
|
||||
curl -fsSL -o /tmp/cmake.tar.gz https://github.com/Kitware/CMake/releases/download/v3.31.12/cmake-3.31.12-linux-aarch64.tar.gz
|
||||
echo "83f8fd91d2038a56556e1400390fcfe42f79602940c494f6c6f1cdae7f9e7f40 /tmp/cmake.tar.gz" | sha256sum -c -
|
||||
tar -xzf /tmp/cmake.tar.gz -C /opt/cmake --strip-components=1 \
|
||||
cmake-3.31.12-linux-aarch64/bin/cmake \
|
||||
cmake-3.31.12-linux-aarch64/bin/cpack \
|
||||
cmake-3.31.12-linux-aarch64/bin/ctest \
|
||||
cmake-3.31.12-linux-aarch64/share
|
||||
else
|
||||
curl -fsSL -o /tmp/cmake.tar.gz https://github.com/Kitware/CMake/releases/download/v3.31.12/cmake-3.31.12-linux-x86_64.tar.gz
|
||||
echo "0dc2e9a6860f06bf10bd8fadc03e35d9eeb4df46e33763a7e480e987758f385c /tmp/cmake.tar.gz" | sha256sum -c -
|
||||
tar -xzf /tmp/cmake.tar.gz -C /opt/cmake --strip-components=1 \
|
||||
cmake-3.31.12-linux-x86_64/bin/cmake \
|
||||
cmake-3.31.12-linux-x86_64/bin/cpack \
|
||||
cmake-3.31.12-linux-x86_64/bin/ctest \
|
||||
cmake-3.31.12-linux-x86_64/share
|
||||
fi
|
||||
rm -f /tmp/cmake.tar.gz
|
||||
ln -sf /opt/cmake/bin/cmake /usr/local/bin/cmake
|
||||
ln -sf /opt/cmake/bin/cpack /usr/local/bin/cpack
|
||||
ln -sf /opt/cmake/bin/ctest /usr/local/bin/ctest
|
||||
hash -r
|
||||
cmake --version
|
||||
fi
|
||||
rm -f /tmp/cmake-required /tmp/cmake-present
|
||||
fi
|
||||
EOT
|
||||
|
||||
RUN git config --global --add safe.directory /LocalAI
|
||||
|
||||
# Prebuild the native engine from a layer that depends on this backend's own
|
||||
|
||||
@@ -536,8 +536,28 @@ message Result {
|
||||
bool success = 2;
|
||||
}
|
||||
|
||||
// EmbeddingLayout describes whether embeddings contains one final vector or
|
||||
// a matrix of per-token vectors. Go-side pooling must never infer this from
|
||||
// tokens/dim alone: a one-token raw matrix and a final vector have the same
|
||||
// shape.
|
||||
enum EmbeddingLayout {
|
||||
EMBEDDING_LAYOUT_UNSPECIFIED = 0;
|
||||
EMBEDDING_LAYOUT_FINAL = 1;
|
||||
EMBEDDING_LAYOUT_PER_TOKEN = 2;
|
||||
}
|
||||
|
||||
message EmbeddingResult {
|
||||
repeated float embeddings = 1;
|
||||
// Shape of the payload above: dim is the embedding width, tokens is the
|
||||
// number of vectors packed into `embeddings` (1 when the backend pooled
|
||||
// server-side, N with pooling:none; total across prompts if a request
|
||||
// carried several). tokens=0/dim=0 means the backend predates shape
|
||||
// reporting. prompt_tokens is the number of prompt tokens evaluated, for
|
||||
// usage accounting.
|
||||
int32 tokens = 2;
|
||||
int32 dim = 3;
|
||||
int32 prompt_tokens = 4;
|
||||
EmbeddingLayout layout = 5;
|
||||
}
|
||||
|
||||
message TranscriptRequest {
|
||||
|
||||
@@ -9,7 +9,7 @@
|
||||
# recipe is a make target (not a prepare.sh) so 'make purge && make' is a clean
|
||||
# rebuild and so the bump bot can see the pin.
|
||||
|
||||
AUDIO_CPP_VERSION?=5a8312ef7b8aa7cf14e9a24ac568cabd8725d68a
|
||||
AUDIO_CPP_VERSION?=43001a7e0f452d80f4588e613f13332940dd4d3a
|
||||
AUDIO_CPP_REPO?=https://github.com/0xShug0/audio.cpp
|
||||
|
||||
CURRENT_MAKEFILE_DIR := $(dir $(abspath $(lastword $(MAKEFILE_LIST))))
|
||||
@@ -77,6 +77,16 @@ endif
|
||||
|
||||
ifeq ($(BUILD_TYPE),cublas)
|
||||
CMAKE_ARGS += -DENGINE_ENABLE_CUDA=ON "-DCMAKE_CUDA_ARCHITECTURES=$(CUDA_ARCHITECTURES)"
|
||||
else ifeq ($(BUILD_TYPE),hipblas)
|
||||
ROCM_HOME ?= /opt/rocm
|
||||
ROCM_PATH ?= /opt/rocm
|
||||
export CXX=$(ROCM_HOME)/llvm/bin/clang++
|
||||
export CC=$(ROCM_HOME)/llvm/bin/clang
|
||||
AMDGPU_TARGETS ?= gfx908,gfx90a,gfx942,gfx950,gfx1030,gfx1100,gfx1101,gfx1102,gfx1151,gfx1200,gfx1201
|
||||
# audio.cpp forwards GPU_TARGETS to CMake's semicolon-delimited HIP list.
|
||||
comma := ,
|
||||
HIP_TARGETS := $(subst $(comma),;,$(AMDGPU_TARGETS))
|
||||
CMAKE_ARGS += -DENGINE_ENABLE_HIP=ON "-DGPU_TARGETS=$(HIP_TARGETS)"
|
||||
else ifeq ($(BUILD_TYPE),vulkan)
|
||||
CMAKE_ARGS += -DENGINE_ENABLE_VULKAN=ON
|
||||
else ifeq ($(UNAME_S),Darwin)
|
||||
|
||||
@@ -29,6 +29,7 @@ const NamedTask kTaskNames[] = {
|
||||
{Task::VoiceDesign, "vdes"},
|
||||
{Task::SpeakerRecognition, "spk"},
|
||||
{Task::Svc, "svc"},
|
||||
{Task::Midi, "midi"},
|
||||
};
|
||||
|
||||
// Accepted on input but never emitted. "spkrec" was this backend's own earlier
|
||||
|
||||
@@ -25,6 +25,7 @@ enum class Task {
|
||||
VoiceDesign,
|
||||
SpeakerRecognition,
|
||||
Svc,
|
||||
Midi,
|
||||
};
|
||||
|
||||
// Mirrors engine::runtime::RunMode.
|
||||
|
||||
@@ -361,7 +361,7 @@ static void test_names_round_trip() {
|
||||
Task::SourceSeparation, Task::AudioGeneration, Task::Tts,
|
||||
Task::VoiceCloning, Task::VoiceConversion,
|
||||
Task::SpeechToSpeech, Task::Alignment, Task::VoiceDesign,
|
||||
Task::SpeakerRecognition, Task::Svc};
|
||||
Task::SpeakerRecognition, Task::Svc, Task::Midi};
|
||||
for (const Task t : all) {
|
||||
Task parsed = Task::Vad;
|
||||
const bool ok = parse_task_name(task_name(t), parsed);
|
||||
|
||||
@@ -69,7 +69,8 @@ static_assert(kEngine(engine::runtime::VoiceTaskKind::VoiceDesign) == 10, "Voice
|
||||
static_assert(kEngine(engine::runtime::VoiceTaskKind::SpeakerRecognition) == 11, "VoiceTaskKind drifted");
|
||||
// The last member. Pinning it pins the member count too, as long as the
|
||||
// enumerators stay contiguous and unassigned, which upstream's declaration is.
|
||||
static_assert(kEngine(engine::runtime::VoiceTaskKind::Svc) == 12,
|
||||
static_assert(kEngine(engine::runtime::VoiceTaskKind::Svc) == 12, "VoiceTaskKind drifted");
|
||||
static_assert(kEngine(engine::runtime::VoiceTaskKind::Midi) == 13,
|
||||
"engine::runtime::VoiceTaskKind gained, lost or reordered a member. "
|
||||
"audiocpp_backend::Task mirrors it positionally: update capability_routing.h, "
|
||||
"to_engine_task and from_engine_task together, then move this pin.");
|
||||
@@ -87,6 +88,7 @@ static_assert(kMirror(Task::Alignment) == 9, "Task drifted from VoiceTaskKind");
|
||||
static_assert(kMirror(Task::VoiceDesign) == 10, "Task drifted from VoiceTaskKind");
|
||||
static_assert(kMirror(Task::SpeakerRecognition) == 11, "Task drifted from VoiceTaskKind");
|
||||
static_assert(kMirror(Task::Svc) == 12, "Task drifted from VoiceTaskKind");
|
||||
static_assert(kMirror(Task::Midi) == 13, "Task drifted from VoiceTaskKind");
|
||||
|
||||
static_assert(static_cast<int>(engine::runtime::RunMode::Offline) == 0, "RunMode drifted");
|
||||
static_assert(static_cast<int>(engine::runtime::RunMode::Streaming) == 1,
|
||||
@@ -101,6 +103,9 @@ engine::core::BackendType parse_backend_type(const std::string &value) {
|
||||
if (value == "cuda") {
|
||||
return engine::core::BackendType::Cuda;
|
||||
}
|
||||
if (value == "hip" || value == "rocm") {
|
||||
return engine::core::BackendType::Hip;
|
||||
}
|
||||
if (value == "vulkan") {
|
||||
return engine::core::BackendType::Vulkan;
|
||||
}
|
||||
@@ -114,7 +119,7 @@ engine::core::BackendType parse_backend_type(const std::string &value) {
|
||||
return engine::core::BackendType::Cpu;
|
||||
}
|
||||
throw ConfigError("audio-cpp: unknown backend option '" + value +
|
||||
"'. Known backends: cpu, cuda, vulkan, metal, best");
|
||||
"'. Known backends: cpu, cuda, hip, rocm, vulkan, metal, best");
|
||||
}
|
||||
|
||||
std::filesystem::path executable_directory() {
|
||||
@@ -241,6 +246,7 @@ engine::runtime::VoiceTaskKind to_engine_task(Task task) {
|
||||
case Task::VoiceDesign: return K::VoiceDesign;
|
||||
case Task::SpeakerRecognition: return K::SpeakerRecognition;
|
||||
case Task::Svc: return K::Svc;
|
||||
case Task::Midi: return K::Midi;
|
||||
}
|
||||
// Unreachable for any valid enumerator. No `default:` label, so -Wswitch
|
||||
// still reports a member this switch stops covering.
|
||||
@@ -263,6 +269,7 @@ Task from_engine_task(engine::runtime::VoiceTaskKind kind) {
|
||||
case K::VoiceDesign: return Task::VoiceDesign;
|
||||
case K::SpeakerRecognition: return Task::SpeakerRecognition;
|
||||
case K::Svc: return Task::Svc;
|
||||
case K::Midi: return Task::Midi;
|
||||
}
|
||||
return Task::Vad;
|
||||
}
|
||||
|
||||
@@ -59,6 +59,12 @@ bool starts_with(const std::string &value, const std::string &prefix) {
|
||||
value.compare(0, prefix.size(), prefix) == 0;
|
||||
}
|
||||
|
||||
bool is_known_backend(const std::string &value) {
|
||||
return value == "cpu" || value == "cuda" || value == "hip" ||
|
||||
value == "rocm" || value == "vulkan" || value == "metal" ||
|
||||
value == "best";
|
||||
}
|
||||
|
||||
} // namespace
|
||||
|
||||
ParsedOptions parse_model_options(const std::vector<std::string> &entries) {
|
||||
@@ -107,6 +113,12 @@ ParsedOptions parse_model_options(const std::vector<std::string> &entries) {
|
||||
} else if (key == "task") {
|
||||
parsed.options.task = value;
|
||||
} else if (key == "backend") {
|
||||
if (!is_known_backend(value)) {
|
||||
parsed.error = "audio-cpp: unknown backend option '" + value +
|
||||
"'. Known backends: cpu, cuda, hip, rocm, "
|
||||
"vulkan, metal, best";
|
||||
return parsed;
|
||||
}
|
||||
parsed.options.backend = value;
|
||||
} else if (key == "model_spec_override") {
|
||||
parsed.options.model_spec_override = value;
|
||||
|
||||
@@ -17,7 +17,7 @@ struct ModelOptions {
|
||||
std::string family;
|
||||
// Pins the audio.cpp task, overriding RPC-based routing. Empty means route.
|
||||
std::string task;
|
||||
// ggml backend: cpu, cuda, vulkan, metal, best.
|
||||
// ggml backend: cpu, cuda, hip (or rocm), vulkan, metal, best.
|
||||
std::string backend = "cpu";
|
||||
int device = 0;
|
||||
// True once a `device:` entry has been seen. 0 is both the default and a
|
||||
|
||||
@@ -75,6 +75,11 @@ static void test_scalar_options() {
|
||||
check(parse_model_options({"live_idle_timeout_ms:0"}).options.live_idle_timeout_ms == 0,
|
||||
"an explicit 0 turns the live idle limit off rather than reverting to "
|
||||
"the default");
|
||||
|
||||
check(parse_model_options({"backend:hip"}).error.empty(),
|
||||
"HIP backend option is accepted");
|
||||
check(parse_model_options({"backend:rocm"}).error.empty(),
|
||||
"ROCm backend alias is accepted");
|
||||
}
|
||||
|
||||
// Values containing colons must survive: split on the FIRST colon only.
|
||||
@@ -124,6 +129,8 @@ static void test_errors() {
|
||||
"negative device is rejected");
|
||||
check(!parse_model_options({"threads:x"}).error.empty(),
|
||||
"non-numeric threads is rejected");
|
||||
check(!parse_model_options({"backend:unknown"}).error.empty(),
|
||||
"unknown compute backend is rejected before model loading");
|
||||
|
||||
// Values too large for int must be rejected, not silently wrapped into a
|
||||
// negative device index that then reaches the ggml backend selector.
|
||||
|
||||
@@ -42,6 +42,7 @@ define bonsai-build
|
||||
rm -rf $(CURRENT_MAKEFILE_DIR)/../bonsai-$(1)-build/patches
|
||||
$(MAKE) -C $(CURRENT_MAKEFILE_DIR)/../bonsai-$(1)-build purge
|
||||
bash $(LLAMA_CPP_DIR)/disable-score-task.sh $(CURRENT_MAKEFILE_DIR)/../bonsai-$(1)-build/grpc-server.cpp
|
||||
bash $(LLAMA_CPP_DIR)/disable-tts-task.sh $(CURRENT_MAKEFILE_DIR)/../bonsai-$(1)-build/grpc-server.cpp
|
||||
$(info $(GREEN)I bonsai build info:$(1)$(RESET))
|
||||
LLAMA_REPO=$(LLAMA_REPO) LLAMA_VERSION=$(BONSAI_VERSION) \
|
||||
$(MAKE) -C $(CURRENT_MAKEFILE_DIR)/../bonsai-$(1)-build llama.cpp
|
||||
@@ -79,6 +80,7 @@ bonsai-cpu-all:
|
||||
rm -rf $(CURRENT_MAKEFILE_DIR)/../bonsai-cpu-all-build/patches
|
||||
$(MAKE) -C $(CURRENT_MAKEFILE_DIR)/../bonsai-cpu-all-build purge
|
||||
bash $(LLAMA_CPP_DIR)/disable-score-task.sh $(CURRENT_MAKEFILE_DIR)/../bonsai-cpu-all-build/grpc-server.cpp
|
||||
bash $(LLAMA_CPP_DIR)/disable-tts-task.sh $(CURRENT_MAKEFILE_DIR)/../bonsai-cpu-all-build/grpc-server.cpp
|
||||
$(info $(GREEN)I bonsai build info:cpu-all-variants$(RESET))
|
||||
LLAMA_REPO=$(LLAMA_REPO) LLAMA_VERSION=$(BONSAI_VERSION) \
|
||||
$(MAKE) -C $(CURRENT_MAKEFILE_DIR)/../bonsai-cpu-all-build llama.cpp
|
||||
|
||||
@@ -69,7 +69,15 @@ target_include_directories(hw_grpc_proto PUBLIC ${CMAKE_CURRENT_BINARY_DIR})
|
||||
|
||||
set(DS4_OBJS "${DS4_DIR}/ds4.o")
|
||||
if(DS4_GPU STREQUAL "cuda")
|
||||
list(APPEND DS4_OBJS "${DS4_DIR}/ds4_cuda.o")
|
||||
list(APPEND DS4_OBJS
|
||||
"${DS4_DIR}/ds4_cuda.o"
|
||||
"${DS4_DIR}/cuda/mmq/ds4_ggml_stubs.o"
|
||||
"${DS4_DIR}/cuda/mmq/ds4_mmq.o"
|
||||
"${DS4_DIR}/cuda/mmq/ds4_mmq_d2r.o"
|
||||
"${DS4_DIR}/cuda/mmq/quantize.o"
|
||||
"${DS4_DIR}/cuda/mmq/mmid.o"
|
||||
"${DS4_DIR}/cuda/mmq/mmvq.o"
|
||||
"${DS4_DIR}/cuda/mmq/ds4_repack.o")
|
||||
elseif(DS4_GPU STREQUAL "metal")
|
||||
list(APPEND DS4_OBJS "${DS4_DIR}/ds4_metal.o")
|
||||
elseif(DS4_GPU STREQUAL "cpu")
|
||||
|
||||
@@ -1,10 +1,10 @@
|
||||
# ds4 backend Makefile.
|
||||
#
|
||||
# Upstream pin lives below as DS4_VERSION?=54b36ed9ba42da31b24f2d1a5feb075c2475dbb1
|
||||
# Upstream pin lives below as DS4_VERSION?=84cc882352757baf628a1776badf7cc54d584e28
|
||||
# (.github/bump_deps.sh) can find and update it - matches the
|
||||
# llama-cpp / ik-llama-cpp / turboquant convention.
|
||||
|
||||
DS4_VERSION?=54b36ed9ba42da31b24f2d1a5feb075c2475dbb1
|
||||
DS4_VERSION?=84cc882352757baf628a1776badf7cc54d584e28
|
||||
DS4_REPO?=https://github.com/antirez/ds4
|
||||
|
||||
CURRENT_MAKEFILE_DIR := $(dir $(abspath $(lastword $(MAKEFILE_LIST))))
|
||||
@@ -23,7 +23,9 @@ CMAKE_ARGS ?= -DCMAKE_BUILD_TYPE=Release
|
||||
# are shared by every GPU mode, so append them unconditionally below.
|
||||
ifeq ($(BUILD_TYPE),cublas)
|
||||
CMAKE_ARGS += -DDS4_GPU=cuda
|
||||
DS4_OBJ_TARGET := ds4.o ds4_cuda.o ds4_distributed.o ds4_tp.o ds4_ssd.o ds4_layer_pack.o
|
||||
DS4_OBJ_TARGET := ds4.o ds4_cuda.o ds4_distributed.o ds4_tp.o ds4_ssd.o ds4_layer_pack.o \
|
||||
cuda/mmq/ds4_ggml_stubs.o cuda/mmq/ds4_mmq.o cuda/mmq/ds4_mmq_d2r.o \
|
||||
cuda/mmq/quantize.o cuda/mmq/mmid.o cuda/mmq/mmvq.o cuda/mmq/ds4_repack.o
|
||||
else ifeq ($(UNAME_S),Darwin)
|
||||
CMAKE_ARGS += -DDS4_GPU=metal
|
||||
DS4_OBJ_TARGET := ds4.o ds4_metal.o ds4_distributed.o ds4_tp.o ds4_ssd.o ds4_layer_pack.o
|
||||
@@ -55,7 +57,7 @@ ds4:
|
||||
# the right per-platform compile flags (Objective-C/Metal on Darwin, nvcc on Linux+CUDA).
|
||||
ds4/ds4.o: ds4
|
||||
ifeq ($(BUILD_TYPE),cublas)
|
||||
+$(MAKE) -C ds4 ds4.o ds4_cuda.o ds4_distributed.o ds4_tp.o ds4_ssd.o ds4_layer_pack.o
|
||||
+$(MAKE) -C ds4 $(DS4_OBJ_TARGET)
|
||||
else ifeq ($(UNAME_S),Darwin)
|
||||
+$(MAKE) -C ds4 ds4.o ds4_metal.o ds4_distributed.o ds4_tp.o ds4_ssd.o ds4_layer_pack.o
|
||||
else
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
|
||||
IK_LLAMA_VERSION?=cb9147fd0d9c08a9a84eee5ac405a73f4e10e3e1
|
||||
IK_LLAMA_VERSION?=8337e4cd3861406fc04e0854b1409cd1b027fbc9
|
||||
LLAMA_REPO?=https://github.com/ikawrakow/ik_llama.cpp
|
||||
|
||||
CMAKE_ARGS?=
|
||||
|
||||
@@ -2565,6 +2565,7 @@ public:
|
||||
grpc::Status Embedding(ServerContext* context, const backend::PredictOptions* request, backend::EmbeddingResult* embeddingResult) {
|
||||
auto identity = checkModelIdentity(request);
|
||||
if (!identity.ok()) return identity;
|
||||
embeddingResult->set_layout(backend::EMBEDDING_LAYOUT_FINAL);
|
||||
json data = parse_options(false, request, llama);
|
||||
const int task_id = llama.queue_tasks.get_new_id();
|
||||
llama.queue_results.add_waiting_task_id(task_id);
|
||||
|
||||
@@ -115,4 +115,14 @@ if(LLAMA_GRPC_BUILD_TESTS)
|
||||
target_include_directories(passthrough_options_test PRIVATE ${CMAKE_CURRENT_SOURCE_DIR})
|
||||
target_compile_features(passthrough_options_test PRIVATE cxx_std_17)
|
||||
add_test(NAME passthrough_options_test COMMAND passthrough_options_test)
|
||||
|
||||
add_executable(tts_request_options_test tts_request_options_test.cpp tts_request_options.h)
|
||||
target_include_directories(tts_request_options_test PRIVATE ${CMAKE_CURRENT_SOURCE_DIR})
|
||||
target_compile_features(tts_request_options_test PRIVATE cxx_std_17)
|
||||
add_test(NAME tts_request_options_test COMMAND tts_request_options_test)
|
||||
|
||||
add_executable(thread_params_test thread_params_test.cpp thread_params.h)
|
||||
target_include_directories(thread_params_test PRIVATE ${CMAKE_CURRENT_SOURCE_DIR})
|
||||
target_compile_features(thread_params_test PRIVATE cxx_std_17)
|
||||
add_test(NAME thread_params_test COMMAND thread_params_test)
|
||||
endif()
|
||||
@@ -1,5 +1,5 @@
|
||||
|
||||
LLAMA_VERSION?=221f0f6356efe2260023208365705ec5d5a7c8f5
|
||||
LLAMA_VERSION?=d59d455fd8ea09e5a2e87ce2a9d668267ffb5ccd
|
||||
LLAMA_REPO?=https://github.com/ggerganov/llama.cpp
|
||||
|
||||
CMAKE_ARGS?=
|
||||
|
||||
Executable
+43
@@ -0,0 +1,43 @@
|
||||
#!/bin/bash
|
||||
# Mark a copied gRPC server as targeting a llama.cpp fork that does not carry
|
||||
# LocalAI's SERVER_TASK_TYPE_TTS patch. The RPCs remain present in the shared
|
||||
# protobuf service, but respond with UNIMPLEMENTED instead of referencing
|
||||
# server task types and mtmd gen-audio APIs absent from those forks.
|
||||
|
||||
set -euo pipefail
|
||||
|
||||
if [[ $# -ne 1 ]]; then
|
||||
echo "usage: $0 <grpc-server.cpp>" >&2
|
||||
exit 2
|
||||
fi
|
||||
|
||||
SRC=$1
|
||||
|
||||
if [[ ! -f "$SRC" ]]; then
|
||||
echo "grpc-server.cpp not found at $SRC" >&2
|
||||
exit 2
|
||||
fi
|
||||
|
||||
if grep -q '^#define LOCALAI_LLAMA_CPP_NO_TTS_TASK' "$SRC"; then
|
||||
echo "==> $SRC already disables the LocalAI TTS task, skipping"
|
||||
exit 0
|
||||
fi
|
||||
|
||||
awk '
|
||||
!done && /^#include/ {
|
||||
print "#define LOCALAI_LLAMA_CPP_NO_TTS_TASK 1"
|
||||
print "// ^ injected by disable-tts-task.sh for an unpatched llama.cpp fork"
|
||||
print ""
|
||||
done = 1
|
||||
}
|
||||
{ print }
|
||||
END {
|
||||
if (!done) {
|
||||
print "disable-tts-task.sh: no #include anchor found" > "/dev/stderr"
|
||||
exit 1
|
||||
}
|
||||
}
|
||||
' "$SRC" > "$SRC.tmp"
|
||||
mv "$SRC.tmp" "$SRC"
|
||||
|
||||
echo "==> LocalAI TTS task disabled in $SRC"
|
||||
@@ -53,8 +53,10 @@
|
||||
#include "arg.h"
|
||||
#include "chat-auto-parser.h"
|
||||
#include "llama_compat.h" // fork-skew switches, generated by prepare.sh
|
||||
#include "thread_params.h"
|
||||
#include "message_content.h"
|
||||
#include "passthrough_options.h"
|
||||
#include "tts_request_options.h"
|
||||
#include <getopt.h>
|
||||
#include <grpcpp/ext/proto_server_reflection_plugin.h>
|
||||
#include <grpcpp/grpcpp.h>
|
||||
@@ -65,6 +67,7 @@
|
||||
#include <atomic>
|
||||
#include <cmath>
|
||||
#include <cstdlib>
|
||||
#include <cstring>
|
||||
#include <fstream>
|
||||
#include <iterator>
|
||||
#include <list>
|
||||
@@ -233,7 +236,15 @@ json parse_options(bool streaming, const backend::PredictOptions* predict, const
|
||||
data["typical_p"] = predict->typicalp();
|
||||
data["temperature"] = predict->temperature();
|
||||
data["repeat_last_n"] = predict->repeat();
|
||||
data["repeat_penalty"] = predict->penalty();
|
||||
// PredictOptions.Penalty is a bare proto float, so a caller that names no
|
||||
// repetition penalty sends 0 rather than omitting the field. Since
|
||||
// llama.cpp 9de0fcf2b, common_sampler_init() rejects a non-positive
|
||||
// penalty_repeat outright (it would divide logits by zero), which turned
|
||||
// every such request into "Failed to initialize samplers". Treat 0 as
|
||||
// "unset" and leave llama.cpp's own neutral default in place.
|
||||
if (predict->penalty() > 0.0f) {
|
||||
data["repeat_penalty"] = predict->penalty();
|
||||
}
|
||||
data["frequency_penalty"] = predict->frequencypenalty();
|
||||
data["presence_penalty"] = predict->presencepenalty();
|
||||
data["mirostat"] = predict->mirostat();
|
||||
@@ -1402,6 +1413,12 @@ static void params_parse(server_context& /*ctx_server*/, const backend::ModelOpt
|
||||
passthrough_draft_gpu_layers);
|
||||
}
|
||||
|
||||
// The library initializer now creates both threadpools before the server
|
||||
// can apply llama_context's fallback for the -1 batch-thread sentinel.
|
||||
params.cpuparams_batch.n_threads = llama_grpc::resolve_batch_threads(
|
||||
params.cpuparams_batch.n_threads,
|
||||
params.cpuparams.n_threads);
|
||||
|
||||
#ifndef LOCALAI_LLAMA_CPP_NO_SCORE_TASK
|
||||
// Score-task suffix forking: reserve seq ids (and recurrent-state cells)
|
||||
// beyond the slots so one scoring call decodes all candidate tails in a
|
||||
@@ -1445,6 +1462,26 @@ static void params_parse(server_context& /*ctx_server*/, const backend::ModelOpt
|
||||
}
|
||||
|
||||
|
||||
#ifndef LOCALAI_LLAMA_CPP_NO_TTS_TASK
|
||||
// MTMD_HELPER_GEN_AUDIO_OUTTYPE_PCM hands back raw float32 samples, but the
|
||||
// WAV header core/backend/tts.go builds around the streamed chunks announces
|
||||
// 16-bit samples, so the wire has to carry s16 or the client decodes floats as
|
||||
// integers and hears noise. The scaling matches write_wav16() in
|
||||
// tools/mtmd/mtmd-helper-gen.cpp, which is what the non-streaming path writes.
|
||||
static std::string tts_pcm_f32_to_s16(const std::string & samples) {
|
||||
const size_t n = samples.size() / sizeof(float);
|
||||
std::string out;
|
||||
out.resize(n * sizeof(int16_t));
|
||||
for (size_t i = 0; i < n; i++) {
|
||||
float v = 0.0f;
|
||||
std::memcpy(&v, samples.data() + i * sizeof(float), sizeof(float));
|
||||
const int16_t s = (int16_t) (std::max(-1.0f, std::min(1.0f, v)) * 32767.0f);
|
||||
std::memcpy(&out[i * sizeof(int16_t)], &s, sizeof(int16_t));
|
||||
}
|
||||
return out;
|
||||
}
|
||||
#endif
|
||||
|
||||
// GRPC Server start
|
||||
class BackendServiceImpl final : public backend::Backend::Service {
|
||||
private:
|
||||
@@ -2089,15 +2126,23 @@ public:
|
||||
|
||||
task.tokens = std::move(inputs[i]);
|
||||
#ifdef LOCALAI_HAS_SERVER_SCHEMA
|
||||
// The schema evaluator no longer takes the per-slot n_ctx: upstream
|
||||
// dropped the parameter and server-schema stopped consulting n_ctx at
|
||||
// all, leaving the context bound to the slot. Forks that predate the
|
||||
// server-schema split still expect it, so only this branch loses it.
|
||||
task.params = server_schema::eval_llama_cmpl_schema(
|
||||
ctx_server.impl->vocab,
|
||||
params_base,
|
||||
ctx_server.get_meta().logit_bias_eog,
|
||||
data);
|
||||
#else
|
||||
task.params = server_task::params_from_json_cmpl(
|
||||
#endif
|
||||
ctx_server.impl->vocab,
|
||||
params_base,
|
||||
ctx_server.get_meta().slot_n_ctx,
|
||||
ctx_server.get_meta().logit_bias_eog,
|
||||
data);
|
||||
#endif
|
||||
task.id_slot = json_value(data, "id_slot", -1);
|
||||
|
||||
// OAI-compat: enable autoparser (PEG-based chat parsing) so that
|
||||
@@ -2659,15 +2704,23 @@ public:
|
||||
|
||||
task.tokens = std::move(inputs[i]);
|
||||
#ifdef LOCALAI_HAS_SERVER_SCHEMA
|
||||
// The schema evaluator no longer takes the per-slot n_ctx: upstream
|
||||
// dropped the parameter and server-schema stopped consulting n_ctx at
|
||||
// all, leaving the context bound to the slot. Forks that predate the
|
||||
// server-schema split still expect it, so only this branch loses it.
|
||||
task.params = server_schema::eval_llama_cmpl_schema(
|
||||
ctx_server.impl->vocab,
|
||||
params_base,
|
||||
ctx_server.get_meta().logit_bias_eog,
|
||||
data);
|
||||
#else
|
||||
task.params = server_task::params_from_json_cmpl(
|
||||
#endif
|
||||
ctx_server.impl->vocab,
|
||||
params_base,
|
||||
ctx_server.get_meta().slot_n_ctx,
|
||||
ctx_server.get_meta().logit_bias_eog,
|
||||
data);
|
||||
#endif
|
||||
task.id_slot = json_value(data, "id_slot", -1);
|
||||
|
||||
// OAI-compat: enable autoparser (PEG-based chat parsing) so that
|
||||
@@ -2865,42 +2918,40 @@ public:
|
||||
return grpc::Status(grpc::StatusCode::INTERNAL, all_results.error->to_json().value("message", "Error in receiving results"));
|
||||
}
|
||||
|
||||
// Collect responses
|
||||
json responses = json::array();
|
||||
// Extract the embeddings typed, straight from the task results (no
|
||||
// JSON round-trip), and report the payload shape alongside the same
|
||||
// flat float array as before: dim is the embedding width, tokens the
|
||||
// number of vectors packed into `embeddings` (1 per prompt when the
|
||||
// server pooled, one per token with pooling:none; summed across
|
||||
// prompts if the request carried several), prompt_tokens the prompt
|
||||
// tokens evaluated, for usage accounting. Consumers seeing 0/0 know
|
||||
// the backend predates shape reporting.
|
||||
int32_t n_vectors = 0;
|
||||
int32_t dim = 0;
|
||||
int32_t prompt_tokens = 0;
|
||||
for (auto & res : all_results.results) {
|
||||
GGML_ASSERT(dynamic_cast<server_task_result_embd*>(res.get()) != nullptr);
|
||||
responses.push_back(res->to_json());
|
||||
}
|
||||
|
||||
std::cout << "[DEBUG] Responses size: " << responses.size() << std::endl;
|
||||
|
||||
// Process the responses and extract embeddings
|
||||
for (const auto & response_elem : responses) {
|
||||
// Check if the response has an "embedding" field
|
||||
if (response_elem.contains("embedding")) {
|
||||
json embedding_data = json_value(response_elem, "embedding", json::array());
|
||||
|
||||
if (embedding_data.is_array() && !embedding_data.empty()) {
|
||||
for (const auto & embedding_vector : embedding_data) {
|
||||
if (embedding_vector.is_array()) {
|
||||
for (const auto & embedding_value : embedding_vector) {
|
||||
embeddingResult->add_embeddings(embedding_value.get<float>());
|
||||
}
|
||||
}
|
||||
}
|
||||
auto * embd_res = dynamic_cast<server_task_result_embd*>(res.get());
|
||||
GGML_ASSERT(embd_res != nullptr);
|
||||
prompt_tokens += embd_res->n_tokens;
|
||||
for (const auto & vec : embd_res->embedding) {
|
||||
for (const float value : vec) {
|
||||
embeddingResult->add_embeddings(value);
|
||||
}
|
||||
} else {
|
||||
// Check if the response itself contains the embedding data directly
|
||||
if (response_elem.is_array()) {
|
||||
for (const auto & embedding_value : response_elem) {
|
||||
embeddingResult->add_embeddings(embedding_value.get<float>());
|
||||
}
|
||||
if (!vec.empty()) {
|
||||
n_vectors++;
|
||||
dim = (int32_t) vec.size();
|
||||
}
|
||||
}
|
||||
}
|
||||
embeddingResult->set_tokens(n_vectors);
|
||||
embeddingResult->set_dim(dim);
|
||||
embeddingResult->set_prompt_tokens(prompt_tokens);
|
||||
embeddingResult->set_layout(
|
||||
llama_pooling_type(ctx_server.get_llama_context()) == LLAMA_POOLING_TYPE_NONE
|
||||
? backend::EMBEDDING_LAYOUT_PER_TOKEN
|
||||
: backend::EMBEDDING_LAYOUT_FINAL);
|
||||
|
||||
|
||||
|
||||
std::cout << "[DEBUG] Embedding vectors: " << n_vectors << " x " << dim << std::endl;
|
||||
|
||||
return grpc::Status::OK;
|
||||
}
|
||||
@@ -2994,6 +3045,229 @@ public:
|
||||
return grpc::Status::OK;
|
||||
}
|
||||
|
||||
#ifndef LOCALAI_LLAMA_CPP_NO_TTS_TASK
|
||||
// Builds the shared TTS task from a request. Returns a non-OK status and
|
||||
// leaves `task` untouched when the request is malformed or the loaded model
|
||||
// cannot synthesise audio.
|
||||
grpc::Status prepareTTSTask(const backend::TTSRequest* request, bool stream, server_task & task) {
|
||||
if (!ctx_server.get_meta().has_cap_tts) {
|
||||
return grpc::Status(grpc::StatusCode::UNIMPLEMENTED,
|
||||
"the loaded model does not support audio generation (no gen-audio mmproj)");
|
||||
}
|
||||
|
||||
std::map<std::string, std::string> params(request->params().begin(), request->params().end());
|
||||
const auto opts = llama_grpc::parse_tts_request_options(
|
||||
request->text(),
|
||||
request->voice(),
|
||||
request->has_language() ? request->language() : std::string(),
|
||||
params);
|
||||
if (!opts.ok) {
|
||||
return grpc::Status(grpc::StatusCode::INVALID_ARGUMENT, opts.error);
|
||||
}
|
||||
|
||||
auto wrapper = mtmd_helper_bitmap_init_from_file(ctx_server.impl->mctx, opts.voice_path.c_str(), false);
|
||||
if (!wrapper.bitmap) {
|
||||
return grpc::Status(grpc::StatusCode::INVALID_ARGUMENT,
|
||||
"failed to read speaker reference audio: " + opts.voice_path);
|
||||
}
|
||||
|
||||
task.tts_inp.set_prompt(opts.text);
|
||||
// core/backend/tts.go always sets TTSRequest.language, so has_language()
|
||||
// is true even when the caller named no language and the string is empty.
|
||||
// gen_audio::inp::get() already maps a stored blank to nullptr, so this
|
||||
// guard is behavior-preserving rather than behavior-fixing. It is kept
|
||||
// so the "unset" intent is visible at the call site instead of resting
|
||||
// on a detail of the helper.
|
||||
if (!opts.language.empty()) {
|
||||
task.tts_inp.set_lang(opts.language);
|
||||
}
|
||||
task.tts_inp.set_speaker_ref(mtmd::bitmap_ptr(wrapper.bitmap));
|
||||
task.tts_inp.data.top_k = opts.top_k;
|
||||
task.tts_inp.data.top_p = opts.top_p;
|
||||
task.tts_inp.data.stream = stream;
|
||||
task.tts_inp.data.out_type = stream
|
||||
? MTMD_HELPER_GEN_AUDIO_OUTTYPE_PCM // Go prepends its own WAV header, see core/backend/tts.go
|
||||
: MTMD_HELPER_GEN_AUDIO_OUTTYPE_WAV;
|
||||
|
||||
task.params.stream = stream;
|
||||
// -1 keeps upstream's 512-frame default. The model does not always emit
|
||||
// its codec EOS, so a short input can otherwise generate the full cap.
|
||||
task.params.n_predict = opts.max_frames > 0 ? opts.max_frames : -1;
|
||||
task.params.sampling = params_base.sampling;
|
||||
// Both values mirror upstream's draft POST /tts handler. Note that the
|
||||
// pair is INERT at this pin: llama_sampler_init_penalties() clamps
|
||||
// penalty_last_n with std::max(penalty_last_n, 0), so -1 means "off",
|
||||
// not "the whole generation", and the penalty sampler is then built
|
||||
// disabled. No repetition penalty is actually applied.
|
||||
//
|
||||
// That is deliberate. Dropping the second line lets the sampling
|
||||
// default of 64 apply and genuinely engages the 1.05 penalty, which was
|
||||
// measured here against the model's habit of never emitting its codec
|
||||
// EOS and running to the frame cap: 0 of 15 short requests ran away
|
||||
// with the penalty inert, 1 of 15 with it active over the last 64
|
||||
// tokens. It does not fix the runaway, so the line stays for parity
|
||||
// with the draft. Use max_frames to bound the output instead.
|
||||
task.params.sampling.penalty_repeat = 1.05f;
|
||||
task.params.sampling.penalty_last_n = -1;
|
||||
if (opts.top_k > 0) {
|
||||
task.params.sampling.top_k = opts.top_k;
|
||||
}
|
||||
if (opts.top_p > 0) {
|
||||
task.params.sampling.top_p = opts.top_p;
|
||||
}
|
||||
|
||||
return grpc::Status::OK;
|
||||
}
|
||||
|
||||
grpc::Status TTS(ServerContext* context, const backend::TTSRequest* request, backend::Result* result) override {
|
||||
auto auth = checkAuth(context);
|
||||
if (!auth.ok()) return auth;
|
||||
auto identity = checkModelIdentity(request);
|
||||
if (!identity.ok()) return identity;
|
||||
if (params_base.model.path.empty()) {
|
||||
return grpc::Status(grpc::StatusCode::FAILED_PRECONDITION, "Model not loaded");
|
||||
}
|
||||
if (request->dst().empty()) {
|
||||
return grpc::Status(grpc::StatusCode::INVALID_ARGUMENT, "dst must name an output file path");
|
||||
}
|
||||
|
||||
server_task task(SERVER_TASK_TYPE_TTS);
|
||||
auto prepared = prepareTTSTask(request, /* stream= */ false, task);
|
||||
if (!prepared.ok()) return prepared;
|
||||
|
||||
auto rd = ctx_server.get_response_reader();
|
||||
task.id = rd.get_new_id();
|
||||
rd.post_task(std::move(task));
|
||||
|
||||
auto should_stop = [context]() { return context->IsCancelled(); };
|
||||
|
||||
std::string audio;
|
||||
while (true) {
|
||||
auto res = rd.next(should_stop);
|
||||
if (!res) {
|
||||
return grpc::Status(grpc::StatusCode::CANCELLED, "TTS request cancelled");
|
||||
}
|
||||
if (res->is_error()) {
|
||||
return grpc::Status(grpc::StatusCode::INTERNAL, res->to_json().dump());
|
||||
}
|
||||
auto * tts_res = dynamic_cast<server_task_result_tts *>(res.get());
|
||||
if (tts_res == nullptr) {
|
||||
return grpc::Status(grpc::StatusCode::INTERNAL, "unexpected result type for a TTS task");
|
||||
}
|
||||
audio.append(tts_res->audio);
|
||||
if (tts_res->final) {
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
std::ofstream out(request->dst(), std::ios::binary | std::ios::trunc);
|
||||
if (!out) {
|
||||
return grpc::Status(grpc::StatusCode::INTERNAL, "failed to open output file: " + request->dst());
|
||||
}
|
||||
out.write(audio.data(), (std::streamsize) audio.size());
|
||||
if (!out) {
|
||||
return grpc::Status(grpc::StatusCode::INTERNAL, "failed to write output file: " + request->dst());
|
||||
}
|
||||
// Buffered data is flushed here, so a full disk or a failing device can
|
||||
// surface for the first time on close. Reporting success then would
|
||||
// leave a truncated file behind under the name the caller will read.
|
||||
out.close();
|
||||
if (!out) {
|
||||
return grpc::Status(grpc::StatusCode::INTERNAL, "failed to close output file: " + request->dst());
|
||||
}
|
||||
|
||||
result->set_success(true);
|
||||
result->set_message("TTS audio generated");
|
||||
return grpc::Status::OK;
|
||||
}
|
||||
|
||||
grpc::Status TTSStream(ServerContext* context, const backend::TTSRequest* request, grpc::ServerWriter<backend::Reply>* writer) override {
|
||||
auto auth = checkAuth(context);
|
||||
if (!auth.ok()) return auth;
|
||||
auto identity = checkModelIdentity(request);
|
||||
if (!identity.ok()) return identity;
|
||||
if (params_base.model.path.empty()) {
|
||||
return grpc::Status(grpc::StatusCode::FAILED_PRECONDITION, "Model not loaded");
|
||||
}
|
||||
|
||||
server_task task(SERVER_TASK_TYPE_TTS);
|
||||
auto prepared = prepareTTSTask(request, /* stream= */ true, task);
|
||||
if (!prepared.ok()) return prepared;
|
||||
|
||||
auto rd = ctx_server.get_response_reader();
|
||||
task.id = rd.get_new_id();
|
||||
rd.post_task(std::move(task));
|
||||
|
||||
auto should_stop = [context]() { return context->IsCancelled(); };
|
||||
|
||||
// core/backend/tts.go:ModelTTSStream builds the WAV header itself from
|
||||
// the sample rate in the first reply's Message, then concatenates every
|
||||
// Reply.Audio verbatim. So the rate goes out once, up front, and the
|
||||
// chunks stay raw PCM.
|
||||
//
|
||||
// Send it before draining rather than off the first audio result: a
|
||||
// chunk needs a whole 72-frame window, about 5.8 s of audio and far
|
||||
// longer in wall time on CPU, and the Go side cannot emit the WAV
|
||||
// header until this reply lands. Waiting would hold the client at zero
|
||||
// bytes for that entire stretch. The rate is a property of the loaded
|
||||
// model, available synchronously, so there is nothing to wait for.
|
||||
{
|
||||
backend::Reply header;
|
||||
const json info = { {"sample_rate", mtmd_gen_audio_get_info(ctx_server.impl->mctx).sample_rate} };
|
||||
header.set_message(info.dump());
|
||||
if (!writer->Write(header)) {
|
||||
return grpc::Status(grpc::StatusCode::CANCELLED, "client closed the TTS stream");
|
||||
}
|
||||
}
|
||||
|
||||
while (true) {
|
||||
auto res = rd.next(should_stop);
|
||||
if (!res) {
|
||||
return grpc::Status(grpc::StatusCode::CANCELLED, "TTS request cancelled");
|
||||
}
|
||||
if (res->is_error()) {
|
||||
return grpc::Status(grpc::StatusCode::INTERNAL, res->to_json().dump());
|
||||
}
|
||||
auto * tts_res = dynamic_cast<server_task_result_tts *>(res.get());
|
||||
if (tts_res == nullptr) {
|
||||
return grpc::Status(grpc::StatusCode::INTERNAL, "unexpected result type for a TTS task");
|
||||
}
|
||||
|
||||
if (!tts_res->audio.empty()) {
|
||||
backend::Reply chunk;
|
||||
chunk.set_audio(tts_pcm_f32_to_s16(tts_res->audio));
|
||||
if (!writer->Write(chunk)) {
|
||||
return grpc::Status(grpc::StatusCode::CANCELLED, "client closed the TTS stream");
|
||||
}
|
||||
}
|
||||
|
||||
if (tts_res->final) {
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
return grpc::Status::OK;
|
||||
}
|
||||
#else
|
||||
grpc::Status TTS(ServerContext* context, const backend::TTSRequest* request, backend::Result* result) override {
|
||||
auto auth = checkAuth(context);
|
||||
if (!auth.ok()) return auth;
|
||||
(void) request;
|
||||
(void) result;
|
||||
return grpc::Status(grpc::StatusCode::UNIMPLEMENTED,
|
||||
"TTS is unavailable in this llama.cpp fork backend");
|
||||
}
|
||||
|
||||
grpc::Status TTSStream(ServerContext* context, const backend::TTSRequest* request, grpc::ServerWriter<backend::Reply>* writer) override {
|
||||
auto auth = checkAuth(context);
|
||||
if (!auth.ok()) return auth;
|
||||
(void) request;
|
||||
(void) writer;
|
||||
return grpc::Status(grpc::StatusCode::UNIMPLEMENTED,
|
||||
"TTSStream is unavailable in this llama.cpp fork backend");
|
||||
}
|
||||
#endif
|
||||
|
||||
// Score returns the model's joint log-probability of each candidate
|
||||
// continuation given a shared prompt.
|
||||
//
|
||||
@@ -3328,9 +3602,15 @@ public:
|
||||
// Populate the response with metrics
|
||||
response->set_slot_id(0);
|
||||
response->set_prompt_json_for_slot("");
|
||||
#if LOCALAI_HAS_SERVER_METRICS
|
||||
response->set_tokens_per_second(res_metrics->metrics.prompt_bucket.n_per_second());
|
||||
response->set_tokens_generated(res_metrics->metrics.predict.count);
|
||||
response->set_prompt_tokens_processed(res_metrics->metrics.prompt.count);
|
||||
#else
|
||||
response->set_tokens_per_second(res_metrics->n_prompt_tokens_processed ? 1.e3 / res_metrics->t_prompt_processing * res_metrics->n_prompt_tokens_processed : 0.);
|
||||
response->set_tokens_generated(res_metrics->n_tokens_predicted_total);
|
||||
response->set_prompt_tokens_processed(res_metrics->n_prompt_tokens_processed_total);
|
||||
#endif
|
||||
|
||||
|
||||
return grpc::Status::OK;
|
||||
|
||||
@@ -1,8 +1,21 @@
|
||||
From 75220a0d74892e3315f4042274b1efa6195868d8 Mon Sep 17 00:00:00 2001
|
||||
From: Codex <codex@local>
|
||||
Date: Mon, 10 Aug 2026 23:05:52 +0000
|
||||
Subject: [PATCH 1/2] score-patch
|
||||
|
||||
---
|
||||
common/common.cpp | 6 +-
|
||||
common/common.h | 3 +
|
||||
tools/CMakeLists.txt | 1 +
|
||||
tools/server/server-context.cpp | 358 +++++++++++++++++++++++++++++++-
|
||||
tools/server/server-task.h | 47 +++++
|
||||
5 files changed, 406 insertions(+), 9 deletions(-)
|
||||
|
||||
diff --git a/common/common.cpp b/common/common.cpp
|
||||
index 8f13217..fc584e1 100644
|
||||
index 2e3f14c..0cec0dc 100644
|
||||
--- a/common/common.cpp
|
||||
+++ b/common/common.cpp
|
||||
@@ -1591,8 +1591,10 @@ struct llama_context_params common_context_params_to_llama(const common_params &
|
||||
@@ -1636,8 +1636,10 @@ struct llama_context_params common_context_params_to_llama(const common_params &
|
||||
auto cparams = llama_context_default_params();
|
||||
|
||||
cparams.n_ctx = params.n_ctx;
|
||||
@@ -13,13 +26,13 @@ index 8f13217..fc584e1 100644
|
||||
+ cparams.n_seq_max = params.n_parallel + params.n_seq_score_forks;
|
||||
+ cparams.n_rs_seq = std::max(params.speculative.need_n_rs_seq(), (uint32_t) std::max(0, params.n_rs_seq));
|
||||
cparams.n_outputs_max = std::max(params.n_outputs_max, 0);
|
||||
cparams.n_outputs_max_per_seq = std::max(params.n_outputs_max_per_seq, 0);
|
||||
cparams.n_batch = params.n_batch;
|
||||
cparams.n_ubatch = params.n_ubatch;
|
||||
diff --git a/common/common.h b/common/common.h
|
||||
index bffc176..e313bd6 100644
|
||||
index 878534d..4001df2 100644
|
||||
--- a/common/common.h
|
||||
+++ b/common/common.h
|
||||
@@ -455,6 +455,9 @@ struct common_params {
|
||||
@@ -445,6 +445,9 @@ struct common_params {
|
||||
int32_t n_keep = 0; // number of tokens to keep from initial prompt
|
||||
int32_t n_chunks = -1; // max number of chunks to process (-1 = unlimited)
|
||||
int32_t n_parallel = 1; // number of parallel sequences to decode
|
||||
@@ -28,7 +41,7 @@ index bffc176..e313bd6 100644
|
||||
+ bool score_enabled = false; // reserve server resources for the Score task type
|
||||
int32_t n_sequences = 1; // number of sequences to decode
|
||||
int32_t n_outputs_max = 0; // max outputs in a batch (0 = n_batch)
|
||||
int32_t grp_attn_n = 1; // group-attention factor
|
||||
int32_t n_outputs_max_per_seq = 1; // max outputs per sequence
|
||||
diff --git a/tools/CMakeLists.txt b/tools/CMakeLists.txt
|
||||
index 780df32..1d2fe8f 100644
|
||||
--- a/tools/CMakeLists.txt
|
||||
@@ -39,28 +52,24 @@ index 780df32..1d2fe8f 100644
|
||||
endif()
|
||||
+add_subdirectory(grpc-server)
|
||||
diff --git a/tools/server/server-context.cpp b/tools/server/server-context.cpp
|
||||
index 715477e..de5bed8 100644
|
||||
index 3b5f6a1..d0e18e6 100644
|
||||
--- a/tools/server/server-context.cpp
|
||||
+++ b/tools/server/server-context.cpp
|
||||
@@ -49,7 +49,16 @@ static uint32_t server_n_outputs_max(const common_params & params) {
|
||||
@@ -48,6 +48,13 @@ static common_speculative_output_limits server_output_limits(const common_params
|
||||
auto result = common_speculative_get_output_limits(
|
||||
params.n_batch, params.n_parallel, common_speculative_n_max(¶ms.speculative));
|
||||
|
||||
const uint32_t n_outputs_per_seq = 1 + common_speculative_n_max(¶ms.speculative);
|
||||
|
||||
- const uint64_t n_outputs = (uint64_t) params.n_parallel * n_outputs_per_seq;
|
||||
+ // score tasks (SERVER_TASK_TYPE_SCORE) output logits for every candidate
|
||||
+ // token, so reserve room for a bounded candidate tail per parallel slot
|
||||
+ if (!params.score_enabled) {
|
||||
+ return std::max<uint32_t>(1, std::min<uint64_t>(n_batch,
|
||||
+ (uint64_t) params.n_parallel * n_outputs_per_seq));
|
||||
+ // Score tasks output logits for every candidate token, so reserve room
|
||||
+ // for a bounded candidate tail per parallel slot.
|
||||
+ if (params.score_enabled) {
|
||||
+ result.per_seq = std::max<int32_t>(result.per_seq, 1 + SERVER_SCORE_MAX_CAND_TOKENS);
|
||||
+ result.total = std::min<int32_t>(params.n_batch, params.n_parallel * result.per_seq);
|
||||
+ }
|
||||
+
|
||||
+ const uint32_t n_outputs_score_seq = 1 + SERVER_SCORE_MAX_CAND_TOKENS;
|
||||
+
|
||||
+ const uint64_t n_outputs = (uint64_t) params.n_parallel * std::max(n_outputs_per_seq, n_outputs_score_seq);
|
||||
|
||||
return std::max<uint32_t>(1, std::min<uint64_t>(n_batch, n_outputs));
|
||||
}
|
||||
@@ -202,6 +211,26 @@ struct server_slot {
|
||||
result.total = std::max<int32_t>(1, result.total);
|
||||
result.per_seq = std::max<int32_t>(1, result.per_seq);
|
||||
return result;
|
||||
@@ -239,6 +246,26 @@ struct server_slot {
|
||||
|
||||
std::vector<completion_token_output> generated_token_probs;
|
||||
|
||||
@@ -87,7 +96,7 @@ index 715477e..de5bed8 100644
|
||||
bool has_next_token = true;
|
||||
bool has_new_line = false;
|
||||
bool truncated = false;
|
||||
@@ -311,6 +340,10 @@ struct server_slot {
|
||||
@@ -341,6 +368,10 @@ struct server_slot {
|
||||
}
|
||||
generated_tokens.clear();
|
||||
generated_token_probs.clear();
|
||||
@@ -97,8 +106,8 @@ index 715477e..de5bed8 100644
|
||||
+ score_divergence = -1;
|
||||
json_schema = json();
|
||||
|
||||
// clear speculative decoding stats
|
||||
@@ -2205,6 +2238,229 @@ private:
|
||||
task_prev = std::move(task);
|
||||
@@ -2271,6 +2302,229 @@ private:
|
||||
queue_results.send(std::move(res));
|
||||
}
|
||||
|
||||
@@ -328,7 +337,7 @@ index 715477e..de5bed8 100644
|
||||
//
|
||||
// Functions to process the task
|
||||
//
|
||||
@@ -2341,6 +2597,7 @@ private:
|
||||
@@ -2407,6 +2661,7 @@ private:
|
||||
case SERVER_TASK_TYPE_INFILL:
|
||||
case SERVER_TASK_TYPE_EMBEDDING:
|
||||
case SERVER_TASK_TYPE_RERANK:
|
||||
@@ -336,7 +345,7 @@ index 715477e..de5bed8 100644
|
||||
{
|
||||
// special case: if input is provided via CLI, tokenize it first
|
||||
// otherwise, no need to tokenize as it's already done inside the HTTP thread
|
||||
@@ -2832,6 +3089,13 @@ private:
|
||||
@@ -2903,6 +3158,13 @@ private:
|
||||
break; // stop any further processing
|
||||
}
|
||||
}
|
||||
@@ -350,7 +359,7 @@ index 715477e..de5bed8 100644
|
||||
}
|
||||
|
||||
void pre_decode() {
|
||||
@@ -3154,6 +3418,16 @@ private:
|
||||
@@ -3222,6 +3484,16 @@ private:
|
||||
n_past = std::min(n_past, slot.alora_invocation_start - 1);
|
||||
}
|
||||
|
||||
@@ -367,7 +376,7 @@ index 715477e..de5bed8 100644
|
||||
const auto n_cache_reuse = slot.task->params.n_cache_reuse;
|
||||
|
||||
const bool can_cache_reuse =
|
||||
@@ -3395,8 +3669,12 @@ private:
|
||||
@@ -3455,8 +3727,12 @@ private:
|
||||
|
||||
bool do_checkpoint = params_base.n_ctx_checkpoints > 0;
|
||||
|
||||
@@ -382,7 +391,7 @@ index 715477e..de5bed8 100644
|
||||
|
||||
// make a checkpoint of the parts of the memory that cannot be rolled back.
|
||||
// checkpoints are created only if:
|
||||
@@ -3463,10 +3741,17 @@ private:
|
||||
@@ -3444,9 +3720,16 @@ private:
|
||||
// embedding requires all tokens in the batch to be output;
|
||||
// MTP also wants logits at every prompt position so the
|
||||
// streaming hook can mirror t_h_nextn into ctx_dft.
|
||||
@@ -395,16 +404,12 @@ index 715477e..de5bed8 100644
|
||||
+ slot.prompt.n_tokens() + 1 < slot.task->n_tokens();
|
||||
add_ok &= batch.add(slot.id,
|
||||
cur_tok,
|
||||
slot.prompt.tokens.pos_next(),
|
||||
- slot.need_embd());
|
||||
+ slot.need_embd() || need_score_logit);
|
||||
/* pos = */ slot.prompt.tokens.pos_next(),
|
||||
- /* output = */ slot.need_embd(),
|
||||
+ /* output = */ slot.need_embd() || need_score_logit,
|
||||
/* is_prompt = */ true);
|
||||
slot.prompt.tokens.push_back(cur_tok);
|
||||
|
||||
slot.n_prompt_tokens_processed++;
|
||||
@@ -3481,6 +3766,32 @@ private:
|
||||
}
|
||||
}
|
||||
|
||||
@@ -3454,2 +3737,28 @@ private:
|
||||
+ // score tasks: break at the shared-prompt boundary so the checkpoint
|
||||
+ // below lands exactly there — the other candidates of the same
|
||||
+ // scoring call re-process only their own tokens. Also break at the
|
||||
@@ -431,10 +436,9 @@ index 715477e..de5bed8 100644
|
||||
+ }
|
||||
+ }
|
||||
+
|
||||
// process the last few tokens of the prompt separately in order to allow for a checkpoint to be created.
|
||||
// create checkpoints that many tokens before the end of the prompt:
|
||||
// - 4 + n_ubatch
|
||||
@@ -3513,6 +3824,15 @@ private:
|
||||
// break at the last user message, or at user messages at least min step past the last checkpoint
|
||||
if (do_checkpoint && spans.is_user_start(slot.prompt.n_tokens())) {
|
||||
@@ -3573,6 +3882,15 @@ private:
|
||||
const bool is_user_start = spans.is_user_start(n_tokens_start);
|
||||
const bool is_last_user_message = n_tokens_start == last_user_pos;
|
||||
|
||||
@@ -450,7 +454,7 @@ index 715477e..de5bed8 100644
|
||||
// entire prompt has been processed
|
||||
if (slot.prompt.n_tokens() == slot.task->n_tokens()) {
|
||||
slot.state = SLOT_STATE_DONE_PROMPT;
|
||||
@@ -3528,8 +3848,8 @@ private:
|
||||
@@ -3588,8 +3906,8 @@ private:
|
||||
slot.init_sampler();
|
||||
} else {
|
||||
// skip ordinary mid-prompt checkpoints, unless the batch starts a user
|
||||
@@ -461,7 +465,7 @@ index 715477e..de5bed8 100644
|
||||
do_checkpoint = false;
|
||||
}
|
||||
}
|
||||
@@ -3546,10 +3866,10 @@ private:
|
||||
@@ -3606,10 +3924,10 @@ private:
|
||||
// do not checkpoint after mtmd chunks
|
||||
do_checkpoint = do_checkpoint && !has_mtmd;
|
||||
|
||||
@@ -474,7 +478,7 @@ index 715477e..de5bed8 100644
|
||||
n_tokens_start > slot.prompt.checkpoints.back().n_tokens + params_base.checkpoint_min_step);
|
||||
SLT_DBG(slot, "main/do_checkpoint = %s, pos_min = %d, pos_max = %d\n", do_checkpoint ? "yes" : "no", pos_min, pos_max);
|
||||
|
||||
@@ -3703,6 +4023,13 @@ private:
|
||||
@@ -3772,6 +4090,13 @@ private:
|
||||
}
|
||||
}
|
||||
|
||||
@@ -488,7 +492,7 @@ index 715477e..de5bed8 100644
|
||||
if (!is_inside_view(slot.i_batch)) {
|
||||
// the required token not in this sub-batch, skip
|
||||
return;
|
||||
@@ -3724,6 +4051,25 @@ private:
|
||||
@@ -3793,6 +4118,25 @@ private:
|
||||
return;
|
||||
}
|
||||
|
||||
@@ -515,7 +519,7 @@ index 715477e..de5bed8 100644
|
||||
|
||||
// prompt evaluated for next-token prediction
|
||||
diff --git a/tools/server/server-task.h b/tools/server/server-task.h
|
||||
index c3eea2e..fb3c178 100644
|
||||
index 6275ec7..5bedf19 100644
|
||||
--- a/tools/server/server-task.h
|
||||
+++ b/tools/server/server-task.h
|
||||
@@ -13,10 +13,25 @@
|
||||
@@ -597,3 +601,5 @@ index c3eea2e..fb3c178 100644
|
||||
struct server_task_result_error : server_task_result {
|
||||
error_type err_type = ERROR_TYPE_SERVER;
|
||||
std::string err_msg;
|
||||
--
|
||||
2.39.5
|
||||
@@ -0,0 +1,845 @@
|
||||
diff --git a/tools/mtmd/mtmd-helper-gen.cpp b/tools/mtmd/mtmd-helper-gen.cpp
|
||||
index 1c58d3ae1..196cbd433 100644
|
||||
--- a/tools/mtmd/mtmd-helper-gen.cpp
|
||||
+++ b/tools/mtmd/mtmd-helper-gen.cpp
|
||||
@@ -50,29 +50,38 @@ static llama_token find_special_token(const llama_vocab * vocab, const std::stri
|
||||
return LLAMA_TOKEN_NULL;
|
||||
}
|
||||
|
||||
+static void put_bytes(std::vector<char> & buf, const void * p, size_t n) {
|
||||
+ const char * c = (const char *) p;
|
||||
+ buf.insert(buf.end(), c, c + n);
|
||||
+}
|
||||
+
|
||||
+// data_sz == UINT32_MAX writes the "unknown length" sentinel (streaming), same as ffmpeg does on a pipe
|
||||
+static void write_wav16_header(std::vector<char> & buf, uint32_t data_sz, int32_t rate) {
|
||||
+ const uint32_t riff_sz = data_sz == UINT32_MAX ? UINT32_MAX : 36 + data_sz;
|
||||
+ const uint32_t fmt_sz = 16, byte_rate = (uint32_t) rate * 2;
|
||||
+ const uint16_t fmt = 1, ch = 1, align = 2, bits = 16;
|
||||
+ const uint32_t rate32 = (uint32_t) rate;
|
||||
+ put_bytes(buf, "RIFF", 4); put_bytes(buf, &riff_sz, 4); put_bytes(buf, "WAVE", 4);
|
||||
+ put_bytes(buf, "fmt ", 4); put_bytes(buf, &fmt_sz, 4);
|
||||
+ put_bytes(buf, &fmt, 2); put_bytes(buf, &ch, 2); put_bytes(buf, &rate32, 4);
|
||||
+ put_bytes(buf, &byte_rate, 4); put_bytes(buf, &align, 2); put_bytes(buf, &bits, 2);
|
||||
+ put_bytes(buf, "data", 4); put_bytes(buf, &data_sz, 4);
|
||||
+}
|
||||
+
|
||||
+static void append_wav16_pcm(std::vector<char> & buf, const float * pcm, size_t n) {
|
||||
+ for (size_t i = 0; i < n; i++) {
|
||||
+ int16_t s = (int16_t) (std::max(-1.0f, std::min(1.0f, pcm[i])) * 32767.0f);
|
||||
+ put_bytes(buf, &s, 2);
|
||||
+ }
|
||||
+}
|
||||
+
|
||||
static bool write_wav16(std::vector<char> & buf, const std::vector<float> & pcm, int32_t rate) {
|
||||
// RIFF chunk sizes are 32-bit; refuse to emit a file with a truncated header
|
||||
if (pcm.size() > ((size_t) UINT32_MAX - 36) / 2) {
|
||||
return false;
|
||||
}
|
||||
- const uint32_t data_sz = (uint32_t) (pcm.size() * 2);
|
||||
- const uint32_t riff_sz = 36 + data_sz;
|
||||
- const uint32_t fmt_sz = 16, byte_rate = (uint32_t) rate * 2;
|
||||
- const uint16_t fmt = 1, ch = 1, align = 2, bits = 16;
|
||||
- const uint32_t rate32 = (uint32_t) rate;
|
||||
- auto put = [&](const void * p, size_t n) {
|
||||
- const char * c = (const char *) p;
|
||||
- buf.insert(buf.end(), c, c + n);
|
||||
- };
|
||||
- put("RIFF", 4); put(&riff_sz, 4); put("WAVE", 4);
|
||||
- put("fmt ", 4); put(&fmt_sz, 4);
|
||||
- put(&fmt, 2); put(&ch, 2); put(&rate32, 4);
|
||||
- put(&byte_rate, 4); put(&align, 2); put(&bits, 2);
|
||||
- put("data", 4); put(&data_sz, 4);
|
||||
- for (float v : pcm) {
|
||||
- int16_t s = (int16_t) (std::max(-1.0f, std::min(1.0f, v)) * 32767.0f);
|
||||
- put(&s, 2);
|
||||
- }
|
||||
+ write_wav16_header(buf, (uint32_t) (pcm.size() * 2), rate);
|
||||
+ append_wav16_pcm(buf, pcm.data(), pcm.size());
|
||||
return true;
|
||||
}
|
||||
|
||||
@@ -92,6 +101,8 @@ public:
|
||||
// set out_stop on end-of-speech, h_state_out must be null if no frame is generated
|
||||
virtual int32_t step_gen(llama_token sampled, const float * h_state_in, const float ** h_state_out, bool * out_stop) = 0;
|
||||
virtual int32_t get_output(int32_t * out_sample_rate, const char ** out_data, size_t * out_data_len, int64_t * out_n_samples) = 0;
|
||||
+ // forces any buffered codes through code2wav now, regardless of window_frames
|
||||
+ virtual int32_t flush() { return 0; }
|
||||
|
||||
protected:
|
||||
llama_context * lctx;
|
||||
@@ -121,6 +132,9 @@ public:
|
||||
prompt_batch.reset();
|
||||
n_prompt = 0;
|
||||
prompt_pos = 0;
|
||||
+ stream = false;
|
||||
+ pcm_sent = 0;
|
||||
+ wav_header_sent = false;
|
||||
}
|
||||
|
||||
int32_t set_input(const mtmd_helper_gen_audio_inp * inp) override {
|
||||
@@ -208,6 +222,7 @@ public:
|
||||
top_p = inp->top_p > 0 ? inp->top_p : def.top_p;
|
||||
seed = inp->seed;
|
||||
out_type = inp->out_type;
|
||||
+ stream = inp->stream;
|
||||
|
||||
// the prompt above holds the whole text stream up to tts_eos, so every generated
|
||||
// frame adds tts_pad on top of the codes embedding
|
||||
@@ -302,31 +317,60 @@ public:
|
||||
}
|
||||
|
||||
int32_t get_output(int32_t * out_sample_rate, const char ** out_data, size_t * out_data_len, int64_t * out_n_samples) override {
|
||||
- if (!flush_gen_wav()) {
|
||||
- return 1;
|
||||
+ *out_sample_rate = info.sample_rate;
|
||||
+
|
||||
+ if (!stream) {
|
||||
+ // one-shot call: force out whatever's left, regardless of window_frames
|
||||
+ if (!flush_gen_wav()) {
|
||||
+ return 1;
|
||||
+ }
|
||||
+ if (out_n_samples) {
|
||||
+ *out_n_samples = (int64_t) audio_pcm.size();
|
||||
+ }
|
||||
+ if (out_type == MTMD_HELPER_GEN_AUDIO_OUTTYPE_PCM) {
|
||||
+ *out_data = (const char *) audio_pcm.data();
|
||||
+ *out_data_len = audio_pcm.size() * sizeof(float);
|
||||
+ return 0;
|
||||
+ }
|
||||
+ out_buf.clear();
|
||||
+ if (!write_wav16(out_buf, audio_pcm, info.sample_rate)) {
|
||||
+ LOG_ERR("mtmd_helper_gen_audio: output too large for WAV\n");
|
||||
+ return 1;
|
||||
+ }
|
||||
+ *out_data = out_buf.data();
|
||||
+ *out_data_len = out_buf.size();
|
||||
+ return 0;
|
||||
}
|
||||
|
||||
- *out_sample_rate = info.sample_rate;
|
||||
+ // streaming: only return audio produced since the previous call
|
||||
+ const size_t n_new = audio_pcm.size() - pcm_sent;
|
||||
if (out_n_samples) {
|
||||
- *out_n_samples = (int64_t) audio_pcm.size();
|
||||
+ *out_n_samples = (int64_t) n_new;
|
||||
}
|
||||
|
||||
if (out_type == MTMD_HELPER_GEN_AUDIO_OUTTYPE_PCM) {
|
||||
- *out_data = (const char *) audio_pcm.data();
|
||||
- *out_data_len = audio_pcm.size() * sizeof(float);
|
||||
+ *out_data = (const char *) (audio_pcm.data() + pcm_sent);
|
||||
+ *out_data_len = n_new * sizeof(float);
|
||||
+ pcm_sent = audio_pcm.size();
|
||||
return 0;
|
||||
}
|
||||
|
||||
out_buf.clear();
|
||||
- if (!write_wav16(out_buf, audio_pcm, info.sample_rate)) {
|
||||
- LOG_ERR("mtmd_helper_gen_audio: output too large for WAV\n");
|
||||
- return 1;
|
||||
+ if (!wav_header_sent) {
|
||||
+ write_wav16_header(out_buf, UINT32_MAX, info.sample_rate);
|
||||
+ wav_header_sent = true;
|
||||
}
|
||||
+ append_wav16_pcm(out_buf, audio_pcm.data() + pcm_sent, n_new);
|
||||
+ pcm_sent = audio_pcm.size();
|
||||
*out_data = out_buf.data();
|
||||
*out_data_len = out_buf.size();
|
||||
return 0;
|
||||
}
|
||||
|
||||
+ int32_t flush() override {
|
||||
+ return flush_gen_wav() ? 0 : 1;
|
||||
+ }
|
||||
+
|
||||
private:
|
||||
bool ensure_cache() {
|
||||
if (specials_ok) {
|
||||
@@ -370,7 +414,7 @@ private:
|
||||
LOG_ERR("mtmd_helper_gen_audio: mmproj has no speaker/audio encoder\n");
|
||||
return false;
|
||||
}
|
||||
- const std::string marker = mtmd_default_marker();
|
||||
+ const std::string marker = mtmd_get_marker(mctx);
|
||||
mtmd_input_text text{ marker.c_str(), marker.size(), false, true };
|
||||
mtmd_input_chunks * chunks = mtmd_input_chunks_init();
|
||||
const mtmd_bitmap * bptr = bitmap;
|
||||
@@ -456,6 +500,9 @@ private:
|
||||
std::vector<float> h_state_buf;
|
||||
mtmd_helper_gen_audio_outtype out_type = MTMD_HELPER_GEN_AUDIO_OUTTYPE_WAV;
|
||||
std::vector<char> out_buf;
|
||||
+ bool stream = false;
|
||||
+ size_t pcm_sent = 0; // samples already returned by get_output()
|
||||
+ bool wav_header_sent = false;
|
||||
};
|
||||
|
||||
// settings that only live in the reference's per-pack yaml, not in the checkpoint
|
||||
@@ -1024,6 +1071,14 @@ void mtmd_helper_gen_audio_reset(mtmd_helper_gen_audio * ctx) {
|
||||
}
|
||||
}
|
||||
|
||||
+struct mtmd_helper_gen_audio_inp mtmd_helper_gen_audio_inp_default(void) {
|
||||
+ mtmd_helper_gen_audio_inp inp{};
|
||||
+ inp.top_k = 50;
|
||||
+ inp.top_p = 1.0f;
|
||||
+ inp.out_type = MTMD_HELPER_GEN_AUDIO_OUTTYPE_WAV;
|
||||
+ return inp;
|
||||
+}
|
||||
+
|
||||
int32_t mtmd_helper_gen_audio_set_input(mtmd_helper_gen_audio * ctx, const mtmd_helper_gen_audio_inp * inp) {
|
||||
if (!ctx->pipeline) {
|
||||
LOG_ERR("mtmd_helper_gen_audio: unsupported or missing gen-audio pipeline\n");
|
||||
@@ -1060,3 +1115,10 @@ int32_t mtmd_helper_gen_audio_get_output(mtmd_helper_gen_audio * ctx, int32_t *
|
||||
}
|
||||
return ctx->pipeline->get_output(out_sample_rate, out_data, out_data_len, out_n_samples);
|
||||
}
|
||||
+
|
||||
+int32_t mtmd_helper_gen_audio_flush(mtmd_helper_gen_audio * ctx) {
|
||||
+ if (!ctx->pipeline) {
|
||||
+ return 1;
|
||||
+ }
|
||||
+ return ctx->pipeline->flush();
|
||||
+}
|
||||
diff --git a/tools/mtmd/mtmd-helper.h b/tools/mtmd/mtmd-helper.h
|
||||
index 832f7171a..3eaa01aab 100644
|
||||
--- a/tools/mtmd/mtmd-helper.h
|
||||
+++ b/tools/mtmd/mtmd-helper.h
|
||||
@@ -175,6 +175,7 @@ enum mtmd_helper_gen_audio_outtype {
|
||||
MTMD_HELPER_GEN_AUDIO_OUTTYPE_WAV, // WAV PCM 16-bit LE, mono
|
||||
};
|
||||
struct mtmd_helper_gen_audio_inp {
|
||||
+ bool stream; // if true, output() must be called after each step_gen()
|
||||
llama_seq_id seq_id;
|
||||
|
||||
const char * prompt;
|
||||
@@ -190,6 +191,8 @@ struct mtmd_helper_gen_audio_inp {
|
||||
enum mtmd_helper_gen_audio_outtype out_type;
|
||||
};
|
||||
|
||||
+MTMD_API struct mtmd_helper_gen_audio_inp mtmd_helper_gen_audio_inp_default(void);
|
||||
+
|
||||
MTMD_API mtmd_helper_gen_audio * mtmd_helper_gen_audio_init(
|
||||
struct llama_context * lctx,
|
||||
struct mtmd_context * mctx);
|
||||
@@ -221,6 +224,8 @@ MTMD_API int32_t mtmd_helper_gen_audio_step_gen(
|
||||
|
||||
// out_data valid until next get_output() or reset() call
|
||||
// out_n_samples (optional, can be NULL) receives the number of generated PCM samples
|
||||
+// if inp->stream is true: returns only audio produced since the previous call, and
|
||||
+// *out_data_len == 0 whenever a full window_frames batch hasn't accumulated yet
|
||||
MTMD_API int32_t mtmd_helper_gen_audio_get_output(
|
||||
mtmd_helper_gen_audio * ctx,
|
||||
int32_t * out_sample_rate,
|
||||
@@ -228,6 +233,10 @@ MTMD_API int32_t mtmd_helper_gen_audio_get_output(
|
||||
size_t * out_data_len,
|
||||
int64_t * out_n_samples);
|
||||
|
||||
+// forces any buffered codes through code2wav now, regardless of window_frames;
|
||||
+// call once when generation has ended, before the last get_output() in stream mode
|
||||
+MTMD_API int32_t mtmd_helper_gen_audio_flush(mtmd_helper_gen_audio * ctx);
|
||||
+
|
||||
#ifdef __cplusplus
|
||||
} // extern "C"
|
||||
#endif
|
||||
@@ -254,8 +263,41 @@ struct mtmd_helper_gen_audio_deleter {
|
||||
};
|
||||
using gen_audio_ptr = std::unique_ptr<mtmd_helper_gen_audio, mtmd_helper_gen_audio_deleter>;
|
||||
struct gen_audio {
|
||||
+
|
||||
+ // sub-struct, RAII wrapper for mtmd_helper_gen_audio_inp
|
||||
+ struct inp {
|
||||
+ mtmd_helper_gen_audio_inp data = mtmd_helper_gen_audio_inp_default();
|
||||
+ std::string prompt_str;
|
||||
+ std::string lang_str;
|
||||
+ mtmd::bitmap_ptr speaker_ref_ptr;
|
||||
+
|
||||
+ inp() = default;
|
||||
+ inp(inp &&) = default;
|
||||
+ inp & operator=(inp &&) = default;
|
||||
+ inp(const inp &) = delete;
|
||||
+ inp & operator=(const inp &) = delete;
|
||||
+
|
||||
+ void set_prompt (std::string p) { prompt_str = std::move(p); }
|
||||
+ void set_lang (std::string l) { lang_str = std::move(l); }
|
||||
+ void set_speaker_ref(mtmd::bitmap_ptr bmp) { speaker_ref_ptr = std::move(bmp); }
|
||||
+
|
||||
+ // pointers are only valid as long as *this is alive
|
||||
+ const mtmd_helper_gen_audio_inp * get() {
|
||||
+ data.prompt = prompt_str.c_str();
|
||||
+ data.prompt_len = prompt_str.size();
|
||||
+ data.lang = lang_str.empty() ? nullptr : lang_str.c_str();
|
||||
+ data.speaker_ref = speaker_ref_ptr.get();
|
||||
+ return &data;
|
||||
+ }
|
||||
+ };
|
||||
+
|
||||
gen_audio_ptr ctx;
|
||||
- gen_audio(struct llama_context * lctx, struct mtmd_context * mctx) : ctx(mtmd_helper_gen_audio_init(lctx, mctx)) {}
|
||||
+ void init(struct llama_context * lctx, struct mtmd_context * mctx) {
|
||||
+ ctx.reset(mtmd_helper_gen_audio_init(lctx, mctx));
|
||||
+ }
|
||||
+ bool valid() const {
|
||||
+ return ctx.get() != nullptr;
|
||||
+ }
|
||||
void reset() {
|
||||
mtmd_helper_gen_audio_reset(ctx.get());
|
||||
}
|
||||
@@ -271,6 +313,9 @@ struct gen_audio {
|
||||
int32_t get_output(int32_t * out_sample_rate, const char ** out_data, size_t * out_data_len, int64_t * out_n_samples = nullptr) {
|
||||
return mtmd_helper_gen_audio_get_output(ctx.get(), out_sample_rate, out_data, out_data_len, out_n_samples);
|
||||
}
|
||||
+ int32_t flush() {
|
||||
+ return mtmd_helper_gen_audio_flush(ctx.get());
|
||||
+ }
|
||||
};
|
||||
|
||||
} // namespace mtmd_helper
|
||||
diff --git a/tools/server/server-context.cpp b/tools/server/server-context.cpp
|
||||
index 9069463fe..b7fa1e534 100644
|
||||
--- a/tools/server/server-context.cpp
|
||||
+++ b/tools/server/server-context.cpp
|
||||
@@ -16,6 +16,7 @@
|
||||
#include "speculative.h"
|
||||
#include "mtmd.h"
|
||||
#include "mtmd-helper.h"
|
||||
+#include "base64.hpp"
|
||||
|
||||
#include <algorithm>
|
||||
#include <cstddef>
|
||||
@@ -41,8 +42,9 @@ constexpr int HTTP_POLLING_SECONDS = 1;
|
||||
|
||||
static common_speculative_output_limits server_output_limits(const common_params & params) {
|
||||
if (params.embedding ||
|
||||
- (params.pooling_type != LLAMA_POOLING_TYPE_UNSPECIFIED && params.pooling_type != LLAMA_POOLING_TYPE_NONE)) {
|
||||
- return { params.n_batch, 1 };
|
||||
+ (params.pooling_type != LLAMA_POOLING_TYPE_UNSPECIFIED && params.pooling_type != LLAMA_POOLING_TYPE_NONE) ||
|
||||
+ !params.mmproj.path.empty()) { // gen-audio (TTS) capability isn't known until the mmproj loads, size generously
|
||||
+ return { params.n_batch, params.n_batch };
|
||||
}
|
||||
|
||||
auto result = common_speculative_get_output_limits(
|
||||
@@ -212,6 +214,30 @@ struct server_slot {
|
||||
mtmd_context * mctx = nullptr;
|
||||
mtmd::batch_ptr mbatch = nullptr;
|
||||
|
||||
+ struct tts_ctx {
|
||||
+ mtmd_helper::gen_audio ctx;
|
||||
+ const float * h_state;
|
||||
+ llama_token sampled;
|
||||
+ int32_t n_decoded;
|
||||
+ bool is_supported() const {
|
||||
+ return ctx.valid();
|
||||
+ }
|
||||
+ void reset() {
|
||||
+ // mtmd_helper_gen_audio_reset() dereferences its argument before it
|
||||
+ // null-checks the pipeline, and the pipeline is only allocated for
|
||||
+ // models that actually carry a gen-audio mmproj. server_slot::reset()
|
||||
+ // runs for every slot of every model, so without this guard any
|
||||
+ // non-TTS model segfaults during slot initialization.
|
||||
+ if (is_supported()) {
|
||||
+ ctx.reset();
|
||||
+ }
|
||||
+ h_state = nullptr;
|
||||
+ sampled = LLAMA_TOKEN_NULL;
|
||||
+ n_decoded = 0;
|
||||
+ }
|
||||
+ };
|
||||
+ tts_ctx tts;
|
||||
+
|
||||
// speculative decoding
|
||||
common_speculative * spec;
|
||||
|
||||
@@ -391,6 +417,8 @@ struct server_slot {
|
||||
|
||||
// clear multimodal state
|
||||
mbatch.reset();
|
||||
+
|
||||
+ tts.reset();
|
||||
}
|
||||
|
||||
void init_sampler() const {
|
||||
@@ -829,6 +857,14 @@ public:
|
||||
mtmd_context * mctx = nullptr;
|
||||
const llama_vocab * vocab = nullptr;
|
||||
|
||||
+ bool has_cap_tts() const {
|
||||
+ return mctx != nullptr && mtmd_gen_audio_get_info(mctx).type != MTMD_GEN_AUDIO_TYPE_NONE;
|
||||
+ }
|
||||
+
|
||||
+ bool has_cap_chat() const {
|
||||
+ return mctx == nullptr || mtmd_helper_model_can_chat(ctx_tgt, mctx);
|
||||
+ }
|
||||
+
|
||||
server_queue queue_tasks;
|
||||
server_response queue_results;
|
||||
|
||||
@@ -1288,6 +1324,10 @@ private:
|
||||
slot.mctx = mctx;
|
||||
slot.prompt.tokens.has_mtmd = mctx != nullptr;
|
||||
|
||||
+ if (has_cap_tts()) {
|
||||
+ slot.tts.ctx.init(ctx_tgt, mctx);
|
||||
+ }
|
||||
+
|
||||
SLT_TRC(slot, "new slot, n_ctx = %d\n", slot.n_ctx);
|
||||
|
||||
slot.callback_on_release = [this](int id_slot) {
|
||||
@@ -1748,6 +1788,28 @@ private:
|
||||
|
||||
SLT_DBG(slot, "launching slot : %s\n", safe_json_to_str(slot.to_json()).c_str());
|
||||
|
||||
+ if (task.type == SERVER_TASK_TYPE_TTS) {
|
||||
+ GGML_ASSERT(has_cap_tts()); // should already checked in route handler
|
||||
+ if (!slot.tts.is_supported()) {
|
||||
+ slot.tts.ctx.init(ctx_tgt, slot.mctx);
|
||||
+ }
|
||||
+
|
||||
+ // TTS slots never enter the shared batch: pre_decode() returns early for
|
||||
+ // them and process_tts_slots() drives them instead, so they skip the
|
||||
+ // prompt-cache bookkeeping that clears this sequence between requests.
|
||||
+ // The gen-audio pipeline always decodes from position 0, and its own
|
||||
+ // reset() only clears host-side buffers, so without this the second and
|
||||
+ // later tasks on a slot decode over the previous request's tokens and
|
||||
+ // step_prompt() fails immediately.
|
||||
+ slot.prompt_clear();
|
||||
+
|
||||
+ task.tts_inp.data.seq_id = slot.id;
|
||||
+ if (slot.tts.ctx.set_input(task.tts_inp.get()) != 0) {
|
||||
+ send_error(task, "failed to process TTS prompt", ERROR_TYPE_SERVER);
|
||||
+ return false;
|
||||
+ }
|
||||
+ }
|
||||
+
|
||||
// initialize samplers
|
||||
if (task.need_sampling()) {
|
||||
try {
|
||||
@@ -1765,6 +1827,9 @@ private:
|
||||
// TODO: getting pre sampling logits is not yet supported with backend sampling
|
||||
use_backend_sampling &= !need_pre_sample_logits;
|
||||
|
||||
+ // TODO: check verify if this actually works with TTS
|
||||
+ use_backend_sampling &= task.type != SERVER_TASK_TYPE_TTS;
|
||||
+
|
||||
// TODO: tmp until backend sampling is fully implemented
|
||||
if (use_backend_sampling) {
|
||||
llama_set_sampler(ctx_tgt, slot.id, common_sampler_get(slot.smpl.get()));
|
||||
@@ -1783,9 +1848,13 @@ private:
|
||||
|
||||
slot.task = std::make_unique<const server_task>(std::move(task));
|
||||
|
||||
- slot.state = slot.task->is_child()
|
||||
- ? SLOT_STATE_WAIT_OTHER // wait for the parent to process prompt
|
||||
- : SLOT_STATE_STARTED;
|
||||
+ if (slot.task->type == SERVER_TASK_TYPE_TTS) {
|
||||
+ slot.state = SLOT_STATE_PROCESSING_PROMPT;
|
||||
+ } else {
|
||||
+ slot.state = slot.task->is_child()
|
||||
+ ? SLOT_STATE_WAIT_OTHER // wait for the parent to process prompt
|
||||
+ : SLOT_STATE_STARTED;
|
||||
+ }
|
||||
|
||||
// reset server kill-switch counter
|
||||
n_empty_consecutive = 0;
|
||||
@@ -2050,6 +2119,18 @@ private:
|
||||
queue_results.send(std::move(res));
|
||||
}
|
||||
|
||||
+ void send_tts_result(server_slot & slot, int32_t sample_rate, const char * data, size_t data_len, bool final) {
|
||||
+ auto res = std::make_unique<server_task_result_tts>();
|
||||
+
|
||||
+ res->id = slot.task->id;
|
||||
+ res->index = slot.task->index;
|
||||
+ res->sample_rate = sample_rate;
|
||||
+ res->audio.assign(data, data_len);
|
||||
+ res->final = final;
|
||||
+
|
||||
+ queue_results.send(std::move(res));
|
||||
+ }
|
||||
+
|
||||
void send_final_response(server_slot & slot) {
|
||||
auto res = std::make_unique<server_task_result_cmpl_final>();
|
||||
|
||||
@@ -2556,6 +2637,7 @@ private:
|
||||
case SERVER_TASK_TYPE_EMBEDDING:
|
||||
case SERVER_TASK_TYPE_RERANK:
|
||||
case SERVER_TASK_TYPE_SCORE:
|
||||
+ case SERVER_TASK_TYPE_TTS:
|
||||
{
|
||||
// special case: if input is provided via CLI, tokenize it first
|
||||
// otherwise, no need to tokenize as it's already done inside the HTTP thread
|
||||
@@ -3007,1 +3089,9 @@ private:
|
||||
+ // note: TTS slots bypass the shared batch entirely
|
||||
+ try {
|
||||
+ process_tts_slots();
|
||||
+ } catch (const std::exception & e) {
|
||||
+ SRV_ERR("process_tts_slots() failed: %s\n", e.what());
|
||||
+ abort_all_slots("process_tts_slots() failed: " + std::string(e.what()));
|
||||
+ }
|
||||
+
|
||||
GGML_ASSERT(batch.slot_batched || batch.size() == 0);
|
||||
@@ -3074,10 +3164,77 @@ private:
|
||||
}
|
||||
}
|
||||
|
||||
+ void process_tts_slots() {
|
||||
+ iterate(slots, [&](server_slot & slot) {
|
||||
+ if (!slot.is_processing() || slot.task->type != SERVER_TASK_TYPE_TTS) {
|
||||
+ return;
|
||||
+ }
|
||||
+
|
||||
+ llama_set_embeddings(ctx_tgt, true);
|
||||
+
|
||||
+ if (slot.state == SLOT_STATE_PROCESSING_PROMPT) {
|
||||
+ const int32_t ret = slot.tts.ctx.step_prompt(llama_n_batch(ctx_tgt));
|
||||
+ if (ret < 0) {
|
||||
+ send_error(slot, "TTS prompt processing failed", ERROR_TYPE_SERVER);
|
||||
+ slot.release();
|
||||
+ } else if (ret == 0) {
|
||||
+ slot.tts.sampled = common_sampler_sample(slot.smpl.get(), ctx_tgt, -1);
|
||||
+ common_sampler_accept(slot.smpl.get(), slot.tts.sampled, true);
|
||||
+ slot.tts.h_state = llama_get_embeddings_ith(ctx_tgt, -1);
|
||||
+ slot.state = SLOT_STATE_GENERATING;
|
||||
+ }
|
||||
+ return;
|
||||
+ }
|
||||
+
|
||||
+ const int32_t n_predict = slot.task->params.n_predict > 0 ? slot.task->params.n_predict : 512;
|
||||
+ if (slot.tts.n_decoded >= n_predict || llama_vocab_is_eog(vocab, slot.tts.sampled)) {
|
||||
+ int32_t sample_rate = 0;
|
||||
+ const char * data = nullptr;
|
||||
+ size_t data_len = 0;
|
||||
+ // generation truly ends here: force out any sub-window remainder still buffered
|
||||
+ if (slot.tts.ctx.flush() != 0 || slot.tts.ctx.get_output(&sample_rate, &data, &data_len) != 0) {
|
||||
+ send_error(slot, "failed to finalize TTS output", ERROR_TYPE_SERVER);
|
||||
+ } else {
|
||||
+ send_tts_result(slot, sample_rate, data, data_len, true);
|
||||
+ }
|
||||
+ slot.release();
|
||||
+ return;
|
||||
+ }
|
||||
+
|
||||
+ const float * h_state_next = nullptr;
|
||||
+ if (slot.tts.ctx.step_gen(slot.tts.sampled, slot.tts.h_state, &h_state_next) != 0) {
|
||||
+ send_error(slot, "TTS generation failed", ERROR_TYPE_SERVER);
|
||||
+ slot.release();
|
||||
+ return;
|
||||
+ }
|
||||
+ slot.tts.h_state = h_state_next;
|
||||
+ slot.tts.n_decoded++;
|
||||
+
|
||||
+ slot.tts.sampled = common_sampler_sample(slot.smpl.get(), ctx_tgt, -1);
|
||||
+ common_sampler_accept(slot.smpl.get(), slot.tts.sampled, true);
|
||||
+
|
||||
+ if (slot.task->params.stream) {
|
||||
+ int32_t sample_rate = 0;
|
||||
+ const char * data = nullptr;
|
||||
+ size_t data_len = 0;
|
||||
+ if (slot.tts.ctx.get_output(&sample_rate, &data, &data_len) != 0) {
|
||||
+ send_error(slot, "TTS streaming output failed", ERROR_TYPE_SERVER);
|
||||
+ slot.release();
|
||||
+ } else if (data_len > 0) {
|
||||
+ send_tts_result(slot, sample_rate, data, data_len, false);
|
||||
+ }
|
||||
+ }
|
||||
+ });
|
||||
+ }
|
||||
+
|
||||
void pre_decode() {
|
||||
// apply context-shift if needed
|
||||
// TODO: simplify and improve
|
||||
iterate(slots, [&](server_slot & slot) {
|
||||
+ if (slot.task && slot.task->type == SERVER_TASK_TYPE_TTS) {
|
||||
+ // TTS slots drive their own decode loop in process_tts_slots(), never enter the shared batch
|
||||
+ return;
|
||||
+ }
|
||||
if (slot.state == SLOT_STATE_GENERATING && slot.prompt.n_tokens() + 1 >= slot.n_ctx) {
|
||||
if (!params_base.ctx_shift) {
|
||||
// this check is redundant (for good)
|
||||
@@ -3150,7 +3307,7 @@ private:
|
||||
|
||||
// determine which slots are generating and drafting
|
||||
iterate(slots, [&](server_slot & slot) {
|
||||
- if (slot.state != SLOT_STATE_GENERATING) {
|
||||
+ if (slot.state != SLOT_STATE_GENERATING || slot.task->type == SERVER_TASK_TYPE_TTS) {
|
||||
return;
|
||||
}
|
||||
|
||||
@@ -3284,7 +3441,7 @@ private:
|
||||
return; // batch is full, skip remaining slots
|
||||
}
|
||||
|
||||
- if (!slot.is_processing()) {
|
||||
+ if (!slot.is_processing() || slot.task->type == SERVER_TASK_TYPE_TTS) {
|
||||
return;
|
||||
}
|
||||
|
||||
@@ -4433,6 +4590,8 @@ server_context_meta server_context::get_meta() const {
|
||||
/* has_inp_image */ impl->chat_params.allow_image,
|
||||
/* has_inp_audio */ impl->chat_params.allow_audio,
|
||||
/* has_inp_video */ impl->chat_params.allow_video,
|
||||
+ /* has_cap_chat */ impl->has_cap_chat(),
|
||||
+ /* has_cap_tts */ impl->has_cap_tts(),
|
||||
/* json_ui_settings */ impl->json_ui_settings,
|
||||
/* slot_n_ctx */ impl->get_slot_n_ctx(),
|
||||
/* pooling_type */ llama_pooling_type(impl->ctx_tgt),
|
||||
@@ -4512,6 +4671,11 @@ std::unique_ptr<server_res_generator> server_routes::handle_completions_impl(
|
||||
|
||||
res->set_req(&req); // will also set spipe if needed
|
||||
|
||||
+ if (!ctx_server.has_cap_chat()) {
|
||||
+ res->error(format_error_response("this server does not support chat/completions", ERROR_TYPE_NOT_SUPPORTED));
|
||||
+ return res;
|
||||
+ }
|
||||
+
|
||||
int32_t sse_ping_interval = params.sse_ping_interval;
|
||||
|
||||
try {
|
||||
@@ -5399,6 +5563,150 @@ void server_routes::init_routes() {
|
||||
return res;
|
||||
};
|
||||
|
||||
+ this->post_tts = [this](const server_http_req & req) {
|
||||
+ auto res = create_response();
|
||||
+ res->set_req(&req); // will also set spipe if needed
|
||||
+
|
||||
+ if (!ctx_server.has_cap_tts()) {
|
||||
+ res->error(format_error_response("this server does not support audio generation", ERROR_TYPE_NOT_SUPPORTED));
|
||||
+ return res;
|
||||
+ }
|
||||
+
|
||||
+ const json body = json::parse(req.body);
|
||||
+
|
||||
+ std::string prompt = json_value(body, "input", json_value(body, "prompt", std::string()));
|
||||
+ if (prompt.empty()) {
|
||||
+ res->error(format_error_response("\"input\" must be a non-empty string", ERROR_TYPE_INVALID_REQUEST));
|
||||
+ return res;
|
||||
+ }
|
||||
+
|
||||
+ const std::string response_format = json_value(body, "response_format", std::string("wav"));
|
||||
+ const bool stream = json_value(body, "stream", false);
|
||||
+
|
||||
+ server_task task(SERVER_TASK_TYPE_TTS);
|
||||
+ task.tts_inp.set_prompt(prompt);
|
||||
+ task.tts_inp.set_lang(json_value(body, "lang", std::string()));
|
||||
+ task.tts_inp.data.top_k = json_value(body, "top_k", 0);
|
||||
+ task.tts_inp.data.top_p = json_value(body, "top_p", 0.0f);
|
||||
+ task.tts_inp.data.stream = stream;
|
||||
+ task.tts_inp.data.out_type = response_format == "pcm"
|
||||
+ ? MTMD_HELPER_GEN_AUDIO_OUTTYPE_PCM
|
||||
+ : MTMD_HELPER_GEN_AUDIO_OUTTYPE_WAV;
|
||||
+ task.params.stream = stream;
|
||||
+ task.params.n_predict = json_value(body, "n_predict", -1);
|
||||
+ task.params.sampling = params.sampling; // baseline defaults, then apply overrides below
|
||||
+ task.params.sampling.penalty_repeat = json_value(body, "repeat_penalty", 1.05f);
|
||||
+ task.params.sampling.penalty_last_n = -1;
|
||||
+ if (task.tts_inp.data.top_k > 0) {
|
||||
+ task.params.sampling.top_k = task.tts_inp.data.top_k;
|
||||
+ }
|
||||
+ if (task.tts_inp.data.top_p > 0) {
|
||||
+ task.params.sampling.top_p = task.tts_inp.data.top_p;
|
||||
+ }
|
||||
+
|
||||
+ // speaker reference: either an uploaded form file ("speaker_ref") or a base64 JSON field ("speaker_ref_b64")
|
||||
+ const unsigned char * speaker_ref_data = nullptr;
|
||||
+ size_t speaker_ref_len = 0;
|
||||
+ std::string speaker_ref_b64_decoded;
|
||||
+
|
||||
+ auto speaker_ref_file = req.files.find("speaker_ref");
|
||||
+ if (speaker_ref_file != req.files.end()) {
|
||||
+ speaker_ref_data = speaker_ref_file->second.data.data();
|
||||
+ speaker_ref_len = speaker_ref_file->second.data.size();
|
||||
+ } else {
|
||||
+ std::string speaker_ref_b64 = json_value(body, "speaker_ref_b64", std::string());
|
||||
+ if (!speaker_ref_b64.empty()) {
|
||||
+ speaker_ref_b64_decoded = base64::decode(speaker_ref_b64);
|
||||
+ speaker_ref_data = (const unsigned char *) speaker_ref_b64_decoded.data();
|
||||
+ speaker_ref_len = speaker_ref_b64_decoded.size();
|
||||
+ }
|
||||
+ }
|
||||
+
|
||||
+ if (speaker_ref_len > 0) {
|
||||
+ auto wrapper = mtmd_helper_bitmap_init_from_buf(ctx_server.mctx, speaker_ref_data, speaker_ref_len, false);
|
||||
+ if (!wrapper.bitmap) {
|
||||
+ res->error(format_error_response("failed to decode \"speaker_ref\"", ERROR_TYPE_INVALID_REQUEST));
|
||||
+ return res;
|
||||
+ }
|
||||
+ task.tts_inp.set_speaker_ref(mtmd::bitmap_ptr(wrapper.bitmap));
|
||||
+ } else {
|
||||
+ // SRV_WRN expands __VA_ARGS__ without the GNU comma-elision extension,
|
||||
+ // so a bare format string leaves a trailing comma and will not compile
|
||||
+ SRV_WRN("%s", "no speaker reference provided, the model may behave randomly\n");
|
||||
+ }
|
||||
+
|
||||
+ auto & rd = res->rd;
|
||||
+ task.id = rd.get_new_id();
|
||||
+ rd.post_task(std::move(task));
|
||||
+
|
||||
+ const std::string content_type = response_format == "pcm" ? "audio/L16" : "audio/wav";
|
||||
+
|
||||
+ if (!stream) {
|
||||
+ auto result = rd.next(req.should_stop);
|
||||
+ if (!result) {
|
||||
+ GGML_ASSERT(req.should_stop());
|
||||
+ return res; // connection is closed
|
||||
+ }
|
||||
+ if (result->is_error()) {
|
||||
+ res->error(result->to_json());
|
||||
+ return res;
|
||||
+ }
|
||||
+ auto * tts_res = dynamic_cast<server_task_result_tts *>(result.get());
|
||||
+ GGML_ASSERT(tts_res != nullptr);
|
||||
+ res->status = 200;
|
||||
+ res->content_type = content_type;
|
||||
+ res->data = std::move(tts_res->audio);
|
||||
+ return res;
|
||||
+ } else {
|
||||
+ auto first_result = rd.next(req.should_stop);
|
||||
+ if (!first_result) {
|
||||
+ GGML_ASSERT(req.should_stop());
|
||||
+ return res; // connection is closed
|
||||
+ }
|
||||
+ if (first_result->is_error()) {
|
||||
+ res->error(first_result->to_json());
|
||||
+ return res;
|
||||
+ }
|
||||
+ auto * first_tts_res = dynamic_cast<server_task_result_tts *>(first_result.get());
|
||||
+ GGML_ASSERT(first_tts_res != nullptr);
|
||||
+
|
||||
+ res->status = 200;
|
||||
+ res->content_type = content_type;
|
||||
+ res->data = std::move(first_tts_res->audio);
|
||||
+ bool is_done = first_tts_res->final;
|
||||
+
|
||||
+ res->set_next([res_this = res.get(), is_done](std::string & output) mutable -> bool {
|
||||
+ if (is_done) {
|
||||
+ return false;
|
||||
+ }
|
||||
+ if (res_this->should_stop()) {
|
||||
+ return false;
|
||||
+ }
|
||||
+ if (!res_this->data.empty()) {
|
||||
+ output = std::move(res_this->data);
|
||||
+ res_this->data.clear();
|
||||
+ return true;
|
||||
+ }
|
||||
+
|
||||
+ server_response_reader & rd = res_this->rd;
|
||||
+ if (!rd.has_next()) {
|
||||
+ return false;
|
||||
+ }
|
||||
+ auto result = rd.next([&res_this]() { return res_this->should_stop(); });
|
||||
+ if (!result || result->is_error()) {
|
||||
+ return false;
|
||||
+ }
|
||||
+ auto * tts_res = dynamic_cast<server_task_result_tts *>(result.get());
|
||||
+ GGML_ASSERT(tts_res != nullptr);
|
||||
+ output = std::move(tts_res->audio);
|
||||
+ is_done = tts_res->final;
|
||||
+ return true;
|
||||
+ });
|
||||
+ }
|
||||
+
|
||||
+ return res;
|
||||
+ };
|
||||
+
|
||||
this->get_lora_adapters = [this](const server_http_req & req) {
|
||||
auto res = create_response();
|
||||
|
||||
diff --git a/tools/server/server-context.h b/tools/server/server-context.h
|
||||
index f9ab1132b..610512678 100644
|
||||
--- a/tools/server/server-context.h
|
||||
+++ b/tools/server/server-context.h
|
||||
@@ -22,6 +22,8 @@ struct server_context_meta {
|
||||
bool has_inp_image;
|
||||
bool has_inp_audio;
|
||||
bool has_inp_video;
|
||||
+ bool has_cap_chat;
|
||||
+ bool has_cap_tts;
|
||||
json json_ui_settings;
|
||||
int slot_n_ctx;
|
||||
enum llama_pooling_type pooling_type;
|
||||
@@ -151,6 +153,7 @@ struct server_routes {
|
||||
server_http_context::handler_t post_embeddings;
|
||||
server_http_context::handler_t post_embeddings_oai;
|
||||
server_http_context::handler_t post_rerank;
|
||||
+ server_http_context::handler_t post_tts;
|
||||
server_http_context::handler_t get_lora_adapters;
|
||||
server_http_context::handler_t post_lora_adapters;
|
||||
|
||||
diff --git a/tools/server/server-task.cpp b/tools/server/server-task.cpp
|
||||
index 1ee677553..939630b8b 100644
|
||||
--- a/tools/server/server-task.cpp
|
||||
+++ b/tools/server/server-task.cpp
|
||||
@@ -1497,6 +1497,17 @@ json server_task_result_rerank::to_json() {
|
||||
};
|
||||
}
|
||||
|
||||
+//
|
||||
+// server_task_result_tts
|
||||
+//
|
||||
+json server_task_result_tts::to_json() {
|
||||
+ return json {
|
||||
+ {"sample_rate", sample_rate},
|
||||
+ {"n_bytes", audio.size()},
|
||||
+ {"final", final},
|
||||
+ };
|
||||
+}
|
||||
+
|
||||
//
|
||||
// server_task_result_error
|
||||
//
|
||||
diff --git a/tools/server/server-task.h b/tools/server/server-task.h
|
||||
index 5bedf1987..e6ca67a65 100644
|
||||
--- a/tools/server/server-task.h
|
||||
+++ b/tools/server/server-task.h
|
||||
@@ -10,6 +10,7 @@
|
||||
|
||||
// TODO: prevent including the whole server-common.h as we only use server_tokens
|
||||
#include "server-common.h"
|
||||
+#include "mtmd-helper.h"
|
||||
|
||||
using json = nlohmann::ordered_json;
|
||||
|
||||
@@ -42,6 +43,7 @@ enum server_task_type {
|
||||
SERVER_TASK_TYPE_SLOT_ERASE,
|
||||
SERVER_TASK_TYPE_GET_LORA,
|
||||
SERVER_TASK_TYPE_SET_LORA,
|
||||
+ SERVER_TASK_TYPE_TTS,
|
||||
};
|
||||
|
||||
// TODO: change this to more generic "response_format" to replace the "format_response_*" in server-common
|
||||
@@ -202,6 +204,9 @@ struct server_task {
|
||||
// used by SERVER_TASK_TYPE_SET_LORA
|
||||
std::map<int, float> set_lora; // mapping adapter ID -> scale
|
||||
|
||||
+ // used by SERVER_TASK_TYPE_TTS
|
||||
+ mtmd_helper::gen_audio::inp tts_inp;
|
||||
+
|
||||
server_task() = default;
|
||||
|
||||
server_task(server_task_type type) : type(type) {}
|
||||
@@ -235,6 +240,7 @@ struct server_task {
|
||||
switch (type) {
|
||||
case SERVER_TASK_TYPE_COMPLETION:
|
||||
case SERVER_TASK_TYPE_INFILL:
|
||||
+ case SERVER_TASK_TYPE_TTS:
|
||||
return true;
|
||||
default:
|
||||
return false;
|
||||
@@ -494,5 +500,15 @@ struct server_task_result_embd : server_task_result {
|
||||
json to_json_oaicompat();
|
||||
};
|
||||
|
||||
+struct server_task_result_tts : server_task_result {
|
||||
+ std::string audio; // raw bytes for this chunk (WAV or PCM, per request's out_type)
|
||||
+ int32_t sample_rate = 0;
|
||||
+ bool final = false; // true for the last chunk of a request
|
||||
+
|
||||
+ virtual bool is_stop() override { return final; }
|
||||
+
|
||||
+ virtual json to_json() override;
|
||||
+};
|
||||
+
|
||||
struct server_task_result_rerank : server_task_result {
|
||||
float score = -1e6;
|
||||
@@ -28,6 +28,13 @@ cp -r message_content_test.cpp llama.cpp/tools/grpc-server/
|
||||
# Generic passthrough parser staging and its standalone regression test.
|
||||
cp -r passthrough_options.h llama.cpp/tools/grpc-server/
|
||||
cp -r passthrough_options_test.cpp llama.cpp/tools/grpc-server/
|
||||
# TTS request validation (included by grpc-server.cpp) and its standalone
|
||||
# regression test.
|
||||
cp -r tts_request_options.h llama.cpp/tools/grpc-server/
|
||||
cp -r tts_request_options_test.cpp llama.cpp/tools/grpc-server/
|
||||
# Thread-count default normalization and its standalone regression test.
|
||||
cp -r thread_params.h llama.cpp/tools/grpc-server/
|
||||
cp -r thread_params_test.cpp llama.cpp/tools/grpc-server/
|
||||
# Parent-death watcher (included by grpc-server.cpp) and its standalone unit
|
||||
# test (run via backend/cpp/run-unit-tests.sh; also buildable under ctest).
|
||||
cp -r parent_watch.h llama.cpp/tools/grpc-server/
|
||||
@@ -49,10 +56,16 @@ else
|
||||
echo "==> llama.cpp predates the load-mode enum, using the legacy mmap/mlock/direct-io booleans"
|
||||
LEGACY_LOAD_MODE=1
|
||||
fi
|
||||
if grep -q "server_metrics metrics;" llama.cpp/tools/server/server-task.h; then
|
||||
HAS_SERVER_METRICS=1
|
||||
else
|
||||
HAS_SERVER_METRICS=0
|
||||
fi
|
||||
cat > llama.cpp/tools/grpc-server/llama_compat.h <<EOF
|
||||
// Generated by backend/cpp/llama-cpp/prepare.sh. Do not edit.
|
||||
#pragma once
|
||||
#define LOCALAI_LEGACY_LOAD_MODE ${LEGACY_LOAD_MODE}
|
||||
#define LOCALAI_HAS_SERVER_METRICS ${HAS_SERVER_METRICS}
|
||||
EOF
|
||||
|
||||
set +e
|
||||
|
||||
@@ -0,0 +1,11 @@
|
||||
#pragma once
|
||||
|
||||
#include <cstdint>
|
||||
|
||||
namespace llama_grpc {
|
||||
|
||||
inline int32_t resolve_batch_threads(int32_t batch_threads, int32_t inference_threads) {
|
||||
return batch_threads < 0 ? inference_threads : batch_threads;
|
||||
}
|
||||
|
||||
} // namespace llama_grpc
|
||||
@@ -0,0 +1,15 @@
|
||||
#include "thread_params.h"
|
||||
|
||||
#include <cstdio>
|
||||
|
||||
int main() {
|
||||
if (llama_grpc::resolve_batch_threads(-1, 4) != 4) {
|
||||
std::fprintf(stderr, "default batch threads did not inherit inference threads\n");
|
||||
return 1;
|
||||
}
|
||||
if (llama_grpc::resolve_batch_threads(2, 4) != 2) {
|
||||
std::fprintf(stderr, "explicit batch threads were overwritten\n");
|
||||
return 1;
|
||||
}
|
||||
return 0;
|
||||
}
|
||||
@@ -0,0 +1,149 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <cstdint>
|
||||
#include <exception>
|
||||
#include <map>
|
||||
#include <string>
|
||||
|
||||
namespace llama_grpc {
|
||||
|
||||
// Validated, parsed form of a backend::TTSRequest, kept free of llama.cpp,
|
||||
// mtmd and gRPC headers so backend/cpp/run-unit-tests.sh can compile it as a
|
||||
// standalone translation unit. grpc-server.cpp turns this into a
|
||||
// mtmd_helper::gen_audio::inp.
|
||||
struct tts_request_options {
|
||||
bool ok = false;
|
||||
std::string error;
|
||||
|
||||
std::string text;
|
||||
std::string voice_path;
|
||||
std::string language;
|
||||
|
||||
// 0 / 0.0f mean "unset": upstream only overrides the sampler defaults when
|
||||
// the value is strictly positive.
|
||||
int32_t top_k = 0;
|
||||
float top_p = 0.0f;
|
||||
|
||||
// Upper bound on generated audio frames, exposed because the model does not
|
||||
// always emit its codec EOS and will otherwise run to the 512-frame default,
|
||||
// which is roughly 41 s at the 12.5 Hz frame rate. 0 means unset, leaving
|
||||
// that default in place.
|
||||
int32_t max_frames = 0;
|
||||
};
|
||||
|
||||
namespace detail {
|
||||
|
||||
// Strict whole-string numeric parsing. std::stoi/stof accept trailing garbage
|
||||
// ("40abc" -> 40), which would silently honour a typo'd request.
|
||||
inline bool parse_whole_int32(const std::string & value, int32_t & out) {
|
||||
if (value.empty()) {
|
||||
return false;
|
||||
}
|
||||
try {
|
||||
size_t consumed = 0;
|
||||
const long parsed = std::stol(value, &consumed);
|
||||
if (consumed != value.size()) {
|
||||
return false;
|
||||
}
|
||||
if (parsed < INT32_MIN || parsed > INT32_MAX) {
|
||||
return false;
|
||||
}
|
||||
out = static_cast<int32_t>(parsed);
|
||||
return true;
|
||||
} catch (const std::exception &) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
inline bool parse_whole_float(const std::string & value, float & out) {
|
||||
if (value.empty()) {
|
||||
return false;
|
||||
}
|
||||
try {
|
||||
size_t consumed = 0;
|
||||
const float parsed = std::stof(value, &consumed);
|
||||
if (consumed != value.size()) {
|
||||
return false;
|
||||
}
|
||||
out = parsed;
|
||||
return true;
|
||||
} catch (const std::exception &) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
inline tts_request_options reject(const std::string & message) {
|
||||
tts_request_options opts;
|
||||
opts.ok = false;
|
||||
opts.error = message;
|
||||
return opts;
|
||||
}
|
||||
|
||||
} // namespace detail
|
||||
|
||||
inline tts_request_options parse_tts_request_options(
|
||||
const std::string & text,
|
||||
const std::string & voice,
|
||||
const std::string & language,
|
||||
const std::map<std::string, std::string> & params) {
|
||||
if (text.empty()) {
|
||||
return detail::reject("text must be a non-empty string");
|
||||
}
|
||||
|
||||
// The Qwen3-TTS Base checkpoints have no built-in speaker. Without a
|
||||
// reference clip the model picks an arbitrary voice, so an unset voice is
|
||||
// a request error rather than a defaulted one.
|
||||
if (voice.empty()) {
|
||||
return detail::reject("voice must name a speaker reference audio file");
|
||||
}
|
||||
|
||||
tts_request_options opts;
|
||||
opts.text = text;
|
||||
opts.voice_path = voice;
|
||||
opts.language = language;
|
||||
|
||||
// Both values are range-checked here rather than left to the caller: the
|
||||
// consumer copies them straight into mtmd_helper::gen_audio::inp, and only
|
||||
// its separate sampler assignment is guarded by "> 0". An out-of-range or
|
||||
// non-finite value would slip past that guard and reach llama.cpp.
|
||||
const auto top_k_it = params.find("top_k");
|
||||
if (top_k_it != params.end()) {
|
||||
if (!detail::parse_whole_int32(top_k_it->second, opts.top_k)) {
|
||||
return detail::reject("top_k must be an integer, got \"" + top_k_it->second + "\"");
|
||||
}
|
||||
if (opts.top_k < 0) {
|
||||
return detail::reject("top_k must be >= 0, got \"" + top_k_it->second + "\"");
|
||||
}
|
||||
}
|
||||
|
||||
const auto top_p_it = params.find("top_p");
|
||||
if (top_p_it != params.end()) {
|
||||
if (!detail::parse_whole_float(top_p_it->second, opts.top_p)) {
|
||||
return detail::reject("top_p must be a number, got \"" + top_p_it->second + "\"");
|
||||
}
|
||||
// Phrased as a negated in-range test, not "p < 0.0f || p > 1.0f",
|
||||
// because every comparison against NaN is false: the obvious form
|
||||
// would accept NaN, and NaN then defeats the consumer's "> 0" guard
|
||||
// too, since that comparison is false as well.
|
||||
if (!(opts.top_p >= 0.0f && opts.top_p <= 1.0f)) {
|
||||
return detail::reject("top_p must be between 0.0 and 1.0, got \"" + top_p_it->second + "\"");
|
||||
}
|
||||
}
|
||||
|
||||
const auto max_frames_it = params.find("max_frames");
|
||||
if (max_frames_it != params.end()) {
|
||||
if (!detail::parse_whole_int32(max_frames_it->second, opts.max_frames)) {
|
||||
return detail::reject("max_frames must be an integer, got \"" + max_frames_it->second + "\"");
|
||||
}
|
||||
if (opts.max_frames < 0) {
|
||||
return detail::reject("max_frames must be >= 0, got \"" + max_frames_it->second + "\"");
|
||||
}
|
||||
}
|
||||
|
||||
opts.ok = true;
|
||||
return opts;
|
||||
}
|
||||
|
||||
} // namespace llama_grpc
|
||||
@@ -0,0 +1,209 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
#include <cstdio>
|
||||
#include <map>
|
||||
#include <string>
|
||||
|
||||
#include "tts_request_options.h"
|
||||
|
||||
static int failures = 0;
|
||||
|
||||
static void check(bool ok, const char * name) {
|
||||
if (!ok) {
|
||||
++failures;
|
||||
std::fprintf(stderr, "FAIL: %s\n", name);
|
||||
}
|
||||
}
|
||||
|
||||
static void test_accepts_a_minimal_valid_request() {
|
||||
const auto opts = llama_grpc::parse_tts_request_options(
|
||||
"Hello world", "/models/voices/ref.wav", "en", {});
|
||||
|
||||
check(opts.ok, "minimal request is accepted");
|
||||
check(opts.error.empty(), "minimal request has no error");
|
||||
check(opts.text == "Hello world", "text passes through");
|
||||
check(opts.voice_path == "/models/voices/ref.wav", "voice path passes through");
|
||||
check(opts.language == "en", "language passes through");
|
||||
check(opts.top_k == 0, "top_k defaults to the unset sentinel");
|
||||
check(opts.top_p == 0.0f, "top_p defaults to the unset sentinel");
|
||||
check(opts.max_frames == 0, "max_frames defaults to the unset sentinel");
|
||||
}
|
||||
|
||||
static void test_rejects_empty_text() {
|
||||
const auto opts = llama_grpc::parse_tts_request_options(
|
||||
"", "/models/voices/ref.wav", "en", {});
|
||||
|
||||
check(!opts.ok, "empty text is rejected");
|
||||
check(opts.error.find("text") != std::string::npos, "empty-text error names the field");
|
||||
}
|
||||
|
||||
static void test_rejects_missing_speaker_reference() {
|
||||
// Qwen3-TTS Base has no built-in speaker; without a reference it produces
|
||||
// an arbitrary voice, so this must be a hard error rather than a surprise.
|
||||
const auto opts = llama_grpc::parse_tts_request_options(
|
||||
"Hello world", "", "en", {});
|
||||
|
||||
check(!opts.ok, "missing voice is rejected");
|
||||
check(opts.error.find("voice") != std::string::npos, "missing-voice error names the field");
|
||||
}
|
||||
|
||||
static void test_parses_sampling_params() {
|
||||
const std::map<std::string, std::string> params{
|
||||
{"top_k", "40"},
|
||||
{"top_p", "0.85"},
|
||||
};
|
||||
const auto opts = llama_grpc::parse_tts_request_options(
|
||||
"Hello world", "/models/voices/ref.wav", "", params);
|
||||
|
||||
check(opts.ok, "sampling params are accepted");
|
||||
check(opts.top_k == 40, "top_k is parsed");
|
||||
check(opts.top_p > 0.849f && opts.top_p < 0.851f, "top_p is parsed");
|
||||
check(opts.language.empty(), "absent language stays empty");
|
||||
}
|
||||
|
||||
static void test_rejects_malformed_sampling_params() {
|
||||
const auto bad_top_k = llama_grpc::parse_tts_request_options(
|
||||
"Hello world", "/models/voices/ref.wav", "", {{"top_k", "forty"}});
|
||||
check(!bad_top_k.ok, "non-numeric top_k is rejected");
|
||||
check(bad_top_k.error.find("top_k") != std::string::npos, "top_k error names the field");
|
||||
|
||||
const auto bad_top_p = llama_grpc::parse_tts_request_options(
|
||||
"Hello world", "/models/voices/ref.wav", "", {{"top_p", ""}});
|
||||
check(!bad_top_p.ok, "empty top_p is rejected");
|
||||
|
||||
const auto trailing = llama_grpc::parse_tts_request_options(
|
||||
"Hello world", "/models/voices/ref.wav", "", {{"top_k", "40abc"}});
|
||||
check(!trailing.ok, "top_k with trailing garbage is rejected");
|
||||
|
||||
const auto trailing_float = llama_grpc::parse_tts_request_options(
|
||||
"Hello world", "/models/voices/ref.wav", "", {{"top_p", "0.8abc"}});
|
||||
check(!trailing_float.ok, "top_p with trailing garbage is rejected");
|
||||
|
||||
// std::stol returns a long, which is wider than int32_t on 64-bit hosts, so
|
||||
// an in-range-for-long value still has to be caught before the narrowing.
|
||||
const auto overflow_top_k = llama_grpc::parse_tts_request_options(
|
||||
"Hello world", "/models/voices/ref.wav", "", {{"top_k", "99999999999"}});
|
||||
check(!overflow_top_k.ok, "top_k beyond int32 range is rejected");
|
||||
check(overflow_top_k.error.find("top_k") != std::string::npos,
|
||||
"top_k overflow error names the field");
|
||||
}
|
||||
|
||||
static void test_rejects_out_of_range_sampling_params() {
|
||||
// These reach mtmd_helper::gen_audio::inp unconditionally downstream, where
|
||||
// the "> 0" sampler guard does not screen them, so they must die here.
|
||||
const auto negative_top_k = llama_grpc::parse_tts_request_options(
|
||||
"Hello world", "/models/voices/ref.wav", "", {{"top_k", "-5"}});
|
||||
check(!negative_top_k.ok, "negative top_k is rejected");
|
||||
check(negative_top_k.error.find("top_k") != std::string::npos,
|
||||
"negative top_k error names the field");
|
||||
|
||||
const auto negative_top_p = llama_grpc::parse_tts_request_options(
|
||||
"Hello world", "/models/voices/ref.wav", "", {{"top_p", "-0.1"}});
|
||||
check(!negative_top_p.ok, "negative top_p is rejected");
|
||||
check(negative_top_p.error.find("top_p") != std::string::npos,
|
||||
"negative top_p error names the field");
|
||||
|
||||
const auto large_top_p = llama_grpc::parse_tts_request_options(
|
||||
"Hello world", "/models/voices/ref.wav", "", {{"top_p", "1.5"}});
|
||||
check(!large_top_p.ok, "top_p above 1.0 is rejected");
|
||||
|
||||
// NaN survives a naive "p < 0.0f || p > 1.0f" range test because every
|
||||
// comparison against NaN is false. This case pins the correct form.
|
||||
const auto nan_top_p = llama_grpc::parse_tts_request_options(
|
||||
"Hello world", "/models/voices/ref.wav", "", {{"top_p", "nan"}});
|
||||
check(!nan_top_p.ok, "NaN top_p is rejected");
|
||||
|
||||
const auto inf_top_p = llama_grpc::parse_tts_request_options(
|
||||
"Hello world", "/models/voices/ref.wav", "", {{"top_p", "inf"}});
|
||||
check(!inf_top_p.ok, "infinite top_p is rejected");
|
||||
}
|
||||
|
||||
static void test_accepts_sampling_param_boundaries() {
|
||||
const auto zero_top_p = llama_grpc::parse_tts_request_options(
|
||||
"Hello world", "/models/voices/ref.wav", "", {{"top_p", "0.0"}});
|
||||
check(zero_top_p.ok, "top_p of 0.0 is accepted");
|
||||
check(zero_top_p.top_p == 0.0f, "top_p of 0.0 round-trips");
|
||||
|
||||
const auto one_top_p = llama_grpc::parse_tts_request_options(
|
||||
"Hello world", "/models/voices/ref.wav", "", {{"top_p", "1.0"}});
|
||||
check(one_top_p.ok, "top_p of 1.0 is accepted");
|
||||
check(one_top_p.top_p == 1.0f, "top_p of 1.0 round-trips");
|
||||
|
||||
const auto zero_top_k = llama_grpc::parse_tts_request_options(
|
||||
"Hello world", "/models/voices/ref.wav", "", {{"top_k", "0"}});
|
||||
check(zero_top_k.ok, "top_k of 0 is accepted");
|
||||
}
|
||||
|
||||
static void test_parses_max_frames() {
|
||||
// The consumer maps a positive value onto n_predict and leaves upstream's
|
||||
// 512-frame default in place when it is unset, so the sentinel matters as
|
||||
// much as the parsed value.
|
||||
const auto opts = llama_grpc::parse_tts_request_options(
|
||||
"Hello world", "/models/voices/ref.wav", "", {{"max_frames", "120"}});
|
||||
|
||||
check(opts.ok, "max_frames is accepted");
|
||||
check(opts.max_frames == 120, "max_frames is parsed");
|
||||
|
||||
const auto absent = llama_grpc::parse_tts_request_options(
|
||||
"Hello world", "/models/voices/ref.wav", "", {{"top_k", "40"}});
|
||||
check(absent.ok, "a request without max_frames is accepted");
|
||||
check(absent.max_frames == 0, "absent max_frames leaves the unset sentinel");
|
||||
|
||||
const auto zero = llama_grpc::parse_tts_request_options(
|
||||
"Hello world", "/models/voices/ref.wav", "", {{"max_frames", "0"}});
|
||||
check(zero.ok, "max_frames of 0 is accepted");
|
||||
check(zero.max_frames == 0, "max_frames of 0 means unset");
|
||||
}
|
||||
|
||||
static void test_rejects_malformed_max_frames() {
|
||||
const auto negative = llama_grpc::parse_tts_request_options(
|
||||
"Hello world", "/models/voices/ref.wav", "", {{"max_frames", "-1"}});
|
||||
check(!negative.ok, "negative max_frames is rejected");
|
||||
check(negative.error.find("max_frames") != std::string::npos,
|
||||
"negative max_frames error names the field");
|
||||
|
||||
const auto non_numeric = llama_grpc::parse_tts_request_options(
|
||||
"Hello world", "/models/voices/ref.wav", "", {{"max_frames", "many"}});
|
||||
check(!non_numeric.ok, "non-numeric max_frames is rejected");
|
||||
check(non_numeric.error.find("max_frames") != std::string::npos,
|
||||
"non-numeric max_frames error names the field");
|
||||
|
||||
const auto trailing = llama_grpc::parse_tts_request_options(
|
||||
"Hello world", "/models/voices/ref.wav", "", {{"max_frames", "120abc"}});
|
||||
check(!trailing.ok, "max_frames with trailing garbage is rejected");
|
||||
|
||||
const auto empty = llama_grpc::parse_tts_request_options(
|
||||
"Hello world", "/models/voices/ref.wav", "", {{"max_frames", ""}});
|
||||
check(!empty.ok, "empty max_frames is rejected");
|
||||
|
||||
const auto overflow = llama_grpc::parse_tts_request_options(
|
||||
"Hello world", "/models/voices/ref.wav", "", {{"max_frames", "99999999999"}});
|
||||
check(!overflow.ok, "max_frames beyond int32 range is rejected");
|
||||
}
|
||||
|
||||
static void test_ignores_unknown_params() {
|
||||
// Unknown keys are backend-specific knobs meant for other TTS engines. A
|
||||
// request routed here must not fail just because it carries them.
|
||||
const auto opts = llama_grpc::parse_tts_request_options(
|
||||
"Hello world", "/models/voices/ref.wav", "", {{"exaggeration", "0.7"}});
|
||||
|
||||
check(opts.ok, "unknown params are ignored, not rejected");
|
||||
}
|
||||
|
||||
int main() {
|
||||
test_accepts_a_minimal_valid_request();
|
||||
test_rejects_empty_text();
|
||||
test_rejects_missing_speaker_reference();
|
||||
test_parses_sampling_params();
|
||||
test_rejects_malformed_sampling_params();
|
||||
test_rejects_out_of_range_sampling_params();
|
||||
test_accepts_sampling_param_boundaries();
|
||||
test_parses_max_frames();
|
||||
test_rejects_malformed_max_frames();
|
||||
test_ignores_unknown_params();
|
||||
|
||||
if (failures == 0) {
|
||||
std::printf("tts_request_options_test: all checks passed\n");
|
||||
}
|
||||
return failures;
|
||||
}
|
||||
@@ -48,6 +48,7 @@ define turboquant-build
|
||||
# stays compiling against vanilla upstream.
|
||||
bash $(CURRENT_MAKEFILE_DIR)/patch-grpc-server.sh $(CURRENT_MAKEFILE_DIR)/../turboquant-$(1)-build/grpc-server.cpp
|
||||
bash $(LLAMA_CPP_DIR)/disable-score-task.sh $(CURRENT_MAKEFILE_DIR)/../turboquant-$(1)-build/grpc-server.cpp
|
||||
bash $(LLAMA_CPP_DIR)/disable-tts-task.sh $(CURRENT_MAKEFILE_DIR)/../turboquant-$(1)-build/grpc-server.cpp
|
||||
$(info $(GREEN)I turboquant build info:$(1)$(RESET))
|
||||
LLAMA_REPO=$(LLAMA_REPO) LLAMA_VERSION=$(TURBOQUANT_VERSION) \
|
||||
$(MAKE) -C $(CURRENT_MAKEFILE_DIR)/../turboquant-$(1)-build llama.cpp
|
||||
@@ -86,6 +87,7 @@ turboquant-cpu-all:
|
||||
$(MAKE) -C $(CURRENT_MAKEFILE_DIR)/../turboquant-cpu-all-build purge
|
||||
bash $(CURRENT_MAKEFILE_DIR)/patch-grpc-server.sh $(CURRENT_MAKEFILE_DIR)/../turboquant-cpu-all-build/grpc-server.cpp
|
||||
bash $(LLAMA_CPP_DIR)/disable-score-task.sh $(CURRENT_MAKEFILE_DIR)/../turboquant-cpu-all-build/grpc-server.cpp
|
||||
bash $(LLAMA_CPP_DIR)/disable-tts-task.sh $(CURRENT_MAKEFILE_DIR)/../turboquant-cpu-all-build/grpc-server.cpp
|
||||
$(info $(GREEN)I turboquant build info:cpu-all-variants$(RESET))
|
||||
LLAMA_REPO=$(LLAMA_REPO) LLAMA_VERSION=$(TURBOQUANT_VERSION) \
|
||||
$(MAKE) -C $(CURRENT_MAKEFILE_DIR)/../turboquant-cpu-all-build llama.cpp
|
||||
|
||||
@@ -8,7 +8,7 @@ JOBS?=$(shell nproc --ignore=1)
|
||||
|
||||
# CrispASR version (release tag)
|
||||
CRISPASR_REPO?=https://github.com/CrispStrobe/CrispASR
|
||||
CRISPASR_VERSION?=fcb79282a6bc52e13d858026c42b24fb6e63c97a
|
||||
CRISPASR_VERSION?=a153b09b37c90cd55cd9336fccbdf3ba7a289596
|
||||
SO_TARGET?=libgocrispasr.so
|
||||
|
||||
CMAKE_ARGS+=-DBUILD_SHARED_LIBS=OFF
|
||||
|
||||
@@ -67,7 +67,16 @@ const defaultTTSSampleRate = 24000
|
||||
// resampling, so the WAV header must match it. Returns ok=false for non-piper
|
||||
// models (key absent) or an unreadable file, letting the caller fall back to
|
||||
// defaultTTSSampleRate.
|
||||
func piperSampleRate(modelPath string) (int, bool) {
|
||||
func piperSampleRate(modelPath string) (rate int, ok bool) {
|
||||
// A malformed metadata length can make gguf-parser-go panic before it can
|
||||
// return an error. Keep a bad voice file from crash-looping the backend.
|
||||
defer func() {
|
||||
if recover() != nil {
|
||||
rate = 0
|
||||
ok = false
|
||||
}
|
||||
}()
|
||||
|
||||
// Only scalar architecture keys are read, so skip the large array metadata
|
||||
// (phoneme map) and mmap the header - same rationale as pkg/vram's reader.
|
||||
f, err := gguf.ParseGGUFFile(modelPath, gguf.UseMMap(), gguf.SkipLargeMetadata())
|
||||
@@ -78,7 +87,7 @@ func piperSampleRate(modelPath string) (int, bool) {
|
||||
if !ok || kv.ValueType != gguf.GGUFMetadataValueTypeUint32 {
|
||||
return 0, false
|
||||
}
|
||||
rate := int(kv.ValueUint32())
|
||||
rate = int(kv.ValueUint32())
|
||||
if rate <= 0 {
|
||||
return 0, false
|
||||
}
|
||||
|
||||
@@ -3,6 +3,7 @@ package main
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/binary"
|
||||
"math"
|
||||
"os"
|
||||
"path/filepath"
|
||||
|
||||
@@ -102,6 +103,24 @@ var _ = Describe("piper sample rate", func() {
|
||||
_, ok := piperSampleRate(p)
|
||||
Expect(ok).To(BeFalse())
|
||||
})
|
||||
|
||||
It("returns ok=false instead of panicking on a malformed string length", func() {
|
||||
p := filepath.Join(GinkgoT().TempDir(), "malformed.gguf")
|
||||
var b bytes.Buffer
|
||||
b.WriteString("GGUF")
|
||||
Expect(binary.Write(&b, binary.LittleEndian, uint32(3))).To(Succeed())
|
||||
Expect(binary.Write(&b, binary.LittleEndian, uint64(0))).To(Succeed())
|
||||
Expect(binary.Write(&b, binary.LittleEndian, uint64(1))).To(Succeed())
|
||||
key := "general.name"
|
||||
Expect(binary.Write(&b, binary.LittleEndian, uint64(len(key)))).To(Succeed())
|
||||
b.WriteString(key)
|
||||
Expect(binary.Write(&b, binary.LittleEndian, ggufTypeString)).To(Succeed())
|
||||
Expect(binary.Write(&b, binary.LittleEndian, uint64(math.MaxInt64))).To(Succeed())
|
||||
Expect(os.WriteFile(p, b.Bytes(), 0o644)).To(Succeed())
|
||||
|
||||
_, ok := piperSampleRate(p)
|
||||
Expect(ok).To(BeFalse())
|
||||
})
|
||||
})
|
||||
|
||||
// End-to-end through the built .so. Gated on CRISPASR_PIPER_MODEL_PATH (a
|
||||
|
||||
@@ -14,7 +14,7 @@ JOBS?=$(shell nproc --ignore=1)
|
||||
# It is kept alive by the upstream tag da2-support (survives a squash-merge);
|
||||
# repoint to the master merge commit once mudler/depth-anything.cpp PR #1 lands.
|
||||
DEPTHANYTHING_REPO?=https://github.com/mudler/depth-anything.cpp.git
|
||||
DEPTHANYTHING_VERSION?=2028b47ac75a8659c6a9aa617baf09be193eb55f
|
||||
DEPTHANYTHING_VERSION?=54abd5c0abfd1f394e01cb3c38f2e3af4daedf85
|
||||
|
||||
ifeq ($(NATIVE),false)
|
||||
CMAKE_ARGS+=-DGGML_NATIVE=OFF
|
||||
|
||||
@@ -38,8 +38,9 @@ type Store struct {
|
||||
// keysAreNormalized stays true until any non-unit-magnitude key
|
||||
// is added; once false, the magnitude-aware fallback path is
|
||||
// used by Find. Re-evaluated only at Set time, never again on
|
||||
// its own — a deletion of the offending key does NOT flip it
|
||||
// back to true (the bookkeeping cost would dominate the gain).
|
||||
// its own — a partial deletion of the offending key does NOT flip
|
||||
// it back to true (the bookkeeping cost would dominate the gain).
|
||||
// An empty store returns to its initial state.
|
||||
keysAreNormalized bool
|
||||
|
||||
// keyLen is the dimension of every stored key. -1 means "no
|
||||
@@ -142,6 +143,10 @@ func (s *Store) StoresDelete(opts *pb.StoresDeleteOptions) error {
|
||||
mergedV = append(mergedV, tailV...)
|
||||
s.keys = mergedK
|
||||
s.values = mergedV
|
||||
if len(s.keys) == 0 {
|
||||
s.keyLen = -1
|
||||
s.keysAreNormalized = true
|
||||
}
|
||||
assert(slices.IsSortedFunc(s.keys, slices.Compare[[]float32]), "Delete: s.keys not sorted post-merge")
|
||||
assert(len(s.keys) == len(s.values), "Delete: keys/values length skew")
|
||||
return nil
|
||||
|
||||
@@ -105,6 +105,46 @@ var _ = Describe("StoresDelete", func() {
|
||||
})).To(Succeed(), "delete of missing key should succeed")
|
||||
Expect(s.keys).To(HaveLen(1))
|
||||
})
|
||||
|
||||
It("reopens the dimension after deleting every key", func() {
|
||||
s := NewStore()
|
||||
oldKey := []float32{2, 0, 0}
|
||||
mustSet(s, [][]float32{oldKey}, [][]byte{[]byte("3d")})
|
||||
Expect(s.keysAreNormalized).To(BeFalse())
|
||||
|
||||
Expect(s.StoresDelete(&pb.StoresDeleteOptions{
|
||||
Keys: wrapKeys([][]float32{oldKey}),
|
||||
})).To(Succeed())
|
||||
Expect(s.keys).To(BeEmpty())
|
||||
Expect(s.keyLen).To(Equal(-1))
|
||||
Expect(s.keysAreNormalized).To(BeTrue())
|
||||
|
||||
newKey := normalizeVec([]float32{1, 1})
|
||||
mustSet(s, [][]float32{newKey}, [][]byte{[]byte("2d")})
|
||||
res, err := s.StoresFind(&pb.StoresFindOptions{
|
||||
Key: &pb.StoresKey{Floats: newKey},
|
||||
TopK: 1,
|
||||
})
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(res.Values).To(HaveLen(1))
|
||||
Expect(string(res.Values[0].Bytes)).To(Equal("2d"))
|
||||
})
|
||||
|
||||
It("retains the dimension after a partial delete", func() {
|
||||
s := NewStore()
|
||||
mustSet(s,
|
||||
[][]float32{{1, 0, 0}, {0, 1, 0}},
|
||||
[][]byte{[]byte("x"), []byte("y")},
|
||||
)
|
||||
Expect(s.StoresDelete(&pb.StoresDeleteOptions{
|
||||
Keys: wrapKeys([][]float32{{1, 0, 0}}),
|
||||
})).To(Succeed())
|
||||
Expect(s.keyLen).To(Equal(3))
|
||||
Expect(s.StoresSet(&pb.StoresSetOptions{
|
||||
Keys: wrapKeys([][]float32{{1, 0}}),
|
||||
Values: wrapValues([][]byte{[]byte("2d")}),
|
||||
})).NotTo(Succeed())
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("StoresFind", func() {
|
||||
|
||||
@@ -0,0 +1,22 @@
|
||||
# Fetched upstream sources
|
||||
sources/
|
||||
|
||||
# CMake build directories
|
||||
build*/
|
||||
|
||||
# Packaging output
|
||||
package/
|
||||
|
||||
# Compiled backend binary. The second name is what a bare `go build ./...` from
|
||||
# this directory produces (it names the binary after the directory), as opposed
|
||||
# to the -o name the Makefile asks for.
|
||||
nemo-speech-cpp-grpc
|
||||
/nemo-speech-cpp
|
||||
|
||||
# Shared libraries staged in-tree by the Makefile (cp from sources/). The
|
||||
# SOVERSION suffix means the payload is libnemo_speech_*.so.1, hence both globs.
|
||||
*.so
|
||||
*.so.*
|
||||
*.dylib
|
||||
|
||||
compile_commands.json
|
||||
@@ -0,0 +1,391 @@
|
||||
# nemo-speech-cpp backend Makefile.
|
||||
#
|
||||
# Upstream pin lives below as NEMO_SPEECH_VERSION so .github/bump_deps.sh can
|
||||
# find and update it, matching the parakeet-cpp / vibevoice-cpp convention.
|
||||
#
|
||||
# Bumping NEMO_SPEECH_VERSION is a no-op on an existing checkout: sources/ is a
|
||||
# directory target, so make only clones when it is missing and never re-checks
|
||||
# out an already-cloned tree. After a bump run 'make purge && make', the same
|
||||
# rule the parakeet-cpp Makefile documents.
|
||||
#
|
||||
# 'build' is the entry point the backend image calls (backend/Dockerfile.golang
|
||||
# runs 'make -C backend/go/$(BACKEND) build' and then copies package/), so it
|
||||
# has to produce the binary and the package, not just the shared libraries.
|
||||
|
||||
NEMO_SPEECH_VERSION?=4f9676226f667d14608487df744f375db87127f8
|
||||
NEMO_SPEECH_REPO?=https://github.com/NVIDIA/NeMo-Speech.cpp
|
||||
|
||||
GOCMD?=go
|
||||
GO_TAGS?=
|
||||
JOBS?=$(shell nproc 2>/dev/null || sysctl -n hw.ncpu 2>/dev/null || echo 4)
|
||||
|
||||
BUILD_TYPE?=
|
||||
NATIVE?=false
|
||||
|
||||
# NEMO_SPEECH_CUBLAS_SHIM defaults ON upstream and builds a drop-in
|
||||
# libcublas.so.13. LocalAI's CUDA images ship the real cuBLAS, so the shim would
|
||||
# shadow it with a slower native GEMM. Always OFF here.
|
||||
CMAKE_ARGS?=-DCMAKE_BUILD_TYPE=Release \
|
||||
-DBUILD_SHARED_LIBS=OFF \
|
||||
-DCMAKE_POSITION_INDEPENDENT_CODE=ON \
|
||||
-DNEMO_SPEECH_CUBLAS_SHIM=OFF \
|
||||
-DNEMO_SPEECH_BUILD_ASR=ON \
|
||||
-DNEMO_SPEECH_BUILD_DIAR=ON \
|
||||
-DNEMO_SPEECH_BUILD_TTS=ON \
|
||||
-DNEMO_SPEECH_BUILD_NMT=ON \
|
||||
-DNEMO_SPEECH_BUILD_CLI=OFF \
|
||||
-DNEMO_SPEECH_BUILD_HTTP=OFF \
|
||||
-DNEMO_SPEECH_BUILD_GRPC=OFF \
|
||||
-DNEMO_SPEECH_WITH_FLASHLIGHT=OFF \
|
||||
-DNEMO_SPEECH_TTS_WITH_ZH=ON \
|
||||
-DNEMO_SPEECH_TTS_WITH_JA=ON
|
||||
|
||||
ifeq ($(NATIVE),false)
|
||||
CMAKE_ARGS+=-DGGML_NATIVE=OFF
|
||||
endif
|
||||
|
||||
# NEMO_SPEECH_TTS_WITH_JA=ON compiles Open JTalk's bundled MeCab, and
|
||||
# mecab/src/dictionary.cpp derives a comparator from std::binary_function, which
|
||||
# C++17 removed. libstdc++ still ships it as deprecated-but-present under
|
||||
# -std=gnu++17, so Linux never notices; libc++ compiles it out and the build dies
|
||||
# with "no template named 'binary_function' in namespace 'std'". Upstream's own
|
||||
# CMakeLists already carries the equivalent workaround for MSVC's STL
|
||||
# (_HAS_AUTO_PTR_ETC plus /FIfunctional) but has no libc++ branch, because
|
||||
# NEMO_SPEECH_TTS_WITH_JA defaults OFF upstream and only LocalAI turns it on.
|
||||
#
|
||||
# libc++ gates the two templates on _LIBCPP_ENABLE_CXX17_REMOVED_UNARY_BINARY_FUNCTION,
|
||||
# and has done since LLVM 16, which is older than any clang Xcode still ships.
|
||||
# The name matters: the older _LIBCPP_ENABLE_CXX17_REMOVED_BINDERS covers
|
||||
# bind1st/bind2nd/ptr_fun/mem_fun and NOT unary_function/binary_function, and the
|
||||
# umbrella _LIBCPP_ENABLE_CXX17_REMOVED_FEATURES no longer exists at all. A wrong
|
||||
# name is silently accepted by the preprocessor and fixes nothing.
|
||||
#
|
||||
# Applied through CMAKE_CXX_FLAGS rather than to the one target because the
|
||||
# tokenizer CMakeLists is upstream's and this tree is a pinned checkout, not a
|
||||
# patched one. Project-wide is also the safer scope: the macro decides whether
|
||||
# libc++'s internal __binary_function alias resolves to std::binary_function or
|
||||
# to __binary_function_keep_layout_base, which is a base class of std::less and
|
||||
# friends, so defining it for a subset of translation units would give those
|
||||
# class templates two spellings in one binary. Both bases are empty and, at
|
||||
# C++17, carry identical members, so the project-wide define changes no layout
|
||||
# and no ABI. On Linux the macro is not a name libstdc++ knows, so the branch is
|
||||
# unreachable there and would be inert even if it were taken.
|
||||
ifeq ($(shell uname -s),Darwin)
|
||||
CXX_COMPAT_FLAGS?=-D_LIBCPP_ENABLE_CXX17_REMOVED_UNARY_BINARY_FUNCTION
|
||||
else
|
||||
CXX_COMPAT_FLAGS?=
|
||||
endif
|
||||
ifneq ($(strip $(CXX_COMPAT_FLAGS)),)
|
||||
CMAKE_ARGS+=-DCMAKE_CXX_FLAGS=$(CXX_COMPAT_FLAGS)
|
||||
endif
|
||||
|
||||
# scripts/build_itn_deps.sh installs the Sparrowhawk/OpenFST runtime here.
|
||||
# NEMO_SPEECH_DEPENDENCY_PREFIX defaults to <src>/.deps upstream, and the ITN
|
||||
# stack goes under its itn/ subdirectory. ITN_MARKER is a real output of that
|
||||
# script (it prints exactly this file on success), so it can drive a make rule.
|
||||
ITN_PREFIX=sources/NeMo-Speech.cpp/.deps/itn
|
||||
ITN_LIB_DIR=$(ITN_PREFIX)/lib
|
||||
ITN_MARKER=$(ITN_LIB_DIR)/libsparrowhawk.so
|
||||
ITN_FST_HEADER=$(ITN_PREFIX)/include/fst/fst.h
|
||||
|
||||
# SentencePiece became a core ASR dependency in 5be7bfb: RNNT context biasing
|
||||
# uses it even when Flashlight and text normalization are disabled. Build the
|
||||
# pinned static archive provided by upstream so every platform gets the same
|
||||
# dependency instead of relying on an undeclared system package.
|
||||
SENTENCEPIECE_PREFIX=sources/NeMo-Speech.cpp/.deps/sentencepiece
|
||||
SENTENCEPIECE_MARKER=$(SENTENCEPIECE_PREFIX)/lib/libsentencepiece.a
|
||||
|
||||
# Linux's ASR CMake block looks in NEMO_SPEECH_DEPENDENCY_PREFIX directly, but
|
||||
# the Apple branch uses generic find_library()/find_path(). Put the same private
|
||||
# prefix on CMake's search path so Darwin consumes the archive built above too.
|
||||
CMAKE_ARGS+=-DCMAKE_PREFIX_PATH=$(abspath $(SENTENCEPIECE_PREFIX))
|
||||
|
||||
ITN_CC?=gcc-12
|
||||
ITN_CXX?=g++-12
|
||||
|
||||
# Pin protoc to the apt one. backend/Dockerfile.golang drops protoc 27.1 into
|
||||
# /usr/local/bin, which precedes /usr/bin on PATH, while libprotobuf-dev is the
|
||||
# distro's (3.21 on noble, 3.12 on jammy). Sparrowhawk resolves protoc from PATH
|
||||
# at make time (configure.ac uses AC_CHECK_PROG, so PROTOC substitutes to the
|
||||
# bare word, and src/proto/Makefile.am invokes $(PROTOC)), and it commits no
|
||||
# pregenerated stubs, so this always runs. Code generated by 27.1 includes
|
||||
# google/protobuf/runtime_version.h and a PROTOBUF_VERSION #error guard that the
|
||||
# older headers do not have, so the mismatch breaks the build. configure honours
|
||||
# a pre-set PROTOC ("Let the user override the test"), which is what this is.
|
||||
ITN_PROTOC?=/usr/bin/protoc
|
||||
|
||||
# Text normalization is Linux-only: Sparrowhawk/OpenFST assume a GNU toolchain
|
||||
# and the gcc-12 pin has no macOS analogue. Documented gap, see the spec.
|
||||
#
|
||||
# An already-configured build tree wins over the platform default. Without that,
|
||||
# a tree configured WITH_NORM=OFF would silently try to reconfigure itself to ON
|
||||
# on the next bare `make test`, which means demanding gcc-12 from a developer who
|
||||
# deliberately built without it. An explicit WITH_NORM= on the command line still
|
||||
# overrides both, since command-line variables beat ?= assignments.
|
||||
CMAKE_CACHE=sources/NeMo-Speech.cpp/build/CMakeCache.txt
|
||||
CACHED_WITH_NORM=$(shell sed -n 's/^NEMO_SPEECH_WITH_NORM:BOOL=//p' $(CMAKE_CACHE) 2>/dev/null)
|
||||
ifeq ($(shell uname -s),Darwin)
|
||||
WITH_NORM?=OFF
|
||||
else ifneq ($(CACHED_WITH_NORM),)
|
||||
WITH_NORM?=$(CACHED_WITH_NORM)
|
||||
else
|
||||
WITH_NORM?=ON
|
||||
endif
|
||||
CMAKE_ARGS+=-DNEMO_SPEECH_WITH_NORM=$(WITH_NORM)
|
||||
|
||||
ifeq ($(BUILD_TYPE),cublas)
|
||||
CMAKE_ARGS+=-DGGML_CUDA=ON
|
||||
else ifeq ($(BUILD_TYPE),vulkan)
|
||||
CMAKE_ARGS+=-DGGML_VULKAN=ON
|
||||
else ifeq ($(BUILD_TYPE),metal)
|
||||
CMAKE_ARGS+=-DGGML_METAL=ON
|
||||
endif
|
||||
|
||||
# ggml-patches/ is a CUDA series. Every kernel it adds lives under
|
||||
# src/ggml-cuda/; the only files it touches outside that directory are enum and
|
||||
# name-table entries in include/ggml.h and src/ggml.c plus, in ggml-cpu, a
|
||||
# supports_op returning false and an abort case for the CUDA-only op. Upstream
|
||||
# agrees: its metal-* and vulkan-* CMake presets inherit the cpu-* ones, which
|
||||
# set NEMO_SPEECH_GGML_PATCHED=OFF, and every use of a patch-only symbol in the
|
||||
# ASR sources sits behind NEMO_SPEECH_FUSED_RELPOS_ATTN /
|
||||
# NEMO_SPEECH_FASTCONFORMER_CUDA_FUSIONS (both force-OFF without GGML_CUDA) or
|
||||
# behind NEMO_SPEECH_GGML_PATCHED itself, which guards a Q8_PLANAR flag write
|
||||
# that a non-CUDA buffer already throws before reaching.
|
||||
#
|
||||
# So on macOS the series buys nothing, and it cannot be applied there anyway:
|
||||
# upstream's scripts/apply-ggml-patches.sh uses mapfile, a bash 4 builtin, and
|
||||
# macOS ships bash 3.2 as the only bash on the runner's PATH. Skip the patch
|
||||
# step and tell cmake the linked ggml is stock, which is exactly upstream's own
|
||||
# Metal configuration. Linux keeps applying the series unchanged.
|
||||
ifeq ($(shell uname -s),Darwin)
|
||||
GGML_PATCHED?=OFF
|
||||
else
|
||||
GGML_PATCHED?=ON
|
||||
endif
|
||||
CMAKE_ARGS+=-DNEMO_SPEECH_GGML_PATCHED=$(GGML_PATCHED)
|
||||
|
||||
.PHONY: nemo-speech-cpp-grpc package build clean purge test all stage-libs patch-ggml engine itn sentencepiece patch-itn-headers
|
||||
|
||||
all: nemo-speech-cpp-grpc package
|
||||
|
||||
sources/NeMo-Speech.cpp:
|
||||
mkdir -p sources
|
||||
cd sources && git clone $(NEMO_SPEECH_REPO) NeMo-Speech.cpp
|
||||
cd sources/NeMo-Speech.cpp && git checkout $(NEMO_SPEECH_VERSION)
|
||||
# NMT links llama.cpp; ja needs open_jtalk; zh needs cppjieba. flashlight and
|
||||
# kenlm are deliberately not initialized, they are out of scope.
|
||||
cd sources/NeMo-Speech.cpp && git submodule update --init --recursive \
|
||||
ggml llama.cpp third_party/open_jtalk third_party/cppjieba third_party/cpp-httplib
|
||||
|
||||
# NEMO_SPEECH_GGML_PATCHED defaults ON and silently assumes the ggml-patches
|
||||
# series is applied. An unpatched checkout builds fine and produces wrong CUDA
|
||||
# encoder output, so a failure here must stop the build rather than warn.
|
||||
#
|
||||
# Upstream's own script is the right tool: it applies the series in filename
|
||||
# order, exits non-zero when a patch does not apply, and decides "already
|
||||
# applied" by comparing the full-series tree hash rather than a timestamp. That
|
||||
# makes it safe to run unconditionally, so there is no sentinel file to go stale
|
||||
# or to wedge the build when deleted.
|
||||
#
|
||||
# Both branches keep the order-only clone prerequisite: it is the only thing
|
||||
# that pulls sources/ in on a WITH_NORM=OFF tree, where the library rule has no
|
||||
# other prerequisite left.
|
||||
ifeq ($(GGML_PATCHED),ON)
|
||||
patch-ggml: | sources/NeMo-Speech.cpp
|
||||
cd sources/NeMo-Speech.cpp && bash scripts/apply-ggml-patches.sh
|
||||
else
|
||||
patch-ggml: | sources/NeMo-Speech.cpp
|
||||
@echo "[ggml-patch] skipped: NEMO_SPEECH_GGML_PATCHED=$(GGML_PATCHED), the series is CUDA-only"
|
||||
endif
|
||||
|
||||
# The Sparrowhawk/OpenFST text-normalization stack, as a target in its own right
|
||||
# keyed on a file the build script actually produces.
|
||||
#
|
||||
# It used to be a side effect of the runtime library rule, which meant make had
|
||||
# no idea whether it existed: once the library was up to date the script could
|
||||
# never run again, so a tree built WITH_NORM=OFF could not be moved to ON, and
|
||||
# anything that needed the ITN prefix was stuck demanding a full clean. As its
|
||||
# own rule it is built on demand, rebuilt independently, and reachable directly
|
||||
# with 'make itn'.
|
||||
#
|
||||
# OpenFST's templates ICE on gcc-13/14 at -O2, hence the gcc-12 pin for this one
|
||||
# step; the runtime itself builds with the image default compiler.
|
||||
$(ITN_MARKER): | sources/NeMo-Speech.cpp
|
||||
@command -v $(ITN_CC) >/dev/null 2>&1 && command -v $(ITN_CXX) >/dev/null 2>&1 || { \
|
||||
echo "ERROR: $(ITN_CC)/$(ITN_CXX) not found, and text normalization needs them:" >&2; \
|
||||
echo " OpenFST's templates ICE on gcc-13 and gcc-14 at -O2." >&2; \
|
||||
echo " Install them, or build this backend with WITH_NORM=OFF." >&2; \
|
||||
exit 1; }
|
||||
# configure's only gate on a preset PROTOC is test -n, so a path that does not
|
||||
# exist is accepted here and surfaces much later as a bare "No such file or
|
||||
# directory" from inside make -C src/proto. Check it up front instead.
|
||||
@command -v $(ITN_PROTOC) >/dev/null 2>&1 || { \
|
||||
echo "ERROR: protoc not found at $(ITN_PROTOC)." >&2; \
|
||||
echo " Install the protobuf-compiler package, whose protoc matches" >&2; \
|
||||
echo " the libprotobuf-dev headers Sparrowhawk compiles against, or" >&2; \
|
||||
echo " point this at a matching one with ITN_PROTOC=/path/to/protoc." >&2; \
|
||||
exit 1; }
|
||||
cd sources/NeMo-Speech.cpp && CC=$(ITN_CC) CXX=$(ITN_CXX) PROTOC=$(ITN_PROTOC) \
|
||||
JOBS=$(JOBS) scripts/build_itn_deps.sh
|
||||
@$(MAKE) --no-print-directory patch-itn-headers
|
||||
|
||||
# OpenFST 1.8.3's FstImpl copy-assignment operator assigns a raw SymbolTable*
|
||||
# (what SymbolTable::Copy() returns) straight to a std::unique_ptr member:
|
||||
#
|
||||
# isymbols_ = impl.isymbols_ ? impl.isymbols_->Copy() : nullptr;
|
||||
#
|
||||
# std::unique_ptr has no operator= taking a raw pointer in any C++ standard, so
|
||||
# that line is ill-formed everywhere. It survived because nothing instantiates
|
||||
# FstImpl::operator=, and gcc <= 13 only checks a template member's body when it
|
||||
# is instantiated. gcc 14 resolves non-dependent operator expressions at template
|
||||
# definition time, so it rejects the line in every translation unit that so much
|
||||
# as includes <fst/fst.h>, with no instantiation involved. Verified: gcc 14.2
|
||||
# fails on a file whose entire content is '#include <fst/fst.h>'.
|
||||
#
|
||||
# That is why this only shows up now. build_itn_deps.sh builds OpenFST with
|
||||
# gcc-12 (its templates ICE on newer gcc at -O2) and upstream's own images build
|
||||
# the runtime with gcc-13, so neither compiler ever sees it. LocalAI's
|
||||
# backend/Dockerfile.golang installs gcc-14 and makes it the default via
|
||||
# update-alternatives, and fst_normalizer.cpp is the one translation unit here
|
||||
# that includes OpenFST, so it is the one that breaks.
|
||||
#
|
||||
# The fix is the same spelling FstImpl::SetInputSymbols already uses 80 lines
|
||||
# further down, and matches the copy constructor's deep-copy intent exactly. It
|
||||
# is applied to the installed prefix rather than to the OpenFST checkout because
|
||||
# the prefix is the only copy the cmake build compiles against; libfst.so is
|
||||
# already linked by this point and cannot contain the function, since no
|
||||
# compiler could ever have emitted it. Only these two lines are affected: gcc 14
|
||||
# reports exactly two errors over the whole OpenFST include closure, both here.
|
||||
#
|
||||
# Guarded on both sides so a pinned-version bump cannot silently no-op it: the
|
||||
# first check fails if neither the broken nor the fixed spelling is present, the
|
||||
# last fails if the broken one survives.
|
||||
patch-itn-headers:
|
||||
@test -f $(ITN_FST_HEADER) || { \
|
||||
echo "ERROR: $(ITN_FST_HEADER) missing; the ITN prefix is not installed." >&2; \
|
||||
exit 1; }
|
||||
@grep -q 'isymbols_ = impl.isymbols_' $(ITN_FST_HEADER) || \
|
||||
grep -q 'isymbols_.reset(impl.isymbols_' $(ITN_FST_HEADER) || { \
|
||||
echo "ERROR: FstImpl::operator= in $(ITN_FST_HEADER) matches neither the" >&2; \
|
||||
echo " known-broken nor the patched form. OpenFST changed upstream;" >&2; \
|
||||
echo " re-check whether this patch is still needed before removing it." >&2; \
|
||||
exit 1; }
|
||||
sed -i -E 's|^([[:space:]]*)([io]symbols_) = (impl\.[io]symbols_ \? impl\.[io]symbols_->Copy\(\) : nullptr);$$|\1\2.reset(\3);|' $(ITN_FST_HEADER)
|
||||
@if grep -q 'symbols_ = impl.[io]symbols_' $(ITN_FST_HEADER); then \
|
||||
echo "ERROR: the FstImpl::operator= patch did not apply to $(ITN_FST_HEADER)." >&2; \
|
||||
exit 1; \
|
||||
fi
|
||||
|
||||
itn: $(ITN_MARKER)
|
||||
|
||||
$(SENTENCEPIECE_MARKER): | sources/NeMo-Speech.cpp
|
||||
# Upstream's license copies use GNU install's -D flag, which BSD install
|
||||
# does not support. Homebrew CMake 4 also rejects SentencePiece's old policy
|
||||
# floor. Patch both incompatibilities before running the helper on Darwin.
|
||||
@if [ "$(shell uname -s)" = Darwin ]; then \
|
||||
cd sources/NeMo-Speech.cpp && \
|
||||
mkdir -p .deps/sentencepiece/share/licenses/nemo-speech/third_party/sentencepiece && \
|
||||
perl -pi \
|
||||
-e 's/install -Dm0644/install -m 0644/g;' \
|
||||
-e 's/-DCMAKE_BUILD_TYPE=Release /-DCMAKE_BUILD_TYPE=Release -DCMAKE_POLICY_VERSION_MINIMUM=3.5 /;' \
|
||||
scripts/build_sentencepiece_static.sh; \
|
||||
fi
|
||||
cd sources/NeMo-Speech.cpp && JOBS=$(JOBS) scripts/build_sentencepiece_static.sh
|
||||
|
||||
sentencepiece: $(SENTENCEPIECE_MARKER)
|
||||
|
||||
# Only a WITH_NORM=ON build needs the ITN stack, and it must exist before cmake
|
||||
# configures, since the WITH_NORM cmake block find_library()s into the prefix
|
||||
# with REQUIRED.
|
||||
NEMO_RUNTIME_PREREQS=$(SENTENCEPIECE_MARKER)
|
||||
ifeq ($(WITH_NORM),ON)
|
||||
NEMO_RUNTIME_PREREQS+=$(ITN_MARKER)
|
||||
endif
|
||||
|
||||
# Upstream sets CMAKE_LIBRARY_OUTPUT_DIRECTORY to ${CMAKE_BINARY_DIR}/bin, so the
|
||||
# shared objects land in build/bin rather than at the top of the build tree.
|
||||
#
|
||||
# patch-ggml is order-only: it is phony and therefore always runs, but an
|
||||
# order-only prerequisite does not mark this target out of date, so an
|
||||
# already-built tree is not relinked on every invocation.
|
||||
sources/NeMo-Speech.cpp/build/bin/libnemo_speech_asr_c.so: $(NEMO_RUNTIME_PREREQS) | patch-ggml
|
||||
cd sources/NeMo-Speech.cpp && cmake -B build -G Ninja $(CMAKE_ARGS)
|
||||
cd sources/NeMo-Speech.cpp && cmake --build build -j$(JOBS)
|
||||
|
||||
# Stage the runtime next to the Go sources so purego.Dlopen finds it during
|
||||
# local development and so package.sh has a single directory to bundle from.
|
||||
#
|
||||
# ASR and NMT build a dedicated _c shared object that links the C++ implementation
|
||||
# in privately. TTS does not: upstream compiles its c_api.cpp straight into
|
||||
# libnemo_speech_tts and only aliases the nemo_speech_tts_c CMake target, so the
|
||||
# TTS C ABI ships without the _c suffix.
|
||||
stage-libs: sources/NeMo-Speech.cpp/build/bin/libnemo_speech_asr_c.so
|
||||
# -a keeps the SOVERSION symlink a symlink instead of duplicating the payload.
|
||||
cp -af sources/NeMo-Speech.cpp/build/bin/libnemo_speech_asr_c.* .
|
||||
cp -af sources/NeMo-Speech.cpp/build/bin/libnemo_speech_tts.* .
|
||||
cp -af sources/NeMo-Speech.cpp/build/bin/libnemo_speech_nmt_c.* .
|
||||
# The _c libraries are thin ABI shims with a DT_NEEDED on the C++
|
||||
# implementation DSO, so dlopen fails without these next to them. TTS needs
|
||||
# no counterpart, its implementation and ABI live in the same object.
|
||||
cp -af sources/NeMo-Speech.cpp/build/bin/libnemo_speech_asr.* .
|
||||
cp -af sources/NeMo-Speech.cpp/build/bin/libnemo_speech_nmt.* .
|
||||
# nemo_speech_text_normalization is STATIC but links sparrowhawk, fstfar and
|
||||
# fst PUBLIC, so those become DT_NEEDED on libnemo_speech_asr.so. They live in
|
||||
# a project-local prefix that nothing else on the system provides, so without
|
||||
# staging them here the packaged backend cannot dlopen at all.
|
||||
#
|
||||
# Keyed on the prefix existing rather than on WITH_NORM, so this stages what
|
||||
# the tree actually built. A WITH_NORM=ON build cannot reach here without the
|
||||
# prefix (the library rule takes ITN_MARKER as a prerequisite), and if a
|
||||
# library that needs Sparrowhawk somehow arrives unstaged, package.sh's
|
||||
# closure guard fails the build rather than shipping it.
|
||||
@if [ -d "$(ITN_LIB_DIR)" ]; then \
|
||||
echo "cp -af $(ITN_LIB_DIR)/*.so* ."; \
|
||||
cp -af $(ITN_LIB_DIR)/*.so* .; \
|
||||
fi
|
||||
|
||||
## Builds the native runtime and stops short of the Go binary. Everything it
|
||||
## touches lives under sources/, a clone pinned by NEMO_SPEECH_VERSION, so
|
||||
## nothing here can observe a change elsewhere in the LocalAI tree.
|
||||
## Dockerfile.golang calls this from a layer that copies in this directory and
|
||||
## nothing else, which keeps the multi-minute ggml/llama.cpp compile in the
|
||||
## registry layer cache across builds whose only change is on the Go side.
|
||||
## Without it that prebuild is skipped and a CUDA build recompiles all of
|
||||
## upstream on every Go-side edit. See .agents/ci-caching.md.
|
||||
engine: stage-libs
|
||||
|
||||
nemo-speech-cpp-grpc: stage-libs
|
||||
# CGO_ENABLED=0 matches whisper / parakeet-cpp / omnivoice-cpp: the runtime is
|
||||
# reached through purego.Dlopen, not cgo, and a static binary is what lets
|
||||
# run.sh route execution through the packaged lib/ld.so.
|
||||
CGO_ENABLED=0 $(GOCMD) build -tags "$(GO_TAGS)" -o nemo-speech-cpp-grpc .
|
||||
|
||||
# The dlopen tests need the staged shared objects on the loader path, the same
|
||||
# way parakeet-cpp sets it up. Depends on stage-libs so that path is not an
|
||||
# empty directory on a clean tree, which would fail the tests confusingly.
|
||||
#
|
||||
# NEMO_SPEECH_REQUIRE_LIBS turns a missing library from a skip into a failure.
|
||||
# The ABI specs are the only thing standing between this backend and silent
|
||||
# memory corruption, so a run that reaches them and quietly skips them is worse
|
||||
# than one that fails: it reports green having checked nothing.
|
||||
test: stage-libs
|
||||
NEMO_SPEECH_REQUIRE_LIBS=1 LD_LIBRARY_PATH=$(CURDIR):$$LD_LIBRARY_PATH $(GOCMD) test ./... -count=1
|
||||
|
||||
package: nemo-speech-cpp-grpc
|
||||
bash package.sh
|
||||
|
||||
# What backend/Dockerfile.golang invokes. It must leave both the binary and a
|
||||
# populated package/ behind, because the final image stage copies package/.
|
||||
build: package
|
||||
|
||||
clean:
|
||||
# Every .so here is staged output (nemo runtime plus, on a WITH_NORM build,
|
||||
# the ITN stack), and the SOVERSION suffix means the payload is *.so.1, so
|
||||
# the globs have to reach past the .so.
|
||||
rm -f nemo-speech-cpp-grpc
|
||||
rm -f *.so *.so.* *.dylib
|
||||
rm -rf package
|
||||
rm -rf sources/NeMo-Speech.cpp/build
|
||||
|
||||
purge: clean
|
||||
rm -rf sources
|
||||
@@ -0,0 +1,423 @@
|
||||
package main
|
||||
|
||||
// purego binds by name at runtime and the config structs cross the ABI by
|
||||
// pointer, so neither a renamed symbol nor a mis-laid-out mirror struct is
|
||||
// visible to the compiler or the linker. Everything here is transcribed from
|
||||
// sources/NeMo-Speech.cpp/include/nemo_speech/{asr,diar,tts,nmt}.h, and
|
||||
// abi_test.go asserts it against the real shared objects.
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"unsafe"
|
||||
|
||||
"github.com/ebitengine/purego"
|
||||
)
|
||||
|
||||
var (
|
||||
asrLib uintptr
|
||||
ttsLib uintptr
|
||||
nmtLib uintptr
|
||||
)
|
||||
|
||||
// ---- ASR ----
|
||||
|
||||
var (
|
||||
ASRCreate func(cfg unsafe.Pointer, out *uintptr) int32
|
||||
ASRDestroy func(recognizer uintptr)
|
||||
ASRRecognizeF32 func(recognizer uintptr, options unsafe.Pointer, samples *float32, nSamples uint64, sampleRate int32, out *uintptr) int32
|
||||
ASRStreamingRecognize func(recognizer uintptr, options unsafe.Pointer, out *uintptr) int32
|
||||
ASRStreamPushF32 func(stream uintptr, samples *float32, nSamples uint64, sampleRate int32) int32
|
||||
ASRStreamForceEndpoint func(stream uintptr) int32
|
||||
ASRStreamFinish func(stream uintptr) int32
|
||||
ASRStreamNext func(stream uintptr, out *uintptr) int32
|
||||
ASRStreamClose func(stream uintptr)
|
||||
ASRRecognitionOptionsDef func() cASRRecognitionOptions
|
||||
|
||||
ASRResultIsFinal func(result uintptr) bool
|
||||
ASRResultAudioProcessed func(result uintptr) float32
|
||||
ASRResultAlternativeCount func(result uintptr) uint64
|
||||
ASRResultTranscript func(result uintptr, alt uint64) string
|
||||
ASRResultConfidence func(result uintptr, alt uint64) float32
|
||||
ASRResultWordCount func(result uintptr, alt uint64) uint64
|
||||
ASRResultWordText func(result uintptr, alt, i uint64) string
|
||||
ASRResultWordStartTime func(result uintptr, alt, i uint64) int32
|
||||
ASRResultWordEndTime func(result uintptr, alt, i uint64) int32
|
||||
ASRResultWordConfidence func(result uintptr, alt, i uint64) float32
|
||||
ASRResultWordSpeakerTag func(result uintptr, alt, i uint64) int32
|
||||
ASRResultLanguageCount func(result uintptr, alt uint64) uint64
|
||||
ASRResultLanguageCode func(result uintptr, alt, i uint64) string
|
||||
ASRResultDestroy func(result uintptr)
|
||||
|
||||
ASRLastError func() string
|
||||
ASRVersion func() string
|
||||
)
|
||||
|
||||
// ---- Diarization (exported from the ASR library) ----
|
||||
|
||||
var (
|
||||
DiarCreate func(cfg unsafe.Pointer, out *uintptr) int32
|
||||
DiarDestroy func(model uintptr)
|
||||
DiarNumSpeakers func(model uintptr) int32
|
||||
DiarSecondsPerFrame func(model uintptr) float64
|
||||
DiarStreamOpen func(model uintptr, out *uintptr) int32
|
||||
DiarStreamPushF32 func(stream uintptr, samples *float32, nSamples uint64, sampleRate int32) int32
|
||||
DiarStreamFinish func(stream uintptr) int32
|
||||
DiarStreamClose func(stream uintptr)
|
||||
// cfg is the optional nemo_speech_diar_segmentation_config (NULL = library
|
||||
// defaults). The two-call count-then-fill pattern is documented on the C
|
||||
// declaration in diar.h.
|
||||
DiarSegments func(stream uintptr, cfg unsafe.Pointer, out unsafe.Pointer, capacity uint64, count *uint64) int32
|
||||
)
|
||||
|
||||
// ---- TTS ----
|
||||
|
||||
var (
|
||||
TTSCreate func(cfg unsafe.Pointer, out *uintptr) int32
|
||||
TTSDestroy func(synthesizer uintptr)
|
||||
TTSSampleRate func(synthesizer uintptr) int32
|
||||
TTSSpeakerCount func(synthesizer uintptr) int32
|
||||
TTSSpeakerName func(synthesizer uintptr, i uint64) string
|
||||
TTSSynthesizeText func(synthesizer uintptr, options unsafe.Pointer, text string, callback uintptr, userData uintptr, statsOut unsafe.Pointer) int32
|
||||
TTSRuntimeConfigDefault func() cTTSRuntimeConfig
|
||||
TTSSynthesisOptionsDefault func() cTTSSynthesisOptions
|
||||
TTSLastError func() string
|
||||
TTSVersion func() string
|
||||
)
|
||||
|
||||
// ---- NMT ----
|
||||
|
||||
var (
|
||||
NMTCreate func(cfg unsafe.Pointer, out *uintptr) int32
|
||||
NMTDestroy func(translator uintptr)
|
||||
NMTTranslate func(translator uintptr, texts *uintptr, nTexts uint64, source, target string, out *uintptr) int32
|
||||
NMTResultCount func(result uintptr) uint64
|
||||
NMTResultText func(result uintptr, i uint64) string
|
||||
NMTResultLanguage func(result uintptr, i uint64) string
|
||||
NMTResultDestroy func(result uintptr)
|
||||
NMTLastError func() string
|
||||
NMTVersion func() string
|
||||
)
|
||||
|
||||
// ---- C struct mirrors ----
|
||||
//
|
||||
// Each mirrors a struct in include/nemo_speech/*.h field for field. The leading
|
||||
// Size field is the C `size_t size` the runtime validates against its own
|
||||
// sizeof, which is what makes a layout mismatch detectable at runtime instead
|
||||
// of silently corrupting memory. Blank fields are System V AMD64 / AAPCS64
|
||||
// padding: C inserts it implicitly, Go does not, so it has to be written out.
|
||||
// See abi_test.go, which pins both every total size and every field offset.
|
||||
|
||||
type cASRBackendConfig struct {
|
||||
Size uintptr
|
||||
GPU int32
|
||||
_ [4]byte // trailing pad to the struct's 8-byte alignment
|
||||
}
|
||||
|
||||
type cASRModelConfig struct {
|
||||
Size uintptr
|
||||
Path uintptr
|
||||
Name uintptr
|
||||
}
|
||||
|
||||
type cASRVADConfig struct {
|
||||
Size uintptr
|
||||
ModelPath uintptr
|
||||
EnableMasking bool
|
||||
_ [3]byte
|
||||
Onset float32
|
||||
Offset float32
|
||||
_ [4]byte
|
||||
}
|
||||
|
||||
type cASRPostprocConfig struct {
|
||||
Size uintptr
|
||||
ProfanityListPath uintptr
|
||||
ITNModelDir uintptr
|
||||
PNCModelPath uintptr
|
||||
}
|
||||
|
||||
type cASRDiarConfig struct {
|
||||
Size uintptr
|
||||
ModelPath uintptr
|
||||
ChunkFrames int32
|
||||
RightContextFrames int32
|
||||
LeftContextFrames int32
|
||||
FIFOFrames int32
|
||||
SpkcacheFrames int32
|
||||
UpdatePeriodFrames int32
|
||||
}
|
||||
|
||||
type cASRRecognizerConfig struct {
|
||||
Size uintptr
|
||||
Backend uintptr
|
||||
Model uintptr
|
||||
Streaming uintptr
|
||||
Decoder uintptr
|
||||
VAD uintptr
|
||||
Endpointing uintptr
|
||||
Postproc uintptr
|
||||
Diar uintptr
|
||||
Batching uintptr
|
||||
}
|
||||
|
||||
type cASRRecognitionOptions struct {
|
||||
Size uintptr
|
||||
RequestID uintptr
|
||||
LanguageCode uintptr
|
||||
InterimResults bool
|
||||
EnableWordTimeOffsets bool
|
||||
EnableAutomaticPunctuation bool
|
||||
VerbatimTranscripts bool
|
||||
ProfanityFilter bool
|
||||
_ [3]byte
|
||||
StopHistoryEouMs int32
|
||||
_ [4]byte
|
||||
SpeechContexts uintptr
|
||||
SpeechContextCount uintptr
|
||||
MaxAlternatives int32
|
||||
EnableSpeakerDiarization bool
|
||||
_ [3]byte
|
||||
MaxSpeakerCount int32
|
||||
_ [4]byte
|
||||
}
|
||||
|
||||
// cDiarModelConfig mirrors nemo_speech_diar_model_config (diar.h). This is the
|
||||
// standalone Sortformer pipeline's own config and is NOT cASRDiarConfig, which
|
||||
// is the diarizer attached to a recognizer: this one carries gpu and preset,
|
||||
// that one does not.
|
||||
//
|
||||
// The six frame counts are sentinel-sensitive. src/asr/c_api.cpp applies each
|
||||
// one only when it is > 0, EXCEPT left_context_frames, which it applies when it
|
||||
// is >= 0. A zero-valued struct would therefore pin the left context to 0
|
||||
// rather than leave the preset's value alone, so loadDiarizer writes -1 into
|
||||
// all six.
|
||||
type cDiarModelConfig struct {
|
||||
Size uintptr
|
||||
ModelPath uintptr
|
||||
GPU int32
|
||||
_ [4]byte // pad to the alignment of the pointer that follows
|
||||
Preset uintptr
|
||||
// Encoder-frame geometry overrides, applied on top of the preset.
|
||||
ChunkFrames int32
|
||||
RightContextFrames int32
|
||||
LeftContextFrames int32
|
||||
FIFOFrames int32
|
||||
SpkcacheFrames int32
|
||||
UpdatePeriodFrames int32
|
||||
}
|
||||
|
||||
// cDiarSegmentationConfig mirrors nemo_speech_diar_segmentation_config
|
||||
// (diar.h): the NeMo ts_vad postprocessing applied when turning per-frame
|
||||
// speaker probabilities into segments.
|
||||
//
|
||||
// onset and offset are float, the four durations are double. That mixture is
|
||||
// the whole reason this mirror needs its offsets pinned: writing all six as
|
||||
// float32 or all six as float64 both produce a struct C would read shifted.
|
||||
type cDiarSegmentationConfig struct {
|
||||
Size uintptr
|
||||
Onset float32
|
||||
Offset float32
|
||||
PadOnsetSec float64
|
||||
PadOffsetSec float64
|
||||
MinGapSec float64
|
||||
MinDurationSec float64
|
||||
}
|
||||
|
||||
// cDiarSegment mirrors nemo_speech_diar_segment (diar.h), the element type
|
||||
// nemo_speech_diar_segments fills.
|
||||
//
|
||||
// It has no leading size field: unlike the config structs it travels from C to
|
||||
// Go, so there is no caller-declared size for the runtime to validate against.
|
||||
// The times are already SECONDS (double), not frame indices, so nothing here
|
||||
// needs the model's seconds-per-frame to be interpreted. Speaker is 1-based,
|
||||
// matching WordInfo.speaker_tag on the ASR surface.
|
||||
type cDiarSegment struct {
|
||||
StartTime float64
|
||||
EndTime float64
|
||||
Speaker int32
|
||||
_ [4]byte // trailing pad to the struct's 8-byte alignment
|
||||
}
|
||||
|
||||
type cTTSModelConfig struct {
|
||||
Size uintptr
|
||||
MagpieModel uintptr
|
||||
CodecModel uintptr
|
||||
TokenizerModelDir uintptr
|
||||
TextNormalizerModelDir uintptr
|
||||
}
|
||||
|
||||
// cTTSRuntimeConfig mirrors nemo_speech_tts_runtime_config. The four backend /
|
||||
// mode fields are C enums, which this toolchain lays out as int32.
|
||||
type cTTSRuntimeConfig struct {
|
||||
Size uintptr
|
||||
Speaker int32
|
||||
Threads int32
|
||||
CodecThreads int32
|
||||
Seed int32
|
||||
Steps int32
|
||||
TopK int32
|
||||
ChunkFrames int32
|
||||
CodecQueueDepth int32
|
||||
CodecHistoryFrames int32
|
||||
CodecFutureFrames int32
|
||||
WindowMs int32
|
||||
Temperature float32
|
||||
OverrideTemperature bool
|
||||
_ [3]byte
|
||||
CFGScale float32
|
||||
OverrideCFGScale bool
|
||||
UseCFG bool
|
||||
UseLocalTransformer bool
|
||||
UseKVCache bool
|
||||
UseStatefulCodec bool
|
||||
CodecCPU bool
|
||||
FlushPartialChunk bool
|
||||
Verbose bool
|
||||
LTBackend int32
|
||||
SamplingBackend int32
|
||||
UMAMode int32
|
||||
LongformMode int32
|
||||
LTFP32 bool
|
||||
_ [7]byte
|
||||
}
|
||||
|
||||
type cTTSSynthesizerConfig struct {
|
||||
Size uintptr
|
||||
Model uintptr
|
||||
Runtime uintptr
|
||||
DefaultLanguageCode uintptr
|
||||
DefaultVoiceName uintptr
|
||||
}
|
||||
|
||||
type cTTSSynthesisOptions struct {
|
||||
Size uintptr
|
||||
RequestID uintptr
|
||||
LanguageCode uintptr
|
||||
Speaker int32
|
||||
Seed int32
|
||||
Steps int32
|
||||
TopK int32
|
||||
Temperature float32
|
||||
OverrideTemperature bool
|
||||
_ [3]byte
|
||||
CFGScale float32
|
||||
OverrideCFGScale bool
|
||||
_ [3]byte
|
||||
VoiceName uintptr
|
||||
OutputSampleRate int32
|
||||
_ [4]byte
|
||||
}
|
||||
|
||||
type cNMTBackendConfig struct {
|
||||
Size uintptr
|
||||
GPU int32
|
||||
_ [4]byte
|
||||
}
|
||||
|
||||
type cNMTModelConfig struct {
|
||||
Size uintptr
|
||||
Path uintptr
|
||||
NCtx int32
|
||||
_ [4]byte
|
||||
}
|
||||
|
||||
type cNMTTranslatorConfig struct {
|
||||
Size uintptr
|
||||
Backend uintptr
|
||||
Model uintptr
|
||||
Generation uintptr
|
||||
Pool uintptr
|
||||
}
|
||||
|
||||
// symbol pairs a Go function pointer with its exported C name. Keeping the
|
||||
// name next to the var means `nm -D libnemo_speech_asr_c.so.1 | grep nemo_speech`
|
||||
// is enough to spot drift after a pin bump.
|
||||
type symbol struct {
|
||||
fn any
|
||||
name string
|
||||
lib *uintptr
|
||||
}
|
||||
|
||||
func symbols() []symbol {
|
||||
return []symbol{
|
||||
{&ASRCreate, "nemo_speech_asr_create", &asrLib},
|
||||
{&ASRDestroy, "nemo_speech_asr_destroy", &asrLib},
|
||||
{&ASRRecognizeF32, "nemo_speech_asr_recognize_f32", &asrLib},
|
||||
{&ASRStreamingRecognize, "nemo_speech_asr_streaming_recognize", &asrLib},
|
||||
{&ASRStreamPushF32, "nemo_speech_asr_stream_push_f32", &asrLib},
|
||||
{&ASRStreamForceEndpoint, "nemo_speech_asr_stream_force_endpoint", &asrLib},
|
||||
{&ASRStreamFinish, "nemo_speech_asr_stream_finish", &asrLib},
|
||||
{&ASRStreamNext, "nemo_speech_asr_stream_next", &asrLib},
|
||||
{&ASRStreamClose, "nemo_speech_asr_stream_close", &asrLib},
|
||||
{&ASRRecognitionOptionsDef, "nemo_speech_asr_recognition_options_default", &asrLib},
|
||||
{&ASRResultIsFinal, "nemo_speech_asr_result_is_final", &asrLib},
|
||||
{&ASRResultAudioProcessed, "nemo_speech_asr_result_audio_processed", &asrLib},
|
||||
{&ASRResultAlternativeCount, "nemo_speech_asr_result_alternative_count", &asrLib},
|
||||
{&ASRResultTranscript, "nemo_speech_asr_result_transcript", &asrLib},
|
||||
{&ASRResultConfidence, "nemo_speech_asr_result_confidence", &asrLib},
|
||||
{&ASRResultWordCount, "nemo_speech_asr_result_word_count", &asrLib},
|
||||
{&ASRResultWordText, "nemo_speech_asr_result_word_text", &asrLib},
|
||||
{&ASRResultWordStartTime, "nemo_speech_asr_result_word_start_time", &asrLib},
|
||||
{&ASRResultWordEndTime, "nemo_speech_asr_result_word_end_time", &asrLib},
|
||||
{&ASRResultWordConfidence, "nemo_speech_asr_result_word_confidence", &asrLib},
|
||||
{&ASRResultWordSpeakerTag, "nemo_speech_asr_result_word_speaker_tag", &asrLib},
|
||||
{&ASRResultLanguageCount, "nemo_speech_asr_result_language_count", &asrLib},
|
||||
{&ASRResultLanguageCode, "nemo_speech_asr_result_language_code", &asrLib},
|
||||
{&ASRResultDestroy, "nemo_speech_asr_result_destroy", &asrLib},
|
||||
{&ASRLastError, "nemo_speech_asr_last_error", &asrLib},
|
||||
{&ASRVersion, "nemo_speech_asr_version", &asrLib},
|
||||
|
||||
{&DiarCreate, "nemo_speech_diar_create", &asrLib},
|
||||
{&DiarDestroy, "nemo_speech_diar_destroy", &asrLib},
|
||||
{&DiarNumSpeakers, "nemo_speech_diar_num_speakers", &asrLib},
|
||||
{&DiarSecondsPerFrame, "nemo_speech_diar_seconds_per_frame", &asrLib},
|
||||
{&DiarStreamOpen, "nemo_speech_diar_stream_open", &asrLib},
|
||||
{&DiarStreamPushF32, "nemo_speech_diar_stream_push_f32", &asrLib},
|
||||
{&DiarStreamFinish, "nemo_speech_diar_stream_finish", &asrLib},
|
||||
{&DiarStreamClose, "nemo_speech_diar_stream_close", &asrLib},
|
||||
{&DiarSegments, "nemo_speech_diar_segments", &asrLib},
|
||||
|
||||
{&TTSCreate, "nemo_speech_tts_create", &ttsLib},
|
||||
{&TTSDestroy, "nemo_speech_tts_destroy", &ttsLib},
|
||||
{&TTSSampleRate, "nemo_speech_tts_sample_rate", &ttsLib},
|
||||
{&TTSSpeakerCount, "nemo_speech_tts_speaker_count", &ttsLib},
|
||||
{&TTSSpeakerName, "nemo_speech_tts_speaker_name", &ttsLib},
|
||||
{&TTSSynthesizeText, "nemo_speech_tts_synthesize_text", &ttsLib},
|
||||
{&TTSRuntimeConfigDefault, "nemo_speech_tts_runtime_config_default", &ttsLib},
|
||||
{&TTSSynthesisOptionsDefault, "nemo_speech_tts_synthesis_options_default", &ttsLib},
|
||||
{&TTSLastError, "nemo_speech_tts_last_error", &ttsLib},
|
||||
{&TTSVersion, "nemo_speech_tts_version", &ttsLib},
|
||||
|
||||
{&NMTCreate, "nemo_speech_nmt_create", &nmtLib},
|
||||
{&NMTDestroy, "nemo_speech_nmt_destroy", &nmtLib},
|
||||
{&NMTTranslate, "nemo_speech_nmt_translate", &nmtLib},
|
||||
{&NMTResultCount, "nemo_speech_nmt_result_count", &nmtLib},
|
||||
{&NMTResultText, "nemo_speech_nmt_result_text", &nmtLib},
|
||||
{&NMTResultLanguage, "nemo_speech_nmt_result_language", &nmtLib},
|
||||
{&NMTResultDestroy, "nemo_speech_nmt_result_destroy", &nmtLib},
|
||||
{&NMTLastError, "nemo_speech_nmt_last_error", &nmtLib},
|
||||
{&NMTVersion, "nemo_speech_nmt_version", &nmtLib},
|
||||
}
|
||||
}
|
||||
|
||||
// registerSymbols binds every entry point. purego panics on a missing symbol,
|
||||
// so this recovers and returns the offending name: after an upstream pin bump a
|
||||
// rename must fail loudly at startup, not at first inference.
|
||||
func registerSymbols() error {
|
||||
for _, s := range symbols() {
|
||||
if err := registerOne(s); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func registerOne(s symbol) (err error) {
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
err = fmt.Errorf("nemo-speech-cpp: binding %q: %v", s.name, r)
|
||||
}
|
||||
}()
|
||||
purego.RegisterLibFunc(s.fn, *s.lib, s.name)
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,317 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"os"
|
||||
"unsafe"
|
||||
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
// requireLibs reports whether a missing shared library must fail the specs
|
||||
// instead of skipping them.
|
||||
//
|
||||
// librariesPresent stats bare filenames relative to the working directory,
|
||||
// while openLibraries resolves them through the loader search path, so the two
|
||||
// can legitimately disagree. Under `make test` that is harmless because the
|
||||
// stage-libs prerequisite puts the .so files in the working directory, but any
|
||||
// other invocation would skip every library-backed spec and still report a
|
||||
// green run. The Makefile sets NEMO_SPEECH_REQUIRE_LIBS=1 so no CI path can
|
||||
// pass on a silent skip; leaving it unset keeps the pure-Go specs runnable on a
|
||||
// checkout with no build.
|
||||
func requireLibs() bool {
|
||||
return os.Getenv("NEMO_SPEECH_REQUIRE_LIBS") == "1"
|
||||
}
|
||||
|
||||
// librariesPresent reports whether a local build is available to bind against.
|
||||
func librariesPresent() bool {
|
||||
for _, n := range []string{
|
||||
libraryName("NEMO_SPEECH_ASR_LIBRARY", "libnemo_speech_asr_c"),
|
||||
libraryName("NEMO_SPEECH_TTS_LIBRARY", "libnemo_speech_tts"),
|
||||
libraryName("NEMO_SPEECH_NMT_LIBRARY", "libnemo_speech_nmt_c"),
|
||||
} {
|
||||
if _, err := os.Stat(n); err != nil {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// layout is one expected number transcribed from the C headers.
|
||||
type layout struct {
|
||||
what string
|
||||
got uintptr
|
||||
want uintptr
|
||||
}
|
||||
|
||||
// The `want` column is what a C compiler reports for the structs in
|
||||
// include/nemo_speech/{asr,diar,tts,nmt}.h under the System V AMD64 / AAPCS64 rules
|
||||
// both supported targets follow. Regenerate after an upstream pin bump with a
|
||||
// throwaway program over the installed headers:
|
||||
//
|
||||
// printf('SIZE %%zu\n', sizeof(nemo_speech_asr_recognition_options));
|
||||
// printf('OFF %%zu\n', offsetof(nemo_speech_asr_recognition_options, max_speaker_count));
|
||||
//
|
||||
// Sizes alone are not enough: two padding mistakes can cancel out and leave the
|
||||
// total unchanged while every field between them reads from the wrong offset,
|
||||
// so each mirror pins its field offsets too.
|
||||
func structSizes() []layout {
|
||||
return []layout{
|
||||
{"cASRBackendConfig", unsafe.Sizeof(cASRBackendConfig{}), 16},
|
||||
{"cASRModelConfig", unsafe.Sizeof(cASRModelConfig{}), 24},
|
||||
{"cASRVADConfig", unsafe.Sizeof(cASRVADConfig{}), 32},
|
||||
{"cASRPostprocConfig", unsafe.Sizeof(cASRPostprocConfig{}), 32},
|
||||
{"cASRDiarConfig", unsafe.Sizeof(cASRDiarConfig{}), 40},
|
||||
{"cASRRecognizerConfig", unsafe.Sizeof(cASRRecognizerConfig{}), 80},
|
||||
{"cASRRecognitionOptions", unsafe.Sizeof(cASRRecognitionOptions{}), 72},
|
||||
{"cDiarModelConfig", unsafe.Sizeof(cDiarModelConfig{}), 56},
|
||||
{"cDiarSegmentationConfig", unsafe.Sizeof(cDiarSegmentationConfig{}), 48},
|
||||
{"cDiarSegment", unsafe.Sizeof(cDiarSegment{}), 24},
|
||||
{"cTTSModelConfig", unsafe.Sizeof(cTTSModelConfig{}), 40},
|
||||
{"cTTSRuntimeConfig", unsafe.Sizeof(cTTSRuntimeConfig{}), 96},
|
||||
{"cTTSSynthesizerConfig", unsafe.Sizeof(cTTSSynthesizerConfig{}), 40},
|
||||
{"cTTSSynthesisOptions", unsafe.Sizeof(cTTSSynthesisOptions{}), 72},
|
||||
{"cNMTBackendConfig", unsafe.Sizeof(cNMTBackendConfig{}), 16},
|
||||
{"cNMTModelConfig", unsafe.Sizeof(cNMTModelConfig{}), 24},
|
||||
{"cNMTTranslatorConfig", unsafe.Sizeof(cNMTTranslatorConfig{}), 40},
|
||||
}
|
||||
}
|
||||
|
||||
func structOffsets() []layout {
|
||||
return []layout{
|
||||
{"cASRBackendConfig.GPU", unsafe.Offsetof(cASRBackendConfig{}.GPU), 8},
|
||||
|
||||
{"cASRModelConfig.Path", unsafe.Offsetof(cASRModelConfig{}.Path), 8},
|
||||
{"cASRModelConfig.Name", unsafe.Offsetof(cASRModelConfig{}.Name), 16},
|
||||
|
||||
{"cASRVADConfig.ModelPath", unsafe.Offsetof(cASRVADConfig{}.ModelPath), 8},
|
||||
{"cASRVADConfig.EnableMasking", unsafe.Offsetof(cASRVADConfig{}.EnableMasking), 16},
|
||||
{"cASRVADConfig.Onset", unsafe.Offsetof(cASRVADConfig{}.Onset), 20},
|
||||
{"cASRVADConfig.Offset", unsafe.Offsetof(cASRVADConfig{}.Offset), 24},
|
||||
|
||||
{"cASRPostprocConfig.ProfanityListPath", unsafe.Offsetof(cASRPostprocConfig{}.ProfanityListPath), 8},
|
||||
{"cASRPostprocConfig.ITNModelDir", unsafe.Offsetof(cASRPostprocConfig{}.ITNModelDir), 16},
|
||||
{"cASRPostprocConfig.PNCModelPath", unsafe.Offsetof(cASRPostprocConfig{}.PNCModelPath), 24},
|
||||
|
||||
{"cASRDiarConfig.ModelPath", unsafe.Offsetof(cASRDiarConfig{}.ModelPath), 8},
|
||||
{"cASRDiarConfig.ChunkFrames", unsafe.Offsetof(cASRDiarConfig{}.ChunkFrames), 16},
|
||||
{"cASRDiarConfig.RightContextFrames", unsafe.Offsetof(cASRDiarConfig{}.RightContextFrames), 20},
|
||||
{"cASRDiarConfig.LeftContextFrames", unsafe.Offsetof(cASRDiarConfig{}.LeftContextFrames), 24},
|
||||
{"cASRDiarConfig.FIFOFrames", unsafe.Offsetof(cASRDiarConfig{}.FIFOFrames), 28},
|
||||
{"cASRDiarConfig.SpkcacheFrames", unsafe.Offsetof(cASRDiarConfig{}.SpkcacheFrames), 32},
|
||||
{"cASRDiarConfig.UpdatePeriodFrames", unsafe.Offsetof(cASRDiarConfig{}.UpdatePeriodFrames), 36},
|
||||
|
||||
{"cASRRecognizerConfig.Backend", unsafe.Offsetof(cASRRecognizerConfig{}.Backend), 8},
|
||||
{"cASRRecognizerConfig.Model", unsafe.Offsetof(cASRRecognizerConfig{}.Model), 16},
|
||||
{"cASRRecognizerConfig.Streaming", unsafe.Offsetof(cASRRecognizerConfig{}.Streaming), 24},
|
||||
{"cASRRecognizerConfig.Decoder", unsafe.Offsetof(cASRRecognizerConfig{}.Decoder), 32},
|
||||
{"cASRRecognizerConfig.VAD", unsafe.Offsetof(cASRRecognizerConfig{}.VAD), 40},
|
||||
{"cASRRecognizerConfig.Endpointing", unsafe.Offsetof(cASRRecognizerConfig{}.Endpointing), 48},
|
||||
{"cASRRecognizerConfig.Postproc", unsafe.Offsetof(cASRRecognizerConfig{}.Postproc), 56},
|
||||
{"cASRRecognizerConfig.Diar", unsafe.Offsetof(cASRRecognizerConfig{}.Diar), 64},
|
||||
{"cASRRecognizerConfig.Batching", unsafe.Offsetof(cASRRecognizerConfig{}.Batching), 72},
|
||||
|
||||
{"cASRRecognitionOptions.RequestID", unsafe.Offsetof(cASRRecognitionOptions{}.RequestID), 8},
|
||||
{"cASRRecognitionOptions.LanguageCode", unsafe.Offsetof(cASRRecognitionOptions{}.LanguageCode), 16},
|
||||
{"cASRRecognitionOptions.InterimResults", unsafe.Offsetof(cASRRecognitionOptions{}.InterimResults), 24},
|
||||
{"cASRRecognitionOptions.EnableWordTimeOffsets", unsafe.Offsetof(cASRRecognitionOptions{}.EnableWordTimeOffsets), 25},
|
||||
{"cASRRecognitionOptions.EnableAutomaticPunctuation", unsafe.Offsetof(cASRRecognitionOptions{}.EnableAutomaticPunctuation), 26},
|
||||
{"cASRRecognitionOptions.VerbatimTranscripts", unsafe.Offsetof(cASRRecognitionOptions{}.VerbatimTranscripts), 27},
|
||||
{"cASRRecognitionOptions.ProfanityFilter", unsafe.Offsetof(cASRRecognitionOptions{}.ProfanityFilter), 28},
|
||||
{"cASRRecognitionOptions.StopHistoryEouMs", unsafe.Offsetof(cASRRecognitionOptions{}.StopHistoryEouMs), 32},
|
||||
{"cASRRecognitionOptions.SpeechContexts", unsafe.Offsetof(cASRRecognitionOptions{}.SpeechContexts), 40},
|
||||
{"cASRRecognitionOptions.SpeechContextCount", unsafe.Offsetof(cASRRecognitionOptions{}.SpeechContextCount), 48},
|
||||
{"cASRRecognitionOptions.MaxAlternatives", unsafe.Offsetof(cASRRecognitionOptions{}.MaxAlternatives), 56},
|
||||
{"cASRRecognitionOptions.EnableSpeakerDiarization", unsafe.Offsetof(cASRRecognitionOptions{}.EnableSpeakerDiarization), 60},
|
||||
{"cASRRecognitionOptions.MaxSpeakerCount", unsafe.Offsetof(cASRRecognitionOptions{}.MaxSpeakerCount), 64},
|
||||
|
||||
{"cDiarModelConfig.ModelPath", unsafe.Offsetof(cDiarModelConfig{}.ModelPath), 8},
|
||||
{"cDiarModelConfig.GPU", unsafe.Offsetof(cDiarModelConfig{}.GPU), 16},
|
||||
{"cDiarModelConfig.Preset", unsafe.Offsetof(cDiarModelConfig{}.Preset), 24},
|
||||
{"cDiarModelConfig.ChunkFrames", unsafe.Offsetof(cDiarModelConfig{}.ChunkFrames), 32},
|
||||
{"cDiarModelConfig.RightContextFrames", unsafe.Offsetof(cDiarModelConfig{}.RightContextFrames), 36},
|
||||
{"cDiarModelConfig.LeftContextFrames", unsafe.Offsetof(cDiarModelConfig{}.LeftContextFrames), 40},
|
||||
{"cDiarModelConfig.FIFOFrames", unsafe.Offsetof(cDiarModelConfig{}.FIFOFrames), 44},
|
||||
{"cDiarModelConfig.SpkcacheFrames", unsafe.Offsetof(cDiarModelConfig{}.SpkcacheFrames), 48},
|
||||
{"cDiarModelConfig.UpdatePeriodFrames", unsafe.Offsetof(cDiarModelConfig{}.UpdatePeriodFrames), 52},
|
||||
|
||||
{"cDiarSegmentationConfig.Onset", unsafe.Offsetof(cDiarSegmentationConfig{}.Onset), 8},
|
||||
{"cDiarSegmentationConfig.Offset", unsafe.Offsetof(cDiarSegmentationConfig{}.Offset), 12},
|
||||
{"cDiarSegmentationConfig.PadOnsetSec", unsafe.Offsetof(cDiarSegmentationConfig{}.PadOnsetSec), 16},
|
||||
{"cDiarSegmentationConfig.PadOffsetSec", unsafe.Offsetof(cDiarSegmentationConfig{}.PadOffsetSec), 24},
|
||||
{"cDiarSegmentationConfig.MinGapSec", unsafe.Offsetof(cDiarSegmentationConfig{}.MinGapSec), 32},
|
||||
{"cDiarSegmentationConfig.MinDurationSec", unsafe.Offsetof(cDiarSegmentationConfig{}.MinDurationSec), 40},
|
||||
|
||||
{"cDiarSegment.StartTime", unsafe.Offsetof(cDiarSegment{}.StartTime), 0},
|
||||
{"cDiarSegment.EndTime", unsafe.Offsetof(cDiarSegment{}.EndTime), 8},
|
||||
{"cDiarSegment.Speaker", unsafe.Offsetof(cDiarSegment{}.Speaker), 16},
|
||||
|
||||
{"cTTSModelConfig.MagpieModel", unsafe.Offsetof(cTTSModelConfig{}.MagpieModel), 8},
|
||||
{"cTTSModelConfig.CodecModel", unsafe.Offsetof(cTTSModelConfig{}.CodecModel), 16},
|
||||
{"cTTSModelConfig.TokenizerModelDir", unsafe.Offsetof(cTTSModelConfig{}.TokenizerModelDir), 24},
|
||||
{"cTTSModelConfig.TextNormalizerModelDir", unsafe.Offsetof(cTTSModelConfig{}.TextNormalizerModelDir), 32},
|
||||
|
||||
{"cTTSRuntimeConfig.Speaker", unsafe.Offsetof(cTTSRuntimeConfig{}.Speaker), 8},
|
||||
{"cTTSRuntimeConfig.Threads", unsafe.Offsetof(cTTSRuntimeConfig{}.Threads), 12},
|
||||
{"cTTSRuntimeConfig.CodecThreads", unsafe.Offsetof(cTTSRuntimeConfig{}.CodecThreads), 16},
|
||||
{"cTTSRuntimeConfig.Seed", unsafe.Offsetof(cTTSRuntimeConfig{}.Seed), 20},
|
||||
{"cTTSRuntimeConfig.Steps", unsafe.Offsetof(cTTSRuntimeConfig{}.Steps), 24},
|
||||
{"cTTSRuntimeConfig.TopK", unsafe.Offsetof(cTTSRuntimeConfig{}.TopK), 28},
|
||||
{"cTTSRuntimeConfig.ChunkFrames", unsafe.Offsetof(cTTSRuntimeConfig{}.ChunkFrames), 32},
|
||||
{"cTTSRuntimeConfig.CodecQueueDepth", unsafe.Offsetof(cTTSRuntimeConfig{}.CodecQueueDepth), 36},
|
||||
{"cTTSRuntimeConfig.CodecHistoryFrames", unsafe.Offsetof(cTTSRuntimeConfig{}.CodecHistoryFrames), 40},
|
||||
{"cTTSRuntimeConfig.CodecFutureFrames", unsafe.Offsetof(cTTSRuntimeConfig{}.CodecFutureFrames), 44},
|
||||
{"cTTSRuntimeConfig.WindowMs", unsafe.Offsetof(cTTSRuntimeConfig{}.WindowMs), 48},
|
||||
{"cTTSRuntimeConfig.Temperature", unsafe.Offsetof(cTTSRuntimeConfig{}.Temperature), 52},
|
||||
{"cTTSRuntimeConfig.OverrideTemperature", unsafe.Offsetof(cTTSRuntimeConfig{}.OverrideTemperature), 56},
|
||||
{"cTTSRuntimeConfig.CFGScale", unsafe.Offsetof(cTTSRuntimeConfig{}.CFGScale), 60},
|
||||
{"cTTSRuntimeConfig.OverrideCFGScale", unsafe.Offsetof(cTTSRuntimeConfig{}.OverrideCFGScale), 64},
|
||||
{"cTTSRuntimeConfig.UseCFG", unsafe.Offsetof(cTTSRuntimeConfig{}.UseCFG), 65},
|
||||
{"cTTSRuntimeConfig.UseLocalTransformer", unsafe.Offsetof(cTTSRuntimeConfig{}.UseLocalTransformer), 66},
|
||||
{"cTTSRuntimeConfig.UseKVCache", unsafe.Offsetof(cTTSRuntimeConfig{}.UseKVCache), 67},
|
||||
{"cTTSRuntimeConfig.UseStatefulCodec", unsafe.Offsetof(cTTSRuntimeConfig{}.UseStatefulCodec), 68},
|
||||
{"cTTSRuntimeConfig.CodecCPU", unsafe.Offsetof(cTTSRuntimeConfig{}.CodecCPU), 69},
|
||||
{"cTTSRuntimeConfig.FlushPartialChunk", unsafe.Offsetof(cTTSRuntimeConfig{}.FlushPartialChunk), 70},
|
||||
{"cTTSRuntimeConfig.Verbose", unsafe.Offsetof(cTTSRuntimeConfig{}.Verbose), 71},
|
||||
{"cTTSRuntimeConfig.LTBackend", unsafe.Offsetof(cTTSRuntimeConfig{}.LTBackend), 72},
|
||||
{"cTTSRuntimeConfig.SamplingBackend", unsafe.Offsetof(cTTSRuntimeConfig{}.SamplingBackend), 76},
|
||||
{"cTTSRuntimeConfig.UMAMode", unsafe.Offsetof(cTTSRuntimeConfig{}.UMAMode), 80},
|
||||
{"cTTSRuntimeConfig.LongformMode", unsafe.Offsetof(cTTSRuntimeConfig{}.LongformMode), 84},
|
||||
{"cTTSRuntimeConfig.LTFP32", unsafe.Offsetof(cTTSRuntimeConfig{}.LTFP32), 88},
|
||||
|
||||
{"cTTSSynthesizerConfig.Model", unsafe.Offsetof(cTTSSynthesizerConfig{}.Model), 8},
|
||||
{"cTTSSynthesizerConfig.Runtime", unsafe.Offsetof(cTTSSynthesizerConfig{}.Runtime), 16},
|
||||
{"cTTSSynthesizerConfig.DefaultLanguageCode", unsafe.Offsetof(cTTSSynthesizerConfig{}.DefaultLanguageCode), 24},
|
||||
{"cTTSSynthesizerConfig.DefaultVoiceName", unsafe.Offsetof(cTTSSynthesizerConfig{}.DefaultVoiceName), 32},
|
||||
|
||||
{"cTTSSynthesisOptions.RequestID", unsafe.Offsetof(cTTSSynthesisOptions{}.RequestID), 8},
|
||||
{"cTTSSynthesisOptions.LanguageCode", unsafe.Offsetof(cTTSSynthesisOptions{}.LanguageCode), 16},
|
||||
{"cTTSSynthesisOptions.Speaker", unsafe.Offsetof(cTTSSynthesisOptions{}.Speaker), 24},
|
||||
{"cTTSSynthesisOptions.Seed", unsafe.Offsetof(cTTSSynthesisOptions{}.Seed), 28},
|
||||
{"cTTSSynthesisOptions.Steps", unsafe.Offsetof(cTTSSynthesisOptions{}.Steps), 32},
|
||||
{"cTTSSynthesisOptions.TopK", unsafe.Offsetof(cTTSSynthesisOptions{}.TopK), 36},
|
||||
{"cTTSSynthesisOptions.Temperature", unsafe.Offsetof(cTTSSynthesisOptions{}.Temperature), 40},
|
||||
{"cTTSSynthesisOptions.OverrideTemperature", unsafe.Offsetof(cTTSSynthesisOptions{}.OverrideTemperature), 44},
|
||||
{"cTTSSynthesisOptions.CFGScale", unsafe.Offsetof(cTTSSynthesisOptions{}.CFGScale), 48},
|
||||
{"cTTSSynthesisOptions.OverrideCFGScale", unsafe.Offsetof(cTTSSynthesisOptions{}.OverrideCFGScale), 52},
|
||||
{"cTTSSynthesisOptions.VoiceName", unsafe.Offsetof(cTTSSynthesisOptions{}.VoiceName), 56},
|
||||
{"cTTSSynthesisOptions.OutputSampleRate", unsafe.Offsetof(cTTSSynthesisOptions{}.OutputSampleRate), 64},
|
||||
|
||||
{"cNMTBackendConfig.GPU", unsafe.Offsetof(cNMTBackendConfig{}.GPU), 8},
|
||||
{"cNMTModelConfig.Path", unsafe.Offsetof(cNMTModelConfig{}.Path), 8},
|
||||
{"cNMTModelConfig.NCtx", unsafe.Offsetof(cNMTModelConfig{}.NCtx), 16},
|
||||
|
||||
{"cNMTTranslatorConfig.Backend", unsafe.Offsetof(cNMTTranslatorConfig{}.Backend), 8},
|
||||
{"cNMTTranslatorConfig.Model", unsafe.Offsetof(cNMTTranslatorConfig{}.Model), 16},
|
||||
{"cNMTTranslatorConfig.Generation", unsafe.Offsetof(cNMTTranslatorConfig{}.Generation), 24},
|
||||
{"cNMTTranslatorConfig.Pool", unsafe.Offsetof(cNMTTranslatorConfig{}.Pool), 32},
|
||||
}
|
||||
}
|
||||
|
||||
var _ = Describe("C struct mirrors", func() {
|
||||
// These need no shared object, so they run on any checkout and catch a
|
||||
// transcription slip the moment it is introduced.
|
||||
It("matches the C sizeof of every mirrored struct", func() {
|
||||
for _, l := range structSizes() {
|
||||
Expect(l.got).To(Equal(l.want), "%s: Go mirror is %d bytes, C says %d", l.what, l.got, l.want)
|
||||
}
|
||||
})
|
||||
|
||||
It("matches the C offset of every mirrored field", func() {
|
||||
for _, l := range structOffsets() {
|
||||
Expect(l.got).To(Equal(l.want), "%s: Go offset %d, C offset %d", l.what, l.got, l.want)
|
||||
}
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("C ABI binding", func() {
|
||||
BeforeEach(func() {
|
||||
if !librariesPresent() {
|
||||
if requireLibs() {
|
||||
cwd, _ := os.Getwd()
|
||||
Fail("NEMO_SPEECH_REQUIRE_LIBS=1 but the shared libraries are not in " + cwd +
|
||||
": these specs are the ABI defence and must not be skipped." +
|
||||
" Run make -C backend/go/nemo-speech-cpp stage-libs")
|
||||
}
|
||||
Skip("shared libraries not built, run make in backend/go/nemo-speech-cpp")
|
||||
}
|
||||
Expect(openLibraries()).To(Succeed())
|
||||
})
|
||||
|
||||
It("resolves every bound symbol", func() {
|
||||
Expect(symbols()).ToNot(BeEmpty())
|
||||
for _, s := range symbols() {
|
||||
Expect(registerOne(s)).To(Succeed())
|
||||
}
|
||||
})
|
||||
|
||||
// The library reports its own sizeof through the size field of each
|
||||
// defaults struct. A Go mirror that disagrees means every field after the
|
||||
// first divergence is read from the wrong offset, which no compiler or
|
||||
// linker check would catch. The three structs below are the only ones with
|
||||
// a defaults entry point, so they are the only ones the runtime can be
|
||||
// asked about directly.
|
||||
It("mirrors the C recognition-options struct layout", func() {
|
||||
def := ASRRecognitionOptionsDef()
|
||||
Expect(def.Size).To(Equal(unsafe.Sizeof(cASRRecognitionOptions{})),
|
||||
"cASRRecognitionOptions does not match the C layout")
|
||||
})
|
||||
|
||||
It("mirrors the C TTS runtime-config struct layout", func() {
|
||||
def := TTSRuntimeConfigDefault()
|
||||
Expect(def.Size).To(Equal(unsafe.Sizeof(cTTSRuntimeConfig{})),
|
||||
"cTTSRuntimeConfig does not match the C layout")
|
||||
})
|
||||
|
||||
It("mirrors the C TTS synthesis-options struct layout", func() {
|
||||
def := TTSSynthesisOptionsDefault()
|
||||
Expect(def.Size).To(Equal(unsafe.Sizeof(cTTSSynthesisOptions{})),
|
||||
"cTTSSynthesisOptions does not match the C layout")
|
||||
})
|
||||
|
||||
// A size match alone cannot see a field read from the wrong offset when two
|
||||
// padding mistakes cancel out, and structOffsets checks the mirrors against
|
||||
// numbers transcribed by the same hand that wrote them. This spec is the
|
||||
// only layer independent of that transcription: it reads values back out of
|
||||
// the running library, so a systematically wrong table cannot hide here.
|
||||
//
|
||||
// Deliberately narrow. An earlier version pinned roughly forty default
|
||||
// values, which would make a legitimate pin bump (threads 4 to 8, or a
|
||||
// flipped flush_partial_chunk) fail with a message that reads like a layout
|
||||
// error. What survives is only the values that are contract, not tuning:
|
||||
//
|
||||
// - max_alternatives is the single non-zero in an otherwise memset-zero
|
||||
// struct, and asr.h documents "<= 1 = 1-best (default)". It pins offset
|
||||
// 56, deep in the tail past the bool run.
|
||||
// - The synthesis-options run of four -1 sentinels, each documented in
|
||||
// tts.h as "< 0 = synthesizer default", pins offsets 24 through 36, and
|
||||
// temperature witnesses that the run stops exactly at offset 40. A
|
||||
// mirror whose tail is shifted by one field spills a -1 into that zero.
|
||||
// - Two -1 sentinels at the ends of the runtime config's long int32 run
|
||||
// pin offset 20 and offset 40 without depending on any tunable.
|
||||
//
|
||||
// Sources: src/asr/c_api.cpp nemo_speech_asr_recognition_options_default,
|
||||
// src/tts/magpietts/runtime.h MagpieRuntimeConfig, src/tts/c_api.cpp
|
||||
// nemo_speech_tts_synthesis_options_default.
|
||||
It("reads the documented default values back through the mirrors", func() {
|
||||
asr := ASRRecognitionOptionsDef()
|
||||
Expect(asr.MaxAlternatives).To(Equal(int32(1)))
|
||||
|
||||
rt := TTSRuntimeConfigDefault()
|
||||
Expect(rt.Seed).To(Equal(int32(-1)))
|
||||
Expect(rt.CodecHistoryFrames).To(Equal(int32(-1)))
|
||||
|
||||
opt := TTSSynthesisOptionsDefault()
|
||||
Expect(opt.Speaker).To(Equal(int32(-1)))
|
||||
Expect(opt.Seed).To(Equal(int32(-1)))
|
||||
Expect(opt.Steps).To(Equal(int32(-1)))
|
||||
Expect(opt.TopK).To(Equal(int32(-1)))
|
||||
Expect(opt.Temperature).To(Equal(float32(0)))
|
||||
})
|
||||
|
||||
It("reports a non-empty version from each library", func() {
|
||||
Expect(ASRVersion()).ToNot(BeEmpty())
|
||||
Expect(TTSVersion()).ToNot(BeEmpty())
|
||||
Expect(NMTVersion()).ToNot(BeEmpty())
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,374 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"runtime"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
"unsafe"
|
||||
|
||||
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
|
||||
"github.com/mudler/xlog"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
)
|
||||
|
||||
// asrWord is one decoded word with its millisecond offsets and 1-based speaker
|
||||
// tag (0 when diarization was not requested).
|
||||
type asrWord struct {
|
||||
Text string
|
||||
Start int32
|
||||
End int32
|
||||
Speaker int32
|
||||
}
|
||||
|
||||
// pinPtr pins v for the lifetime of p and returns its address in the uintptr
|
||||
// form the config structs carry.
|
||||
//
|
||||
// The config structs mirror C, so their pointer members are uintptr, which the
|
||||
// collector does not trace. Everything reachable only through one of them is
|
||||
// therefore invisible to the GC while C is reading it, exactly as described on
|
||||
// cstr, and needs the same pin. runtime.KeepAlive would cover collection but
|
||||
// says nothing about relocation, and the guarantee wanted here is that the
|
||||
// address C holds stays the address of the object.
|
||||
func pinPtr[T any](p *runtime.Pinner, v *T) uintptr {
|
||||
p.Pin(v)
|
||||
// #nosec G103 -- v is pinned into p on the previous line, so its address is
|
||||
// stable and traced for as long as p lives; every caller defers p.Unpin only
|
||||
// after the create call that reads it. One-way, like cstr: nothing converts
|
||||
// this uintptr back to a pointer.
|
||||
return uintptr(unsafe.Pointer(v))
|
||||
}
|
||||
|
||||
// asrDiarConfig builds the config for the diarizer attached to a recognizer.
|
||||
//
|
||||
// Extracted from loadASR for the same reason diarModelConfig was extracted from
|
||||
// loadDiarizer: the six frame counts are sentinel-sensitive and invisible to
|
||||
// every other check in the tree. src/asr/c_api.cpp:151-165 applies five of them
|
||||
// when they are > 0 but applies left_context_frames when it is >= 0, so a
|
||||
// dropped -1 does not fall back to the model's own streaming geometry, it pins
|
||||
// the left context to zero. The struct is the right shape either way, so the
|
||||
// layout assertions in abi_test.go cannot see it and only a spec on this builder
|
||||
// can.
|
||||
//
|
||||
// diarGeometryDefault is shared with the standalone diarizer rather than
|
||||
// restated: it is the same sentinel, from the same rule, in the same runtime.
|
||||
//
|
||||
// modelPath is a C pointer from cstr, not a Go string, and the caller owns its
|
||||
// release.
|
||||
func asrDiarConfig(modelPath uintptr) cASRDiarConfig {
|
||||
return cASRDiarConfig{
|
||||
Size: unsafe.Sizeof(cASRDiarConfig{}),
|
||||
ModelPath: modelPath,
|
||||
ChunkFrames: diarGeometryDefault,
|
||||
RightContextFrames: diarGeometryDefault,
|
||||
LeftContextFrames: diarGeometryDefault,
|
||||
FIFOFrames: diarGeometryDefault,
|
||||
SpkcacheFrames: diarGeometryDefault,
|
||||
UpdatePeriodFrames: diarGeometryDefault,
|
||||
}
|
||||
}
|
||||
|
||||
// loadASR creates the recognizer, attaching VAD, PnC, ITN and diarization when
|
||||
// the corresponding options were set.
|
||||
//
|
||||
// Every field below is assigned by name against include/nemo_speech/asr.h. The
|
||||
// sub-configs are optional pointers: a nil one means "library defaults", which
|
||||
// is why each is populated only when its option was given rather than always
|
||||
// being attached with empty strings.
|
||||
//
|
||||
// Each struct's Size is load-bearing, not decoration. The runtime decides a
|
||||
// field is present with HAS_FIELD (src/asr/c_api.cpp), which tests the caller's
|
||||
// size against offsetof(field) + sizeof(field), so a config sent with Size 0
|
||||
// has every field ignored and the model silently loads with defaults.
|
||||
//
|
||||
// This must not take engineMu: Load is its only caller and already holds it.
|
||||
func (n *NemoSpeech) loadASR(modelFile string) error {
|
||||
// nemo_speech_asr_create deep-copies every const char* into a std::string
|
||||
// (src/asr/c_api.cpp to_config, via str_or_empty) and retains no pointer
|
||||
// afterwards, so pinning for the duration of the create call is both
|
||||
// necessary and sufficient.
|
||||
var pinner runtime.Pinner
|
||||
defer pinner.Unpin()
|
||||
|
||||
pathP, freePath := cstr(modelFile)
|
||||
defer freePath()
|
||||
|
||||
model := cASRModelConfig{Size: unsafe.Sizeof(cASRModelConfig{}), Path: pathP}
|
||||
backend := cASRBackendConfig{Size: unsafe.Sizeof(cASRBackendConfig{}), GPU: n.opts.gpu}
|
||||
|
||||
cfg := cASRRecognizerConfig{
|
||||
Size: unsafe.Sizeof(cASRRecognizerConfig{}),
|
||||
Backend: pinPtr(&pinner, &backend),
|
||||
Model: pinPtr(&pinner, &model),
|
||||
}
|
||||
|
||||
var vad cASRVADConfig
|
||||
if n.opts.vadModel != "" {
|
||||
p, free := cstr(n.opts.vadModel)
|
||||
defer free()
|
||||
vad = cASRVADConfig{Size: unsafe.Sizeof(cASRVADConfig{}), ModelPath: p}
|
||||
cfg.VAD = pinPtr(&pinner, &vad)
|
||||
}
|
||||
|
||||
var postproc cASRPostprocConfig
|
||||
if n.opts.itnDir != "" || n.opts.pncModel != "" {
|
||||
itnP, freeITN := cstr(n.opts.itnDir)
|
||||
defer freeITN()
|
||||
pncP, freePNC := cstr(n.opts.pncModel)
|
||||
defer freePNC()
|
||||
postproc = cASRPostprocConfig{
|
||||
Size: unsafe.Sizeof(cASRPostprocConfig{}),
|
||||
ITNModelDir: itnP,
|
||||
PNCModelPath: pncP,
|
||||
}
|
||||
cfg.Postproc = pinPtr(&pinner, &postproc)
|
||||
}
|
||||
|
||||
var diar cASRDiarConfig
|
||||
if n.opts.diarModel != "" {
|
||||
p, free := cstr(n.opts.diarModel)
|
||||
defer free()
|
||||
diar = asrDiarConfig(p)
|
||||
cfg.Diar = pinPtr(&pinner, &diar)
|
||||
}
|
||||
|
||||
xlog.Info("nemo-speech-cpp: creating recognizer",
|
||||
"gpu", n.opts.gpu,
|
||||
"vad", n.opts.vadModel != "",
|
||||
"pnc", n.opts.pncModel != "",
|
||||
"itn", n.opts.itnDir != "",
|
||||
"diarization", n.opts.diarModel != "")
|
||||
|
||||
// #nosec G103 -- cfg is a local POD struct passed as a pointer for the
|
||||
// duration of this call only; every uintptr member it carries is either a
|
||||
// cstr allocation or a pinPtr address, all pinned above and released by the
|
||||
// defers, and nemo_speech_asr_create deep-copies and retains nothing.
|
||||
if st := ASRCreate(unsafe.Pointer(&cfg), &n.recognizer); st != 0 {
|
||||
return statusErrorf(st, "nemo-speech-cpp: asr create: %s", ASRLastError())
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// recognizeF32 runs one offline decode and returns the result handle, which the
|
||||
// caller must destroy.
|
||||
//
|
||||
// The empty-input guard is here rather than at the call site because &pcm[0]
|
||||
// panics on a zero-length slice: Go never reaches the C side's own "empty
|
||||
// audio" rejection. A silent clip or a truncated upload decodes to zero
|
||||
// samples, which is ordinary input, not an exotic one.
|
||||
//
|
||||
// The caller must hold engineMu.
|
||||
func recognizeF32(recognizer uintptr, opts *cASRRecognitionOptions, pcm []float32, sampleRate int32) (uintptr, error) {
|
||||
if len(pcm) == 0 {
|
||||
return 0, status.Error(codes.InvalidArgument, "nemo-speech-cpp: empty audio")
|
||||
}
|
||||
|
||||
var result uintptr
|
||||
// #nosec G103 -- opts is the caller's live struct, borrowed for this call
|
||||
// only; its LanguageCode is a cstr allocation the caller keeps pinned across
|
||||
// it. &pcm[0] is guarded by the empty check above and the length handed over
|
||||
// is exactly len(pcm), so the runtime cannot read past the slice.
|
||||
if st := ASRRecognizeF32(recognizer, unsafe.Pointer(opts),
|
||||
&pcm[0], uint64(len(pcm)), sampleRate, &result); st != 0 {
|
||||
return 0, statusErrorf(st, "nemo-speech-cpp: recognize: %s", ASRLastError())
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// msToNanos converts a runtime word offset to the wire unit. The runtime
|
||||
// reports milliseconds (src/asr/types.h); TranscriptSegment.start/end and
|
||||
// TranscriptWord.start/end are int64 nanoseconds, which core/backend reads
|
||||
// straight into a time.Duration.
|
||||
func msToNanos(ms int32) int64 {
|
||||
return int64(ms) * int64(time.Millisecond)
|
||||
}
|
||||
|
||||
// extractWords pulls the top alternative's words out of a result handle.
|
||||
func extractWords(result uintptr) []asrWord {
|
||||
if ASRResultAlternativeCount(result) == 0 {
|
||||
return nil
|
||||
}
|
||||
count := ASRResultWordCount(result, 0)
|
||||
words := make([]asrWord, 0, count)
|
||||
for i := uint64(0); i < count; i++ {
|
||||
words = append(words, asrWord{
|
||||
Text: ASRResultWordText(result, 0, i),
|
||||
Start: ASRResultWordStartTime(result, 0, i),
|
||||
End: ASRResultWordEndTime(result, 0, i),
|
||||
Speaker: ASRResultWordSpeakerTag(result, 0, i),
|
||||
})
|
||||
}
|
||||
return words
|
||||
}
|
||||
|
||||
// wordsRequested reports whether the caller asked for word-level timestamps.
|
||||
// The OpenAI transcription API gates word timings behind
|
||||
// timestamp_granularities[] containing "word" and defaults to segment level
|
||||
// otherwise; every backend here follows that contract (see
|
||||
// backend/go/parakeet-cpp).
|
||||
func wordsRequested(granularities []string) bool {
|
||||
for _, g := range granularities {
|
||||
if strings.EqualFold(strings.TrimSpace(g), "word") {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// wordsToSegments groups words into one segment per consecutive speaker run.
|
||||
// Without diarization every word carries speaker 0, so this collapses to a
|
||||
// single segment.
|
||||
//
|
||||
// The boundary is a CHANGE of speaker, not the first appearance of one: a
|
||||
// conversation that returns to an earlier speaker has to start a new turn
|
||||
// rather than reopen the old one.
|
||||
//
|
||||
// withWords additionally attaches the per-word timings that
|
||||
// core/backend/transcript.go turns into the response's word list. It is off by
|
||||
// default because the OpenAI contract asks for word timestamps explicitly, and
|
||||
// a long transcript pays for every word twice otherwise.
|
||||
func wordsToSegments(words []asrWord, withWords bool) []*pb.TranscriptSegment {
|
||||
if len(words) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
var segs []*pb.TranscriptSegment
|
||||
start := 0
|
||||
flush := func(end int) {
|
||||
run := words[start:end]
|
||||
texts := make([]string, 0, len(run))
|
||||
for _, w := range run {
|
||||
texts = append(texts, w.Text)
|
||||
}
|
||||
seg := &pb.TranscriptSegment{
|
||||
// #nosec G115 -- TranscriptSegment.Id is int32 on the wire, and segs
|
||||
// holds one entry per speaker run over the words of a single decode
|
||||
// result, which exhausts memory long before it reaches 2^31.
|
||||
Id: int32(len(segs)),
|
||||
Text: strings.Join(texts, " "),
|
||||
Start: msToNanos(run[0].Start),
|
||||
End: msToNanos(run[len(run)-1].End),
|
||||
}
|
||||
// The speaker tag is 1-based with 0 meaning untagged, so an undiarized
|
||||
// run must stay unlabelled rather than be attributed to a speaker "0".
|
||||
if run[0].Speaker > 0 {
|
||||
seg.Speaker = strconv.Itoa(int(run[0].Speaker))
|
||||
}
|
||||
if withWords {
|
||||
seg.Words = wordsToProto(run)
|
||||
}
|
||||
segs = append(segs, seg)
|
||||
}
|
||||
|
||||
for i := 1; i < len(words); i++ {
|
||||
if words[i].Speaker != words[start].Speaker {
|
||||
flush(i)
|
||||
start = i
|
||||
}
|
||||
}
|
||||
flush(len(words))
|
||||
return segs
|
||||
}
|
||||
|
||||
// AudioTranscription decodes the audio at req.Dst and returns one offline
|
||||
// transcription.
|
||||
//
|
||||
// The whole body runs inside withEngine, so the family check and the C calls
|
||||
// that trust the handle happen under a single acquisition of engineMu. Decoding
|
||||
// the audio is in there too: pkg/grpc/server.go already serialises RPCs on this
|
||||
// backend through base.SingleThread, so the lock costs no concurrency, and the
|
||||
// alternative (check, unlock, decode, relock) is the exact gap Free can land in.
|
||||
func (n *NemoSpeech) AudioTranscription(ctx context.Context, req *pb.TranscriptRequest) (pb.TranscriptResult, error) {
|
||||
var out *pb.TranscriptResult
|
||||
if err := n.withEngine(familyASR, func() error {
|
||||
r, err := n.transcribe(req)
|
||||
out = r
|
||||
return err
|
||||
}); err != nil {
|
||||
return pb.TranscriptResult{}, err
|
||||
}
|
||||
// transcribe returns a non-nil result whenever it returns a nil error, so
|
||||
// this cannot fire today. It is a guard rather than a comment because the
|
||||
// alternative to stating the invariant is a nil dereference in an RPC
|
||||
// handler if a later edit ever adds a success path that forgets to set it.
|
||||
if out == nil {
|
||||
return pb.TranscriptResult{}, status.Error(codes.Internal,
|
||||
"nemo-speech-cpp: transcription produced no result")
|
||||
}
|
||||
|
||||
// Assembled field by field rather than dereferenced: the RPC signature
|
||||
// returns the proto message by value, but the message embeds a mutex, so
|
||||
// copying the struct is a copylocks violation. Every backend in this tree
|
||||
// gets around it the same way, by only ever returning a composite literal.
|
||||
return pb.TranscriptResult{
|
||||
Text: out.Text,
|
||||
Segments: out.Segments,
|
||||
Language: out.Language,
|
||||
Duration: out.Duration,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// transcribe is AudioTranscription's body. The caller must hold engineMu.
|
||||
func (n *NemoSpeech) transcribe(req *pb.TranscriptRequest) (*pb.TranscriptResult, error) {
|
||||
if req.GetDst() == "" {
|
||||
return nil, status.Error(codes.InvalidArgument,
|
||||
"nemo-speech-cpp: TranscriptRequest.dst (audio path) is required")
|
||||
}
|
||||
|
||||
pcm, sampleRate, err := decodeAudioMono16k(req.GetDst())
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument,
|
||||
"nemo-speech-cpp: read audio: %v", err)
|
||||
}
|
||||
// Rejected here, before anything crosses the ABI, and not only inside
|
||||
// recognizeF32: a silent or truncated upload decodes to zero samples, and
|
||||
// there is no point building options and pinning strings for a request
|
||||
// that cannot produce a transcript. recognizeF32 keeps its own guard as a
|
||||
// precondition on the function.
|
||||
if len(pcm) == 0 {
|
||||
return nil, status.Error(codes.InvalidArgument, "nemo-speech-cpp: empty audio")
|
||||
}
|
||||
|
||||
// A per-request language wins over the model-level default; both may be
|
||||
// empty, which the runtime reads as auto/model default.
|
||||
language := req.GetLanguage()
|
||||
if language == "" {
|
||||
language = n.opts.languageCode
|
||||
}
|
||||
langP, freeLang := cstr(language)
|
||||
defer freeLang()
|
||||
|
||||
opts := ASRRecognitionOptionsDef()
|
||||
opts.LanguageCode = langP
|
||||
// Segments are built out of word offsets, so they are always asked for.
|
||||
opts.EnableWordTimeOffsets = true
|
||||
// Keyed on the recognizer owning a diar model, not on req.Diarize: asr.h
|
||||
// documents that a request asking for diarization from a recognizer created
|
||||
// without one fails with INVALID_ARGUMENT, and setting diar_model is already
|
||||
// the operator's opt-in.
|
||||
opts.EnableSpeakerDiarization = n.opts.diarModel != ""
|
||||
|
||||
result, err := recognizeF32(n.recognizer, &opts, pcm, sampleRate)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer ASRResultDestroy(result)
|
||||
|
||||
out := &pb.TranscriptResult{
|
||||
Text: ASRResultTranscript(result, 0),
|
||||
Segments: wordsToSegments(extractWords(result),
|
||||
wordsRequested(req.GetTimestampGranularities())),
|
||||
}
|
||||
// Multilingual models report what they decided the audio was; monolingual
|
||||
// ones report nothing, and an empty language is better than echoing back
|
||||
// whatever the caller guessed.
|
||||
if ASRResultLanguageCount(result, 0) > 0 {
|
||||
out.Language = ASRResultLanguageCode(result, 0, 0)
|
||||
}
|
||||
if sampleRate > 0 {
|
||||
out.Duration = float32(len(pcm)) / float32(sampleRate)
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
@@ -0,0 +1,526 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
"unsafe"
|
||||
|
||||
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
|
||||
"github.com/mudler/xlog"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
)
|
||||
|
||||
// streamChunkSamples is one push into a streaming session. At 16 kHz mono 1600
|
||||
// samples is 100 ms, short enough that the decoder is polled often enough to
|
||||
// see an endpoint promptly and short enough that a cancelled request stops
|
||||
// within one push.
|
||||
const streamChunkSamples = 1600
|
||||
|
||||
// The rates nemo_speech_asr_stream_push_f32 will resample from (asr.h). Outside
|
||||
// this range the runtime has nothing to do with the audio, and 0 is NOT
|
||||
// "unknown": it means "these samples are already at the model rate".
|
||||
const (
|
||||
minStreamSampleRate = 8000
|
||||
maxStreamSampleRate = 96000
|
||||
// TranscriptLiveConfig.sample_rate documents 0 as 16 kHz, which is a
|
||||
// different meaning from the C API's 0, so it is resolved before the push.
|
||||
defaultLiveSampleRate = 16000
|
||||
)
|
||||
|
||||
// streamResult is one result lifted out of C memory. Everything is copied
|
||||
// before nemo_speech_asr_result_destroy runs, so a streamResult outlives the
|
||||
// handle it came from.
|
||||
type streamResult struct {
|
||||
Text string
|
||||
Final bool
|
||||
Words []asrWord
|
||||
}
|
||||
|
||||
// asrSession is the streaming half of the ASR C API, narrowed to the four
|
||||
// entry points the two streaming RPCs use.
|
||||
//
|
||||
// It is an interface because there is no NeMo GGUF small enough to keep in the
|
||||
// tree, so the loops on top of it (chunking, the need-more-audio drain, the
|
||||
// live config/reset protocol) would otherwise have no test at all. The seam is
|
||||
// at the ABI, not at the model: a fake session scripts what the C API returns,
|
||||
// it does not pretend to transcribe anything.
|
||||
type asrSession interface {
|
||||
// push buffers audio. It does not decode; next drives that.
|
||||
push(pcm []float32, sampleRate int32) error
|
||||
// finish flushes the decoder tail. The end-of-stream final then comes back
|
||||
// from next.
|
||||
finish() error
|
||||
// next pulls one result. ok=false means the decoder needs more audio,
|
||||
// which is a pause in the stream and not an error or an end.
|
||||
next() (result streamResult, ok bool, err error)
|
||||
close()
|
||||
}
|
||||
|
||||
// sessionOpener creates a session for one language. n.openSession is the
|
||||
// C-backed implementation.
|
||||
type sessionOpener func(language string) (asrSession, error)
|
||||
|
||||
// cSession is the real asrSession, over one nemo_speech_asr_stream.
|
||||
type cSession struct {
|
||||
handle uintptr
|
||||
}
|
||||
|
||||
func (s *cSession) push(pcm []float32, sampleRate int32) error {
|
||||
// &pcm[0] panics on an empty slice, and an empty frame is ordinary input
|
||||
// from a live caller: it is a keepalive, not audio.
|
||||
if len(pcm) == 0 {
|
||||
return nil
|
||||
}
|
||||
if st := ASRStreamPushF32(s.handle, &pcm[0], uint64(len(pcm)), sampleRate); st != 0 {
|
||||
return statusErrorf(st, "nemo-speech-cpp: stream push: %s", ASRLastError())
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *cSession) finish() error {
|
||||
if st := ASRStreamFinish(s.handle); st != 0 {
|
||||
return statusErrorf(st, "nemo-speech-cpp: stream finish: %s", ASRLastError())
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *cSession) next() (streamResult, bool, error) {
|
||||
var handle uintptr
|
||||
if st := ASRStreamNext(s.handle, &handle); st != 0 {
|
||||
return streamResult{}, false, statusErrorf(st, "nemo-speech-cpp: stream next: %s", ASRLastError())
|
||||
}
|
||||
// OK with a NULL handle is the documented "need more audio". Reading it as
|
||||
// an error aborts every stream at the first gap; reading it as "keep
|
||||
// pulling" spins forever.
|
||||
if handle == 0 {
|
||||
return streamResult{}, false, nil
|
||||
}
|
||||
// Destroyed here rather than by the caller: everything below is copied out
|
||||
// of C memory into Go values, so nothing survives that would need it, and
|
||||
// a caller that returned early would otherwise leak the result.
|
||||
defer ASRResultDestroy(handle)
|
||||
|
||||
return streamResult{
|
||||
Text: ASRResultTranscript(handle, 0),
|
||||
Final: ASRResultIsFinal(handle),
|
||||
Words: extractWords(handle),
|
||||
}, true, nil
|
||||
}
|
||||
|
||||
func (s *cSession) close() { ASRStreamClose(s.handle) }
|
||||
|
||||
// openSession starts a streaming recognition on the loaded recognizer.
|
||||
//
|
||||
// The caller must hold engineMu.
|
||||
//
|
||||
// nemo_speech_asr_streaming_recognize copies the options (src/asr/c_api.cpp
|
||||
// to_options) and keeps no pointer into them, so the language buffer only has
|
||||
// to stay pinned across this call, exactly as in loadASR.
|
||||
func (n *NemoSpeech) openSession(language string) (asrSession, error) {
|
||||
// A per-request language wins over the model-level default; both may be
|
||||
// empty, which the runtime reads as auto/model default.
|
||||
if language == "" {
|
||||
language = n.opts.languageCode
|
||||
}
|
||||
langP, freeLang := cstr(language)
|
||||
defer freeLang()
|
||||
|
||||
opts := ASRRecognitionOptionsDef()
|
||||
opts.LanguageCode = langP
|
||||
// Segments and the live word list are built out of word offsets, so they
|
||||
// are always asked for.
|
||||
opts.EnableWordTimeOffsets = true
|
||||
// Keyed on the recognizer owning a diar model rather than on the request:
|
||||
// asr.h documents that asking a recognizer created without one for
|
||||
// diarization fails with INVALID_ARGUMENT.
|
||||
opts.EnableSpeakerDiarization = n.opts.diarModel != ""
|
||||
// interim_results is left off deliberately. The runtime emits interims from
|
||||
// next() regardless of it, and they are filtered here rather than
|
||||
// forwarded: see streamPCM's emit for why the wire contract cannot carry
|
||||
// them.
|
||||
|
||||
var handle uintptr
|
||||
// #nosec G103 -- opts is a local POD struct borrowed for this call only, and
|
||||
// its one uintptr member (LanguageCode) is the cstr allocation pinned by the
|
||||
// deferred freeLang above. to_options copies the struct, so nothing here
|
||||
// outlives the call.
|
||||
if st := ASRStreamingRecognize(n.recognizer, unsafe.Pointer(&opts), &handle); st != 0 {
|
||||
return nil, statusErrorf(st, "nemo-speech-cpp: streaming recognize: %s", ASRLastError())
|
||||
}
|
||||
xlog.Debug("nemo-speech-cpp: streaming session open", "language", language)
|
||||
return &cSession{handle: handle}, nil
|
||||
}
|
||||
|
||||
// chunkPCM slices pcm into fixed-size chunks, leaving the final chunk short
|
||||
// rather than padding it: silence padding would push audio the caller never
|
||||
// sent through the encoder and shift the tail word timings.
|
||||
func chunkPCM(pcm []float32, size int) [][]float32 {
|
||||
if len(pcm) == 0 {
|
||||
return nil
|
||||
}
|
||||
out := make([][]float32, 0, (len(pcm)+size-1)/size)
|
||||
for off := 0; off < len(pcm); off += size {
|
||||
out = append(out, pcm[off:min(off+size, len(pcm))])
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// drain pulls every result the session currently has, handing each to emit.
|
||||
// It returns when the session reports it needs more audio, which is the loop's
|
||||
// only terminating condition.
|
||||
func drain(sess asrSession, emit func(streamResult) error) error {
|
||||
for {
|
||||
r, ok, err := sess.next()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
if err := emit(r); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// streamPCM drives one whole clip through an open session, emitting each
|
||||
// finalized utterance as a delta and closing with the assembled result.
|
||||
//
|
||||
// Only finals become deltas, and the reason is the wire contract:
|
||||
// TranscriptStreamResponse.delta is newly-FINALIZED text that consumers
|
||||
// CONCATENATE (core/http/endpoints/openai/transcription.go, and the realtime
|
||||
// semantic-VAD path). An interim is the decoder's running hypothesis for the
|
||||
// utterance in flight, so forwarding "he", "hell", "hello", "Hello." would
|
||||
// assemble to "hehellhelloHello." rather than to the transcript. That the
|
||||
// runtime also postprocesses finals only (build_result_ in
|
||||
// src/asr/recognizer.cpp runs ITN and strip_formatting on the final, so it
|
||||
// rewrites rather than extends the interim) means there is no diffing trick
|
||||
// that would rescue them either.
|
||||
//
|
||||
// The cost is that the first delta of an utterance arrives at its endpoint
|
||||
// rather than mid-word.
|
||||
func streamPCM(ctx context.Context, sess asrSession, pcm []float32, sampleRate int32, wantWords bool, results chan<- *pb.TranscriptStreamResponse) error {
|
||||
if len(pcm) == 0 {
|
||||
return status.Error(codes.InvalidArgument, "nemo-speech-cpp: empty audio")
|
||||
}
|
||||
|
||||
var (
|
||||
full strings.Builder
|
||||
segments []*pb.TranscriptSegment
|
||||
// sawEndpoint records a final that arrived before the tail flush, i.e.
|
||||
// a real endpoint rather than the end of the file.
|
||||
sawEndpoint bool
|
||||
flushing bool
|
||||
tailText string
|
||||
)
|
||||
|
||||
emit := func(r streamResult) error {
|
||||
if !r.Final {
|
||||
return nil
|
||||
}
|
||||
if flushing {
|
||||
tailText += r.Text
|
||||
} else {
|
||||
sawEndpoint = true
|
||||
}
|
||||
if r.Text == "" {
|
||||
return nil
|
||||
}
|
||||
|
||||
// The separator is part of the delta, not added when assembling the
|
||||
// final text, so concatenating the deltas reproduces FinalResult.Text
|
||||
// exactly. Utterance transcripts carry no leading or trailing space of
|
||||
// their own (the runner clears its buffer at each endpoint).
|
||||
delta := r.Text
|
||||
if full.Len() > 0 {
|
||||
delta = " " + delta
|
||||
}
|
||||
full.WriteString(delta)
|
||||
|
||||
// One segment run per utterance, renumbered into the running sequence.
|
||||
// wordsToSegments splits a run further on a speaker change, so a
|
||||
// diarized utterance contributes one segment per turn.
|
||||
segs := wordsToSegments(r.Words, wantWords)
|
||||
if len(segs) == 0 {
|
||||
// Word offsets were requested but a decoder head may still return
|
||||
// none; a segment carrying just the text beats dropping it.
|
||||
segs = []*pb.TranscriptSegment{{Text: r.Text}}
|
||||
}
|
||||
for _, s := range segs {
|
||||
// #nosec G115 -- TranscriptSegment.Id is int32 on the wire, and
|
||||
// segments holds one entry per speaker run per finalized utterance of
|
||||
// a single request, which exhausts memory long before it reaches 2^31.
|
||||
s.Id = int32(len(segments))
|
||||
segments = append(segments, s)
|
||||
}
|
||||
|
||||
results <- &pb.TranscriptStreamResponse{Delta: delta}
|
||||
return nil
|
||||
}
|
||||
|
||||
for _, chunk := range chunkPCM(pcm, streamChunkSamples) {
|
||||
// The RPC body holds engineMu for the whole stream, so Free waits on
|
||||
// it. Without this check a client that disconnected mid-file would pin
|
||||
// the model against unload until the whole clip had been pushed.
|
||||
if err := ctx.Err(); err != nil {
|
||||
return status.Error(codes.Canceled, "nemo-speech-cpp: transcription cancelled")
|
||||
}
|
||||
if err := sess.push(chunk, sampleRate); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := drain(sess, emit); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
flushing = true
|
||||
if err := sess.finish(); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := drain(sess, emit); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
final := &pb.TranscriptResult{
|
||||
Text: full.String(),
|
||||
Segments: segments,
|
||||
// The tail flush returns whatever the decoder was still holding.
|
||||
// Nothing held back after at least one endpoint means the last
|
||||
// endpoint consumed the audio, which is what "the clip ended on an
|
||||
// utterance boundary" means here. Text coming back means it ended
|
||||
// mid-utterance.
|
||||
Eou: sawEndpoint && tailText == "",
|
||||
}
|
||||
if sampleRate > 0 {
|
||||
final.Duration = float32(len(pcm)) / float32(sampleRate)
|
||||
}
|
||||
results <- &pb.TranscriptStreamResponse{FinalResult: final}
|
||||
return nil
|
||||
}
|
||||
|
||||
// runLive drives one bidirectional live session. The protocol is the one
|
||||
// documented on the RPC in backend.proto: a Config first, a ready ack once the
|
||||
// session is open, deltas as utterances finalize, and a terminal result when
|
||||
// the caller closes its send side.
|
||||
//
|
||||
// There is no context here on purpose. The gRPC host closes `in` when the
|
||||
// stream context is cancelled (pkg/grpc/server.go's recv pump), so ranging
|
||||
// over it is what stops this loop, and that is also what releases engineMu for
|
||||
// a waiting Free.
|
||||
func runLive(open sessionOpener, in <-chan *pb.TranscriptLiveRequest, out chan<- *pb.TranscriptLiveResponse) error {
|
||||
first, ok := <-in
|
||||
if !ok {
|
||||
// The caller closed without sending anything. Nothing was opened, so
|
||||
// there is nothing to report.
|
||||
return nil
|
||||
}
|
||||
cfg := first.GetConfig()
|
||||
if cfg == nil {
|
||||
return status.Error(codes.InvalidArgument,
|
||||
"nemo-speech-cpp: the first live message must carry a config")
|
||||
}
|
||||
rate, err := liveSampleRate(cfg)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
sess, err := open(cfg.GetLanguage())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
// A mid-stream config replaces sess, so this closes whichever session is
|
||||
// current when the RPC unwinds.
|
||||
defer func() { sess.close() }()
|
||||
|
||||
// Callers block on the first Recv waiting for this and degrade to
|
||||
// non-live transcription when it does not arrive, so it goes out before
|
||||
// any audio is read.
|
||||
out <- &pb.TranscriptLiveResponse{Ready: true}
|
||||
|
||||
var (
|
||||
full strings.Builder
|
||||
flushing bool
|
||||
)
|
||||
emit := func(r streamResult) error {
|
||||
// Finals only, for the same reason as streamPCM: an interim is a
|
||||
// hypothesis the final rewrites, and delta is newly-finalized text.
|
||||
if !r.Final || (r.Text == "" && len(r.Words) == 0) {
|
||||
return nil
|
||||
}
|
||||
|
||||
// The separator goes INTO the delta, exactly as in streamPCM, because
|
||||
// the live consumer is the one that actually concatenates: the realtime
|
||||
// semantic-VAD path joins the accumulated deltas with the empty string
|
||||
// and only clears them at a turn reset, never at an endpoint. Adding
|
||||
// the space when assembling the terminal text instead would make the
|
||||
// running caption read "one.two." while the committed transcript read
|
||||
// "one. two.".
|
||||
delta := r.Text
|
||||
if delta != "" && full.Len() > 0 {
|
||||
delta = " " + delta
|
||||
}
|
||||
full.WriteString(delta)
|
||||
|
||||
out <- &pb.TranscriptLiveResponse{
|
||||
Delta: delta,
|
||||
// A final that arrives while audio is still coming IS the model's
|
||||
// endpoint: the decoder resets its utterance there and the next one
|
||||
// starts fresh, which is the turn boundary the realtime detector
|
||||
// waits on. The final that comes back from the tail flush is the
|
||||
// end of the STREAM, not a user yielding a turn, so it carries no
|
||||
// eou even though the send side has already closed.
|
||||
Eou: !flushing,
|
||||
Words: wordsToProto(r.Words),
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
for req := range in {
|
||||
switch payload := req.GetPayload().(type) {
|
||||
case *pb.TranscriptLiveRequest_Config:
|
||||
// A rate cannot change inside a stream (asr.h) and the decoder
|
||||
// keeps utterance state, so a reconfigure has to be a fresh
|
||||
// session rather than a reconfigured one.
|
||||
newRate, err := liveSampleRate(payload.Config)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
// Opened before the old one is closed so a failure here leaves a
|
||||
// live session for the deferred close, not a dangling handle.
|
||||
next, err := open(payload.Config.GetLanguage())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
sess.close()
|
||||
sess, rate = next, newRate
|
||||
full.Reset()
|
||||
case *pb.TranscriptLiveRequest_Audio:
|
||||
pcm := payload.Audio.GetPcm()
|
||||
if len(pcm) == 0 {
|
||||
continue
|
||||
}
|
||||
if err := sess.push(pcm, rate); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := drain(sess, emit); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Send side closed: flush the tail and emit the terminal result. Like the
|
||||
// other backends' live path this carries Text only; per-utterance segments
|
||||
// and the duration are the file path's concern.
|
||||
flushing = true
|
||||
if err := sess.finish(); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := drain(sess, emit); err != nil {
|
||||
return err
|
||||
}
|
||||
// Not trimmed: the terminal text is the verbatim concatenation of the
|
||||
// deltas, which is the invariant the concatenating consumers rely on. The
|
||||
// first delta never carries the separator, so there is no leading space to
|
||||
// trim off in the first place.
|
||||
out <- &pb.TranscriptLiveResponse{
|
||||
FinalResult: &pb.TranscriptResult{Text: full.String()},
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// liveSampleRate resolves TranscriptLiveConfig.sample_rate to the rate the C
|
||||
// API is given. The proto's 0 means 16 kHz; the C API's 0 means "already at the
|
||||
// model rate", so the two cannot be forwarded to each other.
|
||||
func liveSampleRate(cfg *pb.TranscriptLiveConfig) (int32, error) {
|
||||
rate := cfg.GetSampleRate()
|
||||
if rate == 0 {
|
||||
return defaultLiveSampleRate, nil
|
||||
}
|
||||
if rate < minStreamSampleRate || rate > maxStreamSampleRate {
|
||||
return 0, status.Errorf(codes.InvalidArgument,
|
||||
"nemo-speech-cpp: unsupported live sample_rate %d (accepted: 0 or %d-%d Hz)",
|
||||
rate, minStreamSampleRate, maxStreamSampleRate)
|
||||
}
|
||||
return rate, nil
|
||||
}
|
||||
|
||||
// wordsToProto converts decoded words to the wire form. TranscriptWord.start
|
||||
// and .end are int64 nanoseconds; the runtime reports milliseconds.
|
||||
func wordsToProto(words []asrWord) []*pb.TranscriptWord {
|
||||
if len(words) == 0 {
|
||||
return nil
|
||||
}
|
||||
out := make([]*pb.TranscriptWord, len(words))
|
||||
for i, w := range words {
|
||||
out[i] = &pb.TranscriptWord{
|
||||
Text: w.Text,
|
||||
Start: msToNanos(w.Start),
|
||||
End: msToNanos(w.End),
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// AudioTranscriptionStream decodes the audio at req.Dst through the streaming
|
||||
// recognizer, emitting each finalized utterance as it lands.
|
||||
//
|
||||
// The body runs inside withEngine for the reason documented on withEngine, and
|
||||
// that holds engineMu for the whole stream: Free waits rather than destroying
|
||||
// the recognizer under a half-finished stream. streamPCM honours ctx so the
|
||||
// wait is bounded by the client's disconnect rather than by its silence.
|
||||
func (n *NemoSpeech) AudioTranscriptionStream(ctx context.Context, req *pb.TranscriptRequest, results chan *pb.TranscriptStreamResponse) error {
|
||||
// The host ranges over this channel and only returns once it closes, so
|
||||
// every path out of here, rejection included, has to close it.
|
||||
defer close(results)
|
||||
|
||||
return n.withEngine(familyASR, func() error {
|
||||
return n.transcribeStream(ctx, req, results)
|
||||
})
|
||||
}
|
||||
|
||||
// transcribeStream is AudioTranscriptionStream's body. The caller must hold
|
||||
// engineMu.
|
||||
func (n *NemoSpeech) transcribeStream(ctx context.Context, req *pb.TranscriptRequest, results chan<- *pb.TranscriptStreamResponse) error {
|
||||
if req.GetDst() == "" {
|
||||
return status.Error(codes.InvalidArgument,
|
||||
"nemo-speech-cpp: TranscriptRequest.dst (audio path) is required")
|
||||
}
|
||||
// Checked before the decode so a client that has already gone away does
|
||||
// not pay for an ffmpeg run, and so a cancellation is never reported as a
|
||||
// broken file.
|
||||
if err := ctx.Err(); err != nil {
|
||||
return status.Error(codes.Canceled, "nemo-speech-cpp: transcription cancelled")
|
||||
}
|
||||
|
||||
pcm, sampleRate, err := decodeAudioMono16k(req.GetDst())
|
||||
if err != nil {
|
||||
return status.Errorf(codes.InvalidArgument, "nemo-speech-cpp: read audio: %v", err)
|
||||
}
|
||||
// Before the session is opened, for the same reason as the offline path:
|
||||
// there is no transcript to be had from zero samples, and opening a stream
|
||||
// only to close it again asks the runtime to allocate decoder state for
|
||||
// nothing.
|
||||
if len(pcm) == 0 {
|
||||
return status.Error(codes.InvalidArgument, "nemo-speech-cpp: empty audio")
|
||||
}
|
||||
|
||||
sess, err := n.openSession(req.GetLanguage())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer sess.close()
|
||||
|
||||
return streamPCM(ctx, sess, pcm, sampleRate,
|
||||
wordsRequested(req.GetTimestampGranularities()), results)
|
||||
}
|
||||
|
||||
// AudioTranscriptionLive serves the bidirectional live RPC over one streaming
|
||||
// session. See runLive for the protocol and withEngine for the locking.
|
||||
func (n *NemoSpeech) AudioTranscriptionLive(in <-chan *pb.TranscriptLiveRequest, out chan<- *pb.TranscriptLiveResponse) error {
|
||||
defer close(out)
|
||||
|
||||
return n.withEngine(familyASR, func() error {
|
||||
return runLive(n.openSession, in, out)
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,699 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"path/filepath"
|
||||
"time"
|
||||
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
|
||||
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
|
||||
)
|
||||
|
||||
// fakeSession is a scripted asrSession. It stands in for the streaming C API,
|
||||
// not for a model: no NeMo GGUF is small enough to keep in the tree, and the
|
||||
// need-more-audio drain is the easiest thing in this file to get subtly wrong
|
||||
// (a mishandled NULL either spins forever or drops every result).
|
||||
//
|
||||
// script is one batch of results per drain. next() hands back the current
|
||||
// batch one result at a time and then reports "need more audio" exactly once,
|
||||
// which advances to the next batch. That is precisely the C contract:
|
||||
// nemo_speech_asr_stream_next returns OK with a NULL handle when the decoder
|
||||
// has consumed the buffered audio, and the loop must resume after the next
|
||||
// push rather than treat it as the end of the stream.
|
||||
type fakeSession struct {
|
||||
script [][]streamResult
|
||||
batch int
|
||||
pos int
|
||||
|
||||
pushed [][]float32
|
||||
rates []int32
|
||||
finished int
|
||||
closed int
|
||||
|
||||
pushErr error
|
||||
finishErr error
|
||||
nextErr error
|
||||
}
|
||||
|
||||
func (f *fakeSession) push(pcm []float32, sampleRate int32) error {
|
||||
if f.pushErr != nil {
|
||||
return f.pushErr
|
||||
}
|
||||
f.pushed = append(f.pushed, pcm)
|
||||
f.rates = append(f.rates, sampleRate)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (f *fakeSession) finish() error {
|
||||
if f.finishErr != nil {
|
||||
return f.finishErr
|
||||
}
|
||||
f.finished++
|
||||
return nil
|
||||
}
|
||||
|
||||
func (f *fakeSession) next() (streamResult, bool, error) {
|
||||
if f.nextErr != nil {
|
||||
return streamResult{}, false, f.nextErr
|
||||
}
|
||||
if f.batch >= len(f.script) {
|
||||
return streamResult{}, false, nil
|
||||
}
|
||||
if f.pos >= len(f.script[f.batch]) {
|
||||
f.batch++
|
||||
f.pos = 0
|
||||
return streamResult{}, false, nil
|
||||
}
|
||||
r := f.script[f.batch][f.pos]
|
||||
f.pos++
|
||||
return r, true, nil
|
||||
}
|
||||
|
||||
func (f *fakeSession) close() { f.closed++ }
|
||||
|
||||
// samples returns the flat concatenation of everything pushed, so a spec can
|
||||
// assert the whole clip reached the engine without caring how it was sliced.
|
||||
func (f *fakeSession) samples() []float32 {
|
||||
var out []float32
|
||||
for _, c := range f.pushed {
|
||||
out = append(out, c...)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// collect drains a response channel into a slice. The channels are unbuffered
|
||||
// in the specs on purpose: a producer that stops honouring cancellation would
|
||||
// otherwise fill a buffer and look healthy.
|
||||
func collect[T any](ch chan T) chan []T {
|
||||
done := make(chan []T, 1)
|
||||
go func() {
|
||||
var got []T
|
||||
for v := range ch {
|
||||
got = append(got, v)
|
||||
}
|
||||
done <- got
|
||||
}()
|
||||
return done
|
||||
}
|
||||
|
||||
var _ = Describe("chunkPCM", func() {
|
||||
It("splits into equal chunks when evenly divisible", func() {
|
||||
chunks := chunkPCM(make([]float32, 400), 100)
|
||||
Expect(chunks).To(HaveLen(4))
|
||||
for _, c := range chunks {
|
||||
Expect(c).To(HaveLen(100))
|
||||
}
|
||||
})
|
||||
|
||||
// Padding the tail with silence would push phantom audio through the
|
||||
// encoder and shift the tail word timings, so the final chunk stays short.
|
||||
It("makes the final chunk short rather than padding it", func() {
|
||||
chunks := chunkPCM(make([]float32, 250), 100)
|
||||
Expect(chunks).To(HaveLen(3))
|
||||
Expect(chunks[2]).To(HaveLen(50))
|
||||
})
|
||||
|
||||
It("returns one chunk when the input is shorter than the chunk size", func() {
|
||||
chunks := chunkPCM(make([]float32, 10), 100)
|
||||
Expect(chunks).To(HaveLen(1))
|
||||
Expect(chunks[0]).To(HaveLen(10))
|
||||
})
|
||||
|
||||
It("returns nothing for empty input", func() {
|
||||
Expect(chunkPCM(nil, 100)).To(BeEmpty())
|
||||
Expect(chunkPCM([]float32{}, 100)).To(BeEmpty())
|
||||
})
|
||||
|
||||
// Every spec above works on all-zero audio, so none of them can tell a
|
||||
// correct slicing from one that reorders or repeats windows. Audio fed out
|
||||
// of order still decodes, it just decodes to nonsense.
|
||||
It("preserves sample order across the chunk boundaries", func() {
|
||||
pcm := []float32{1, 2, 3, 4, 5}
|
||||
chunks := chunkPCM(pcm, 2)
|
||||
Expect(chunks).To(HaveLen(3))
|
||||
Expect(chunks[0]).To(Equal([]float32{1, 2}))
|
||||
Expect(chunks[1]).To(Equal([]float32{3, 4}))
|
||||
Expect(chunks[2]).To(Equal([]float32{5}))
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("drain", func() {
|
||||
It("emits every result in a batch and stops on need-more-audio", func() {
|
||||
sess := &fakeSession{script: [][]streamResult{
|
||||
{{Text: "a"}, {Text: "b", Final: true}},
|
||||
{{Text: "c"}},
|
||||
}}
|
||||
var got []string
|
||||
Expect(drain(sess, func(r streamResult) error {
|
||||
got = append(got, r.Text)
|
||||
return nil
|
||||
})).To(Succeed())
|
||||
Expect(got).To(Equal([]string{"a", "b"}))
|
||||
})
|
||||
|
||||
// The NULL handle is a pause, not an end: the next drain, after more audio
|
||||
// has been pushed, must pick the stream back up.
|
||||
It("resumes on the next drain after a need-more-audio pause", func() {
|
||||
sess := &fakeSession{script: [][]streamResult{{{Text: "a"}}, {{Text: "b"}}}}
|
||||
var got []string
|
||||
emit := func(r streamResult) error { got = append(got, r.Text); return nil }
|
||||
Expect(drain(sess, emit)).To(Succeed())
|
||||
Expect(drain(sess, emit)).To(Succeed())
|
||||
Expect(got).To(Equal([]string{"a", "b"}))
|
||||
})
|
||||
|
||||
It("returns nothing and no error for a stream with no results ready", func() {
|
||||
var got []string
|
||||
Expect(drain(&fakeSession{}, func(r streamResult) error {
|
||||
got = append(got, r.Text)
|
||||
return nil
|
||||
})).To(Succeed())
|
||||
Expect(got).To(BeEmpty())
|
||||
})
|
||||
|
||||
It("propagates a failure from the runtime", func() {
|
||||
sess := &fakeSession{nextErr: errors.New("boom")}
|
||||
Expect(drain(sess, func(streamResult) error { return nil })).To(MatchError(ContainSubstring("boom")))
|
||||
})
|
||||
|
||||
It("stops pulling once emit fails", func() {
|
||||
sess := &fakeSession{script: [][]streamResult{{{Text: "a"}, {Text: "b"}}}}
|
||||
Expect(drain(sess, func(streamResult) error {
|
||||
return errors.New("send failed")
|
||||
})).To(MatchError(ContainSubstring("send failed")))
|
||||
Expect(sess.pos).To(Equal(1))
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("streamPCM", func() {
|
||||
streamWords := func(ctx context.Context, sess asrSession, pcm []float32, rate int32, wantWords bool) ([]*pb.TranscriptStreamResponse, error) {
|
||||
GinkgoHelper()
|
||||
results := make(chan *pb.TranscriptStreamResponse)
|
||||
done := collect(results)
|
||||
err := streamPCM(ctx, sess, pcm, rate, wantWords, results)
|
||||
close(results)
|
||||
return <-done, err
|
||||
}
|
||||
stream := func(ctx context.Context, sess asrSession, pcm []float32, rate int32) ([]*pb.TranscriptStreamResponse, error) {
|
||||
GinkgoHelper()
|
||||
return streamWords(ctx, sess, pcm, rate, false)
|
||||
}
|
||||
|
||||
It("pushes the whole clip in chunks at the clip's own sample rate", func() {
|
||||
sess := &fakeSession{}
|
||||
pcm := make([]float32, streamChunkSamples*2+7)
|
||||
_, err := stream(context.Background(), sess, pcm, 16000)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(sess.pushed).To(HaveLen(3))
|
||||
Expect(sess.samples()).To(HaveLen(len(pcm)))
|
||||
for _, r := range sess.rates {
|
||||
Expect(r).To(Equal(int32(16000)))
|
||||
}
|
||||
})
|
||||
|
||||
It("finishes the stream once, after the last chunk", func() {
|
||||
sess := &fakeSession{}
|
||||
_, err := stream(context.Background(), sess, make([]float32, 10), 16000)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(sess.finished).To(Equal(1))
|
||||
})
|
||||
|
||||
// Interims are the decoder's running hypothesis for the utterance in
|
||||
// flight. The wire contract is that delta is newly FINALIZED text and that
|
||||
// concatenating the deltas reproduces the transcript, so forwarding an
|
||||
// interim would duplicate every word it later re-sends inside the final.
|
||||
It("emits a delta per final and nothing for interims", func() {
|
||||
sess := &fakeSession{script: [][]streamResult{{
|
||||
{Text: "hel"},
|
||||
{Text: "hello"},
|
||||
{Text: "Hello.", Final: true},
|
||||
}}}
|
||||
got, err := stream(context.Background(), sess, make([]float32, 10), 16000)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
|
||||
var deltas []string
|
||||
for _, r := range got {
|
||||
if r.GetDelta() != "" {
|
||||
deltas = append(deltas, r.GetDelta())
|
||||
}
|
||||
}
|
||||
Expect(deltas).To(Equal([]string{"Hello."}))
|
||||
})
|
||||
|
||||
It("reproduces the final transcript by concatenating the deltas", func() {
|
||||
sess := &fakeSession{script: [][]streamResult{{
|
||||
{Text: "One.", Final: true},
|
||||
{Text: "Two.", Final: true},
|
||||
}}}
|
||||
got, err := stream(context.Background(), sess, make([]float32, 10), 16000)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
|
||||
var joined string
|
||||
var final *pb.TranscriptResult
|
||||
for _, r := range got {
|
||||
joined += r.GetDelta()
|
||||
if r.GetFinalResult() != nil {
|
||||
final = r.GetFinalResult()
|
||||
}
|
||||
}
|
||||
Expect(final).ToNot(BeNil())
|
||||
Expect(final.GetText()).To(Equal("One. Two."))
|
||||
Expect(joined).To(Equal(final.GetText()))
|
||||
})
|
||||
|
||||
It("sends the terminal final result last and only once", func() {
|
||||
sess := &fakeSession{script: [][]streamResult{{{Text: "hi", Final: true}}}}
|
||||
got, err := stream(context.Background(), sess, make([]float32, 10), 16000)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(got).ToNot(BeEmpty())
|
||||
|
||||
var finals int
|
||||
for _, r := range got {
|
||||
if r.GetFinalResult() != nil {
|
||||
finals++
|
||||
}
|
||||
}
|
||||
Expect(finals).To(Equal(1))
|
||||
Expect(got[len(got)-1].GetFinalResult()).ToNot(BeNil())
|
||||
})
|
||||
|
||||
It("reports the clip duration in seconds", func() {
|
||||
sess := &fakeSession{}
|
||||
got, err := stream(context.Background(), sess, make([]float32, 8000), 16000)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(got[len(got)-1].GetFinalResult().GetDuration()).To(BeNumerically("~", 0.5, 1e-6))
|
||||
})
|
||||
|
||||
It("builds per-utterance segments with nanosecond timestamps", func() {
|
||||
sess := &fakeSession{script: [][]streamResult{{
|
||||
{Text: "one", Final: true, Words: []asrWord{{Text: "one", Start: 0, End: 500}}},
|
||||
{Text: "two", Final: true, Words: []asrWord{{Text: "two", Start: 900, End: 1400}}},
|
||||
}}}
|
||||
got, err := stream(context.Background(), sess, make([]float32, 10), 16000)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
|
||||
segs := got[len(got)-1].GetFinalResult().GetSegments()
|
||||
Expect(segs).To(HaveLen(2))
|
||||
Expect(segs[0].GetId()).To(Equal(int32(0)))
|
||||
Expect(segs[1].GetId()).To(Equal(int32(1)))
|
||||
Expect(time.Duration(segs[1].GetStart())).To(Equal(900 * time.Millisecond))
|
||||
Expect(time.Duration(segs[1].GetEnd())).To(Equal(1400 * time.Millisecond))
|
||||
})
|
||||
|
||||
// core/backend/transcript.go builds the response's word list out of
|
||||
// TranscriptSegment.Words, so leaving it unset makes
|
||||
// timestamp_granularities: ["word"] come back empty.
|
||||
It("attaches the word timings only when they were asked for", func() {
|
||||
script := func() [][]streamResult {
|
||||
return [][]streamResult{{{Text: "one", Final: true,
|
||||
Words: []asrWord{{Text: "one", Start: 100, End: 500}}}}}
|
||||
}
|
||||
|
||||
got, err := streamWords(context.Background(), &fakeSession{script: script()}, make([]float32, 10), 16000, true)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
segs := got[len(got)-1].GetFinalResult().GetSegments()
|
||||
Expect(segs[0].GetWords()).To(HaveLen(1))
|
||||
Expect(segs[0].GetWords()[0].GetText()).To(Equal("one"))
|
||||
Expect(time.Duration(segs[0].GetWords()[0].GetStart())).To(Equal(100 * time.Millisecond))
|
||||
|
||||
got, err = streamWords(context.Background(), &fakeSession{script: script()}, make([]float32, 10), 16000, false)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
segs = got[len(got)-1].GetFinalResult().GetSegments()
|
||||
Expect(segs[0].GetText()).To(Equal("one"))
|
||||
Expect(segs[0].GetWords()).To(BeEmpty())
|
||||
})
|
||||
|
||||
// The flush that nemo_speech_asr_stream_finish triggers returns whatever the
|
||||
// decoder was still holding. Nothing held back means the last endpoint
|
||||
// consumed the audio, which is exactly "the clip ended on an utterance
|
||||
// boundary"; text coming back means it ended mid-utterance.
|
||||
It("marks eou when the tail flush had nothing left to emit", func() {
|
||||
sess := &fakeSession{script: [][]streamResult{
|
||||
{{Text: "done.", Final: true}},
|
||||
{{Text: "", Final: true}},
|
||||
}}
|
||||
got, err := stream(context.Background(), sess, make([]float32, 10), 16000)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(got[len(got)-1].GetFinalResult().GetEou()).To(BeTrue())
|
||||
})
|
||||
|
||||
It("does not mark eou when the tail flush produced text", func() {
|
||||
sess := &fakeSession{script: [][]streamResult{
|
||||
{{Text: "done.", Final: true}},
|
||||
{{Text: "and more", Final: true}},
|
||||
}}
|
||||
got, err := stream(context.Background(), sess, make([]float32, 10), 16000)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(got[len(got)-1].GetFinalResult().GetEou()).To(BeFalse())
|
||||
})
|
||||
|
||||
// The RPC body runs inside withEngine, so it holds the engine mutex for the
|
||||
// whole stream and Free waits on it. A loop that ignored cancellation would
|
||||
// pin the model against unload for as long as a disconnected client's audio
|
||||
// takes to push.
|
||||
It("stops promptly when the request context is cancelled", func() {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
sess := &fakeSession{}
|
||||
_, err := stream(ctx, sess, make([]float32, streamChunkSamples*4), 16000)
|
||||
Expect(status.Code(err)).To(Equal(codes.Canceled))
|
||||
Expect(sess.pushed).To(BeEmpty())
|
||||
})
|
||||
|
||||
It("reports a push failure", func() {
|
||||
sess := &fakeSession{pushErr: errors.New("push blew up")}
|
||||
_, err := stream(context.Background(), sess, make([]float32, 10), 16000)
|
||||
Expect(err).To(MatchError(ContainSubstring("push blew up")))
|
||||
})
|
||||
|
||||
It("reports a finish failure", func() {
|
||||
sess := &fakeSession{finishErr: errors.New("finish blew up")}
|
||||
_, err := stream(context.Background(), sess, make([]float32, 10), 16000)
|
||||
Expect(err).To(MatchError(ContainSubstring("finish blew up")))
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("runLive", func() {
|
||||
// live drives runLive against a fake opener and returns everything the RPC
|
||||
// wrote plus the sessions it opened.
|
||||
live := func(reqs []*pb.TranscriptLiveRequest, script ...[][]streamResult) ([]*pb.TranscriptLiveResponse, []*fakeSession, error) {
|
||||
GinkgoHelper()
|
||||
var opened []*fakeSession
|
||||
open := func(language string) (asrSession, error) {
|
||||
s := &fakeSession{}
|
||||
if len(opened) < len(script) {
|
||||
s.script = script[len(opened)]
|
||||
}
|
||||
opened = append(opened, s)
|
||||
return s, nil
|
||||
}
|
||||
|
||||
in := make(chan *pb.TranscriptLiveRequest)
|
||||
out := make(chan *pb.TranscriptLiveResponse)
|
||||
done := collect(out)
|
||||
go func() {
|
||||
defer close(in)
|
||||
for _, r := range reqs {
|
||||
in <- r
|
||||
}
|
||||
}()
|
||||
err := runLive(open, in, out)
|
||||
close(out)
|
||||
return <-done, opened, err
|
||||
}
|
||||
|
||||
cfg := func(rate int32) *pb.TranscriptLiveRequest {
|
||||
return &pb.TranscriptLiveRequest{Payload: &pb.TranscriptLiveRequest_Config{
|
||||
Config: &pb.TranscriptLiveConfig{SampleRate: rate},
|
||||
}}
|
||||
}
|
||||
audio := func(pcm ...float32) *pb.TranscriptLiveRequest {
|
||||
return &pb.TranscriptLiveRequest{Payload: &pb.TranscriptLiveRequest_Audio{
|
||||
Audio: &pb.TranscriptLiveAudio{Pcm: pcm},
|
||||
}}
|
||||
}
|
||||
|
||||
It("requires the first message to carry a config", func() {
|
||||
_, opened, err := live([]*pb.TranscriptLiveRequest{audio(1, 2, 3)})
|
||||
Expect(status.Code(err)).To(Equal(codes.InvalidArgument))
|
||||
Expect(opened).To(BeEmpty())
|
||||
})
|
||||
|
||||
It("returns without error when the caller closes without sending anything", func() {
|
||||
got, opened, err := live(nil)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(got).To(BeEmpty())
|
||||
Expect(opened).To(BeEmpty())
|
||||
})
|
||||
|
||||
// Callers block on the first Recv waiting for this ack, and degrade to
|
||||
// non-live transcription when it does not arrive.
|
||||
It("acknowledges a successful open before any transcript", func() {
|
||||
got, _, err := live([]*pb.TranscriptLiveRequest{cfg(0)})
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(got).ToNot(BeEmpty())
|
||||
Expect(got[0].GetReady()).To(BeTrue())
|
||||
})
|
||||
|
||||
// The proto documents 0 as "16 kHz". The C API reads 0 as "these samples
|
||||
// are already at the model rate" and skips resampling, so forwarding the
|
||||
// zero through would silently mean something else.
|
||||
It("resolves the default sample rate to 16 kHz before pushing", func() {
|
||||
_, opened, err := live([]*pb.TranscriptLiveRequest{cfg(0), audio(1, 2, 3)})
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(opened).To(HaveLen(1))
|
||||
Expect(opened[0].rates).To(Equal([]int32{16000}))
|
||||
})
|
||||
|
||||
It("pushes at the configured sample rate", func() {
|
||||
_, opened, err := live([]*pb.TranscriptLiveRequest{cfg(8000), audio(1, 2, 3)})
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(opened[0].rates).To(Equal([]int32{8000}))
|
||||
Expect(opened[0].samples()).To(Equal([]float32{1, 2, 3}))
|
||||
})
|
||||
|
||||
It("rejects a sample rate the runtime cannot resample", func() {
|
||||
_, opened, err := live([]*pb.TranscriptLiveRequest{cfg(4000)})
|
||||
Expect(status.Code(err)).To(Equal(codes.InvalidArgument))
|
||||
Expect(opened).To(BeEmpty())
|
||||
})
|
||||
|
||||
It("ignores an empty audio frame instead of pushing it", func() {
|
||||
_, opened, err := live([]*pb.TranscriptLiveRequest{cfg(0), audio()})
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(opened[0].pushed).To(BeEmpty())
|
||||
})
|
||||
|
||||
It("streams a delta with its words and marks the utterance boundary", func() {
|
||||
got, _, err := live(
|
||||
[]*pb.TranscriptLiveRequest{cfg(0), audio(1)},
|
||||
[][]streamResult{{
|
||||
{Text: "partial"},
|
||||
{Text: "Hello there.", Final: true, Words: []asrWord{
|
||||
{Text: "Hello", Start: 100, End: 400},
|
||||
{Text: "there", Start: 400, End: 900},
|
||||
}},
|
||||
}},
|
||||
)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
|
||||
var deltas []*pb.TranscriptLiveResponse
|
||||
for _, r := range got {
|
||||
if r.GetDelta() != "" {
|
||||
deltas = append(deltas, r)
|
||||
}
|
||||
}
|
||||
Expect(deltas).To(HaveLen(1))
|
||||
Expect(deltas[0].GetDelta()).To(Equal("Hello there."))
|
||||
Expect(deltas[0].GetEou()).To(BeTrue())
|
||||
Expect(deltas[0].GetWords()).To(HaveLen(2))
|
||||
Expect(time.Duration(deltas[0].GetWords()[1].GetStart())).To(Equal(400 * time.Millisecond))
|
||||
Expect(time.Duration(deltas[0].GetWords()[1].GetEnd())).To(Equal(900 * time.Millisecond))
|
||||
})
|
||||
|
||||
It("finishes and closes the session when the caller closes the send side", func() {
|
||||
got, opened, err := live(
|
||||
[]*pb.TranscriptLiveRequest{cfg(0), audio(1)},
|
||||
[][]streamResult{{{Text: "one.", Final: true}}, {{Text: "two.", Final: true}}},
|
||||
)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(opened[0].finished).To(Equal(1))
|
||||
Expect(opened[0].closed).To(Equal(1))
|
||||
Expect(got[len(got)-1].GetFinalResult()).ToNot(BeNil())
|
||||
Expect(got[len(got)-1].GetFinalResult().GetText()).To(Equal("one. two."))
|
||||
})
|
||||
|
||||
// The live path is the one with a consumer that really concatenates: the
|
||||
// realtime semantic-VAD path joins the accumulated deltas with the empty
|
||||
// string and clears them only at a turn reset, never at an utterance
|
||||
// boundary. A separator added when assembling the terminal text instead of
|
||||
// inside the delta makes the running caption read "one.two." while the
|
||||
// committed transcript reads "one. two.".
|
||||
It("reproduces the final transcript by concatenating the deltas", func() {
|
||||
got, _, err := live(
|
||||
[]*pb.TranscriptLiveRequest{cfg(0), audio(1)},
|
||||
[][]streamResult{{{Text: "one.", Final: true}}, {{Text: "two.", Final: true}}},
|
||||
)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
|
||||
var joined string
|
||||
var final *pb.TranscriptResult
|
||||
for _, r := range got {
|
||||
joined += r.GetDelta()
|
||||
if r.GetFinalResult() != nil {
|
||||
final = r.GetFinalResult()
|
||||
}
|
||||
}
|
||||
Expect(final).ToNot(BeNil())
|
||||
Expect(final.GetText()).To(Equal("one. two."))
|
||||
Expect(joined).To(Equal(final.GetText()))
|
||||
})
|
||||
|
||||
// Eou is the model's endpoint, which is a user yielding the turn. The final
|
||||
// that comes back from the tail flush is the end of the stream: the send
|
||||
// side has already closed, so reporting a turn boundary there tells the
|
||||
// turn detector something that did not happen.
|
||||
It("marks the endpoint finals but not the tail flush", func() {
|
||||
got, _, err := live(
|
||||
[]*pb.TranscriptLiveRequest{cfg(0), audio(1)},
|
||||
[][]streamResult{{{Text: "one.", Final: true}}, {{Text: "two.", Final: true}}},
|
||||
)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
|
||||
var eous []bool
|
||||
for _, r := range got {
|
||||
if r.GetDelta() != "" {
|
||||
eous = append(eous, r.GetEou())
|
||||
}
|
||||
}
|
||||
Expect(eous).To(Equal([]bool{true, false}))
|
||||
})
|
||||
|
||||
// A rate cannot change inside a stream and the decoder keeps no state
|
||||
// across a reset, so a second config has to be a fresh session, not a
|
||||
// reconfigured one.
|
||||
It("opens a fresh session on a mid-stream config and drops the old transcript", func() {
|
||||
got, opened, err := live(
|
||||
[]*pb.TranscriptLiveRequest{cfg(0), audio(1), cfg(0), audio(2)},
|
||||
[][]streamResult{{{Text: "dropped.", Final: true}}},
|
||||
[][]streamResult{{{Text: "kept.", Final: true}}},
|
||||
)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(opened).To(HaveLen(2))
|
||||
Expect(opened[0].closed).To(Equal(1))
|
||||
Expect(got[len(got)-1].GetFinalResult().GetText()).To(Equal("kept."))
|
||||
})
|
||||
|
||||
It("reports a push failure and still closes the session", func() {
|
||||
var opened []*fakeSession
|
||||
open := func(string) (asrSession, error) {
|
||||
s := &fakeSession{pushErr: errors.New("push blew up")}
|
||||
opened = append(opened, s)
|
||||
return s, nil
|
||||
}
|
||||
in := make(chan *pb.TranscriptLiveRequest, 2)
|
||||
in <- cfg(0)
|
||||
in <- audio(1, 2)
|
||||
close(in)
|
||||
out := make(chan *pb.TranscriptLiveResponse, 8)
|
||||
|
||||
err := runLive(open, in, out)
|
||||
Expect(err).To(MatchError(ContainSubstring("push blew up")))
|
||||
Expect(opened[0].closed).To(Equal(1))
|
||||
})
|
||||
|
||||
It("propagates a failure to open the session", func() {
|
||||
open := func(string) (asrSession, error) { return nil, errors.New("no streaming here") }
|
||||
in := make(chan *pb.TranscriptLiveRequest, 1)
|
||||
in <- cfg(0)
|
||||
close(in)
|
||||
out := make(chan *pb.TranscriptLiveResponse, 8)
|
||||
Expect(runLive(open, in, out)).To(MatchError(ContainSubstring("no streaming here")))
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("AudioTranscriptionStream", func() {
|
||||
run := func(ctx context.Context, n *NemoSpeech, req *pb.TranscriptRequest) ([]*pb.TranscriptStreamResponse, error) {
|
||||
GinkgoHelper()
|
||||
results := make(chan *pb.TranscriptStreamResponse)
|
||||
done := collect(results)
|
||||
err := n.AudioTranscriptionStream(ctx, req, results)
|
||||
return <-done, err
|
||||
}
|
||||
|
||||
// The RPC owns the channel: the gRPC host ranges over it and only returns
|
||||
// once it closes, so a rejection path that forgets to close hangs the call
|
||||
// instead of failing it.
|
||||
It("closes the results channel on every rejection path", func() {
|
||||
for _, n := range []*NemoSpeech{{fam: familyTTS}, {fam: familyASR}, {}} {
|
||||
_, err := run(context.Background(), n, &pb.TranscriptRequest{})
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(n.engineMu.TryLock()).To(BeTrue())
|
||||
n.engineMu.Unlock()
|
||||
}
|
||||
})
|
||||
|
||||
It("refuses a model loaded as another family", func() {
|
||||
n := &NemoSpeech{fam: familyTTS}
|
||||
_, err := run(context.Background(), n, &pb.TranscriptRequest{Dst: "x.wav"})
|
||||
Expect(status.Code(err)).To(Equal(codes.Unimplemented))
|
||||
Expect(err.Error()).To(ContainSubstring("tts"))
|
||||
})
|
||||
|
||||
It("requires a destination path", func() {
|
||||
n := &NemoSpeech{fam: familyASR}
|
||||
_, err := run(context.Background(), n, &pb.TranscriptRequest{})
|
||||
Expect(status.Code(err)).To(Equal(codes.InvalidArgument))
|
||||
})
|
||||
|
||||
// Cancellation is checked before the decode so a client that has already
|
||||
// gone away does not pay for an ffmpeg run, and so the check cannot be
|
||||
// mistaken for the decode failing.
|
||||
It("returns cancelled without touching the audio", func() {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
n := &NemoSpeech{fam: familyASR}
|
||||
_, err := run(ctx, n, &pb.TranscriptRequest{
|
||||
Dst: filepath.Join(GinkgoT().TempDir(), "absent.wav"),
|
||||
})
|
||||
Expect(status.Code(err)).To(Equal(codes.Canceled))
|
||||
})
|
||||
|
||||
It("reports an audio file it cannot read", func() {
|
||||
n := &NemoSpeech{fam: familyASR}
|
||||
_, err := run(context.Background(), n, &pb.TranscriptRequest{
|
||||
Dst: filepath.Join(GinkgoT().TempDir(), "absent.wav"),
|
||||
})
|
||||
Expect(status.Code(err)).To(Equal(codes.InvalidArgument))
|
||||
})
|
||||
|
||||
// Same ordering constraint as the offline path: a clip that decodes to no
|
||||
// samples has to be refused before a session is opened, which is also
|
||||
// before any bound entry point is called. Nothing is loaded here, so a
|
||||
// guard placed after the open would panic instead of failing.
|
||||
It("refuses a decodable clip that carries no samples, before opening a session", func() {
|
||||
path := filepath.Join(GinkgoT().TempDir(), "silence.wav")
|
||||
writeMono16kWAV(path, 0)
|
||||
|
||||
n := &NemoSpeech{fam: familyASR}
|
||||
var err error
|
||||
Expect(func() {
|
||||
_, err = run(context.Background(), n, &pb.TranscriptRequest{Dst: path})
|
||||
}).ToNot(Panic())
|
||||
Expect(status.Code(err)).To(Equal(codes.InvalidArgument))
|
||||
Expect(err.Error()).To(ContainSubstring("empty audio"))
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("AudioTranscriptionLive", func() {
|
||||
It("refuses a model loaded as another family and closes the output", func() {
|
||||
n := &NemoSpeech{fam: familyNMT}
|
||||
in := make(chan *pb.TranscriptLiveRequest)
|
||||
close(in)
|
||||
out := make(chan *pb.TranscriptLiveResponse)
|
||||
done := collect(out)
|
||||
|
||||
err := n.AudioTranscriptionLive(in, out)
|
||||
Expect(status.Code(err)).To(Equal(codes.Unimplemented))
|
||||
Expect(<-done).To(BeEmpty())
|
||||
Expect(n.engineMu.TryLock()).To(BeTrue())
|
||||
n.engineMu.Unlock()
|
||||
})
|
||||
|
||||
It("refuses an unloaded model", func() {
|
||||
n := &NemoSpeech{}
|
||||
in := make(chan *pb.TranscriptLiveRequest)
|
||||
close(in)
|
||||
out := make(chan *pb.TranscriptLiveResponse)
|
||||
done := collect(out)
|
||||
|
||||
err := n.AudioTranscriptionLive(in, out)
|
||||
Expect(status.Code(err)).To(Equal(codes.Unimplemented))
|
||||
Expect(<-done).To(BeEmpty())
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,367 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"math"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"time"
|
||||
"unsafe"
|
||||
|
||||
"github.com/go-audio/audio"
|
||||
"github.com/go-audio/wav"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
|
||||
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
|
||||
)
|
||||
|
||||
// writeMono16kWAV writes `frames` samples of 16 kHz mono 16-bit silence.
|
||||
// That is already AudioToWav's target format, so the decode path copies the
|
||||
// file through instead of shelling out to ffmpeg, which the test host may not
|
||||
// have.
|
||||
func writeMono16kWAV(path string, frames int) {
|
||||
GinkgoHelper()
|
||||
f, err := os.Create(path)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
enc := wav.NewEncoder(f, 16000, 16, 1, 1)
|
||||
Expect(enc.Write(&audio.IntBuffer{
|
||||
Format: &audio.Format{NumChannels: 1, SampleRate: 16000},
|
||||
SourceBitDepth: 16,
|
||||
Data: make([]int, frames),
|
||||
})).To(Succeed())
|
||||
Expect(enc.Close()).To(Succeed())
|
||||
Expect(f.Close()).To(Succeed())
|
||||
}
|
||||
|
||||
var _ = Describe("wordsToSegments", func() {
|
||||
It("groups words into one segment per speaker run", func() {
|
||||
words := []asrWord{
|
||||
{Text: "hello", Start: 0, End: 400, Speaker: 1},
|
||||
{Text: "there", Start: 400, End: 800, Speaker: 1},
|
||||
{Text: "hi", Start: 900, End: 1200, Speaker: 2},
|
||||
}
|
||||
segs := wordsToSegments(words, false)
|
||||
Expect(segs).To(HaveLen(2))
|
||||
Expect(segs[0].Text).To(Equal("hello there"))
|
||||
Expect(segs[1].Text).To(Equal("hi"))
|
||||
})
|
||||
|
||||
// A run is bounded by a CHANGE of speaker, not by the speaker id being new.
|
||||
// Grouping that keyed on the id itself (a map, or a comparison against the
|
||||
// first word) would merge the two A turns into one segment spanning B, and
|
||||
// the three-word spec above cannot see that because it never returns to an
|
||||
// earlier speaker.
|
||||
It("starts a new segment when an earlier speaker takes another turn", func() {
|
||||
words := []asrWord{
|
||||
{Text: "one", Start: 0, End: 100, Speaker: 1},
|
||||
{Text: "two", Start: 100, End: 200, Speaker: 2},
|
||||
{Text: "three", Start: 200, End: 300, Speaker: 1},
|
||||
}
|
||||
segs := wordsToSegments(words, false)
|
||||
Expect(segs).To(HaveLen(3))
|
||||
Expect(segs[0].Text).To(Equal("one"))
|
||||
Expect(segs[1].Text).To(Equal("two"))
|
||||
Expect(segs[2].Text).To(Equal("three"))
|
||||
})
|
||||
|
||||
// TranscriptSegment.start/end are int64 nanoseconds, not seconds:
|
||||
// core/backend/transcript.go reads them straight into a time.Duration. The
|
||||
// runtime reports word offsets in milliseconds (src/asr/types.h:46).
|
||||
It("converts millisecond word times to nanoseconds", func() {
|
||||
words := []asrWord{{Text: "a", Start: 1500, End: 2250, Speaker: 0}}
|
||||
segs := wordsToSegments(words, false)
|
||||
Expect(segs).To(HaveLen(1))
|
||||
Expect(time.Duration(segs[0].Start)).To(Equal(1500 * time.Millisecond))
|
||||
Expect(time.Duration(segs[0].End)).To(Equal(2250 * time.Millisecond))
|
||||
})
|
||||
|
||||
It("spans a segment from its first word's start to its last word's end", func() {
|
||||
words := []asrWord{
|
||||
{Text: "a", Start: 100, End: 200, Speaker: 0},
|
||||
{Text: "b", Start: 500, End: 900, Speaker: 0},
|
||||
}
|
||||
segs := wordsToSegments(words, false)
|
||||
Expect(segs).To(HaveLen(1))
|
||||
Expect(time.Duration(segs[0].Start)).To(Equal(100 * time.Millisecond))
|
||||
Expect(time.Duration(segs[0].End)).To(Equal(900 * time.Millisecond))
|
||||
})
|
||||
|
||||
It("produces a single segment when no speaker tags are present", func() {
|
||||
words := []asrWord{
|
||||
{Text: "a", Start: 0, End: 100, Speaker: 0},
|
||||
{Text: "b", Start: 100, End: 200, Speaker: 0},
|
||||
}
|
||||
segs := wordsToSegments(words, false)
|
||||
Expect(segs).To(HaveLen(1))
|
||||
Expect(segs[0].Text).To(Equal("a b"))
|
||||
})
|
||||
|
||||
It("returns no segments for no words", func() {
|
||||
Expect(wordsToSegments(nil, false)).To(BeEmpty())
|
||||
Expect(wordsToSegments([]asrWord{}, false)).To(BeEmpty())
|
||||
})
|
||||
|
||||
It("numbers the segments from zero in order", func() {
|
||||
words := []asrWord{
|
||||
{Text: "a", Speaker: 1},
|
||||
{Text: "b", Speaker: 2},
|
||||
{Text: "c", Speaker: 3},
|
||||
}
|
||||
segs := wordsToSegments(words, false)
|
||||
Expect(segs).To(HaveLen(3))
|
||||
for i, s := range segs {
|
||||
Expect(s.Id).To(Equal(int32(i)))
|
||||
}
|
||||
})
|
||||
|
||||
// TranscriptSegment.Words is what core/backend/transcript.go turns into the
|
||||
// response's word list, so an unset one makes timestamp_granularities:
|
||||
// ["word"] come back empty however good the timings were.
|
||||
It("attaches the per-word timings only when they were asked for", func() {
|
||||
words := []asrWord{
|
||||
{Text: "a", Start: 0, End: 100},
|
||||
{Text: "b", Start: 100, End: 250},
|
||||
}
|
||||
with := wordsToSegments(words, true)
|
||||
Expect(with[0].Words).To(HaveLen(2))
|
||||
Expect(with[0].Words[1].Text).To(Equal("b"))
|
||||
Expect(time.Duration(with[0].Words[1].Start)).To(Equal(100 * time.Millisecond))
|
||||
Expect(time.Duration(with[0].Words[1].End)).To(Equal(250 * time.Millisecond))
|
||||
|
||||
Expect(wordsToSegments(words, false)[0].Words).To(BeEmpty())
|
||||
})
|
||||
|
||||
// A speaker change splits the run, and each segment must carry only its own
|
||||
// words rather than the whole utterance's.
|
||||
It("gives each speaker run only its own words", func() {
|
||||
segs := wordsToSegments([]asrWord{
|
||||
{Text: "a", Speaker: 1},
|
||||
{Text: "b", Speaker: 2},
|
||||
}, true)
|
||||
Expect(segs).To(HaveLen(2))
|
||||
Expect(segs[0].Words).To(HaveLen(1))
|
||||
Expect(segs[0].Words[0].Text).To(Equal("a"))
|
||||
Expect(segs[1].Words[0].Text).To(Equal("b"))
|
||||
})
|
||||
|
||||
// The C ABI documents the speaker tag as 1-based with 0 meaning "untagged",
|
||||
// so a run of untagged words must not come back attributed to a speaker
|
||||
// literally named "0".
|
||||
It("labels a diarized run and leaves an untagged one unlabelled", func() {
|
||||
Expect(wordsToSegments([]asrWord{{Text: "a", Speaker: 2}}, false)[0].Speaker).To(Equal("2"))
|
||||
Expect(wordsToSegments([]asrWord{{Text: "a", Speaker: 0}}, false)[0].Speaker).To(BeEmpty())
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("wordsRequested", func() {
|
||||
It("recognises the OpenAI word granularity in any casing or padding", func() {
|
||||
Expect(wordsRequested([]string{"word"})).To(BeTrue())
|
||||
Expect(wordsRequested([]string{"segment", " Word "})).To(BeTrue())
|
||||
})
|
||||
|
||||
It("defaults to segment level", func() {
|
||||
Expect(wordsRequested(nil)).To(BeFalse())
|
||||
Expect(wordsRequested([]string{"segment"})).To(BeFalse())
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("recognizeF32", func() {
|
||||
// &pcm[0] panics on a zero-length slice, and a silent or empty upload is
|
||||
// ordinary input rather than an exotic one. The C side rejects empty audio
|
||||
// too, but Go never gets that far.
|
||||
It("refuses empty audio instead of indexing an empty slice", func() {
|
||||
// A zero options struct is enough: the guard has to fire before the
|
||||
// options are ever handed across the ABI, and building real ones would
|
||||
// need the library bound, which this spec deliberately does not.
|
||||
opts := cASRRecognitionOptions{}
|
||||
for _, pcm := range [][]float32{nil, {}} {
|
||||
var (
|
||||
handle uintptr
|
||||
err error
|
||||
)
|
||||
Expect(func() { handle, err = recognizeF32(0, &opts, pcm, 16000) }).ToNot(Panic())
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(status.Code(err)).To(Equal(codes.InvalidArgument))
|
||||
Expect(err.Error()).To(ContainSubstring("empty audio"))
|
||||
Expect(handle).To(BeZero())
|
||||
}
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("AudioTranscription", func() {
|
||||
// The gate has to fire before anything expensive: a model loaded as TTS
|
||||
// cannot transcribe whatever the request says, and reading the audio first
|
||||
// would report a file problem for a configuration one.
|
||||
It("refuses a model loaded as another family, before it reads the audio", func() {
|
||||
n := &NemoSpeech{fam: familyTTS}
|
||||
_, err := n.AudioTranscription(context.Background(), &pb.TranscriptRequest{
|
||||
Dst: filepath.Join(GinkgoT().TempDir(), "does-not-exist.wav"),
|
||||
})
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(status.Code(err)).To(Equal(codes.Unimplemented))
|
||||
Expect(err.Error()).To(ContainSubstring("tts"))
|
||||
})
|
||||
|
||||
It("refuses an unloaded model", func() {
|
||||
n := &NemoSpeech{}
|
||||
_, err := n.AudioTranscription(context.Background(), &pb.TranscriptRequest{Dst: "ignored.wav"})
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(status.Code(err)).To(Equal(codes.Unimplemented))
|
||||
})
|
||||
|
||||
// A lock leaked on a rejection path deadlocks the next request rather than
|
||||
// failing it, which is far harder to diagnose than the failure itself.
|
||||
It("releases the engine lock on every rejection path", func() {
|
||||
n := &NemoSpeech{fam: familyTTS}
|
||||
_, err := n.AudioTranscription(context.Background(), &pb.TranscriptRequest{Dst: "x.wav"})
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(n.engineMu.TryLock()).To(BeTrue())
|
||||
n.engineMu.Unlock()
|
||||
})
|
||||
|
||||
It("reports an audio file it cannot read", func() {
|
||||
n := &NemoSpeech{fam: familyASR}
|
||||
_, err := n.AudioTranscription(context.Background(), &pb.TranscriptRequest{
|
||||
Dst: filepath.Join(GinkgoT().TempDir(), "absent.wav"),
|
||||
})
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(status.Code(err)).To(Equal(codes.InvalidArgument))
|
||||
Expect(n.engineMu.TryLock()).To(BeTrue())
|
||||
n.engineMu.Unlock()
|
||||
})
|
||||
|
||||
It("requires a destination path", func() {
|
||||
n := &NemoSpeech{fam: familyASR}
|
||||
_, err := n.AudioTranscription(context.Background(), &pb.TranscriptRequest{})
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(status.Code(err)).To(Equal(codes.InvalidArgument))
|
||||
})
|
||||
|
||||
// The whole rejection path end to end, on the input that actually reaches
|
||||
// it: a silent or truncated upload decodes to zero samples, and the guard
|
||||
// has to fire between the decode and the ABI. Nothing is loaded here (no
|
||||
// recognizer, and the specs that bind the library may not have run), so
|
||||
// this also pins the ORDER: a guard placed after the options are built
|
||||
// calls a nil-bound entry point and panics rather than failing.
|
||||
It("refuses a decodable clip that carries no samples", func() {
|
||||
path := filepath.Join(GinkgoT().TempDir(), "silence.wav")
|
||||
writeMono16kWAV(path, 0)
|
||||
|
||||
n := &NemoSpeech{fam: familyASR, recognizer: 0}
|
||||
var err error
|
||||
Expect(func() {
|
||||
_, err = n.AudioTranscription(context.Background(), &pb.TranscriptRequest{Dst: path})
|
||||
}).ToNot(Panic())
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(status.Code(err)).To(Equal(codes.InvalidArgument))
|
||||
Expect(err.Error()).To(ContainSubstring("empty audio"))
|
||||
Expect(n.engineMu.TryLock()).To(BeTrue())
|
||||
n.engineMu.Unlock()
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("sampleRateOf", func() {
|
||||
// 0 is not "unknown" to this runtime: nemo_speech_asr_recognize_f32 and
|
||||
// nemo_speech_asr_stream_push_f32 both read a 0 rate as "these samples are
|
||||
// already at the model rate" and skip resampling. Falling back to it for an
|
||||
// undecodable header would silently pitch-shift the audio instead of
|
||||
// failing, so an unknown rate has to be an error.
|
||||
It("rejects a buffer whose format the decoder did not fill in", func() {
|
||||
_, err := sampleRateOf(&audio.IntBuffer{})
|
||||
Expect(err).To(HaveOccurred())
|
||||
})
|
||||
|
||||
It("rejects a non-positive sample rate", func() {
|
||||
_, err := sampleRateOf(&audio.IntBuffer{Format: &audio.Format{SampleRate: 0, NumChannels: 1}})
|
||||
Expect(err).To(HaveOccurred())
|
||||
})
|
||||
|
||||
// The WAV header carries the sample rate as an unsigned 32-bit field, which
|
||||
// go-audio widens to int. Anything above the int32 range therefore passes a
|
||||
// "> 0" test and then narrows to a NEGATIVE rate, which the runtime would take
|
||||
// as a resampling ratio rather than reject. The failure is silent, so the
|
||||
// bound is asserted rather than left to the caller.
|
||||
//
|
||||
// Written as a conversion plus one rather than as the constant MaxInt32+1:
|
||||
// the untyped form does not fit an int on a 32-bit build and would not
|
||||
// compile there, while this wraps to a negative rate the same guard rejects.
|
||||
It("rejects a rate that would not survive the narrowing to int32", func() {
|
||||
_, err := sampleRateOf(&audio.IntBuffer{
|
||||
Format: &audio.Format{SampleRate: int(math.MaxInt32) + 1, NumChannels: 1},
|
||||
})
|
||||
Expect(err).To(HaveOccurred())
|
||||
})
|
||||
|
||||
It("returns the decoded rate", func() {
|
||||
rate, err := sampleRateOf(&audio.IntBuffer{Format: &audio.Format{SampleRate: 22050, NumChannels: 1}})
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(rate).To(Equal(int32(22050)))
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("decodeAudioMono16k", func() {
|
||||
It("decodes a 16 kHz mono WAV to float32 samples at its own rate", func() {
|
||||
path := filepath.Join(GinkgoT().TempDir(), "silence.wav")
|
||||
writeMono16kWAV(path, 800)
|
||||
|
||||
pcm, rate, err := decodeAudioMono16k(path)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(rate).To(Equal(int32(16000)))
|
||||
Expect(pcm).To(HaveLen(800))
|
||||
})
|
||||
|
||||
// A zero-frame WAV is what a truncated upload decodes to, and it is the
|
||||
// input recognizeF32's guard exists for.
|
||||
It("decodes a WAV with no frames to an empty slice", func() {
|
||||
path := filepath.Join(GinkgoT().TempDir(), "empty.wav")
|
||||
writeMono16kWAV(path, 0)
|
||||
|
||||
pcm, _, err := decodeAudioMono16k(path)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(pcm).To(BeEmpty())
|
||||
})
|
||||
|
||||
It("reports a file that does not exist", func() {
|
||||
_, _, err := decodeAudioMono16k(filepath.Join(GinkgoT().TempDir(), "nope.wav"))
|
||||
Expect(err).To(HaveOccurred())
|
||||
})
|
||||
})
|
||||
|
||||
// The six frame counts on the recognizer-attached diarizer are
|
||||
// sentinel-sensitive and invisible to every other check in the tree.
|
||||
// src/asr/c_api.cpp:151-165 applies five of them when they are > 0 but applies
|
||||
// left_context_frames when it is >= 0, so a dropped -1 does not fall back to
|
||||
// the model's own streaming geometry, it pins the left context to zero. The
|
||||
// struct is the right shape either way, so abi_test.go's layout assertions
|
||||
// cannot see it.
|
||||
var _ = Describe("asrDiarConfig", func() {
|
||||
It("keeps the model path it was given", func() {
|
||||
Expect(asrDiarConfig(42).ModelPath).To(Equal(uintptr(42)))
|
||||
})
|
||||
|
||||
// A config sent with the wrong size has every field past it ignored by
|
||||
// HAS_FIELD, and the diarizer attaches with defaults instead of failing.
|
||||
It("declares the size the runtime validates against", func() {
|
||||
Expect(asrDiarConfig(42).Size).To(Equal(unsafe.Sizeof(cASRDiarConfig{})))
|
||||
})
|
||||
|
||||
It("leaves every frame count at the sentinel that means default", func() {
|
||||
cfg := asrDiarConfig(42)
|
||||
Expect(cfg.ChunkFrames).To(Equal(diarGeometryDefault))
|
||||
Expect(cfg.RightContextFrames).To(Equal(diarGeometryDefault))
|
||||
Expect(cfg.LeftContextFrames).To(Equal(diarGeometryDefault))
|
||||
Expect(cfg.FIFOFrames).To(Equal(diarGeometryDefault))
|
||||
Expect(cfg.SpkcacheFrames).To(Equal(diarGeometryDefault))
|
||||
Expect(cfg.UpdatePeriodFrames).To(Equal(diarGeometryDefault))
|
||||
})
|
||||
|
||||
// Stated separately from the field-by-field assertions above: the whole
|
||||
// group is only "unset" to the runtime while the sentinel stays negative,
|
||||
// and zero is a value it would apply to the left context.
|
||||
It("uses a negative sentinel, not zero", func() {
|
||||
Expect(diarGeometryDefault).To(BeNumerically("<", 0))
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,86 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"math"
|
||||
"os"
|
||||
"path/filepath"
|
||||
|
||||
"github.com/go-audio/audio"
|
||||
"github.com/go-audio/wav"
|
||||
"github.com/mudler/LocalAI/pkg/utils"
|
||||
)
|
||||
|
||||
// decodeAudioMono16k converts an arbitrary audio file to 16 kHz mono PCM and
|
||||
// returns the float32 samples together with the rate they are actually at.
|
||||
//
|
||||
// pkg/utils exposes the ffmpeg normalisation (AudioToWav) but no decode, so
|
||||
// every Go ASR backend pairs it with go-audio itself. This mirrors
|
||||
// backend/go/parakeet-cpp rather than adding a shared helper: the backends
|
||||
// differ in what they need back (parakeet wants a duration, this one wants the
|
||||
// sample rate to hand to the runtime), so a shared signature would be a
|
||||
// lowest-common-denominator of both.
|
||||
func decodeAudioMono16k(path string) ([]float32, int32, error) {
|
||||
dir, err := os.MkdirTemp("", "nemo-speech")
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
defer func() { _ = os.RemoveAll(dir) }()
|
||||
|
||||
// A WAV already at 16 kHz mono 16-bit is hardlinked or copied through
|
||||
// without spawning ffmpeg, so the common case costs nothing.
|
||||
converted := filepath.Join(dir, "converted.wav")
|
||||
if err := utils.AudioToWav(path, converted); err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
|
||||
// #nosec G304 -- converted is filepath.Join of a directory this function just
|
||||
// created with os.MkdirTemp and a constant basename. The request-controlled
|
||||
// path is the INPUT to AudioToWav and never reaches this open.
|
||||
fh, err := os.Open(converted)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
defer func() { _ = fh.Close() }()
|
||||
|
||||
buf, err := wav.NewDecoder(fh).FullPCMBuffer()
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
|
||||
// The rate is read back from the decoded file rather than assumed to be
|
||||
// 16000. AudioToWav always lands there today, but the runtime resamples
|
||||
// anything from 8 to 96 kHz off this number, so a wrong one would not fail,
|
||||
// it would silently pitch-shift the audio and quietly degrade the transcript.
|
||||
rate, err := sampleRateOf(buf)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
return buf.AsFloat32Buffer().Data, rate, nil
|
||||
}
|
||||
|
||||
// sampleRateOf reads the decoded rate back off the buffer.
|
||||
//
|
||||
// It is an error rather than a zero fallback because 0 is not "unknown" to this
|
||||
// runtime: nemo_speech_asr_recognize_f32 and nemo_speech_asr_stream_push_f32
|
||||
// both read a 0 rate as "these samples are already at the model rate" and skip
|
||||
// resampling (include/nemo_speech/asr.h). Handing 0 over for a header the
|
||||
// decoder could not read would not fail, it would silently pitch-shift the
|
||||
// audio and quietly degrade the transcript, which is the same failure the
|
||||
// caller comment warns about for a wrong rate.
|
||||
//
|
||||
// The upper bound is what makes the narrowing to int32 safe rather than merely
|
||||
// unlikely. go-audio reads the WAV header's sample rate as an unsigned 32-bit
|
||||
// field into an int, so on a 64-bit build a header claiming more than 2^31-1
|
||||
// survives the "> 0" test and then narrows to a NEGATIVE rate, which the runtime
|
||||
// would take as a resampling ratio. Nothing this backend decodes can reach that
|
||||
// today (AudioToWav either passes through a WAV it has confirmed is exactly
|
||||
// 16 kHz or runs ffmpeg with -ar 16000), but that is a property of a helper in
|
||||
// another package, and this function exists precisely because the rate is read
|
||||
// back rather than assumed.
|
||||
func sampleRateOf(buf *audio.IntBuffer) (int32, error) {
|
||||
if buf.Format == nil || buf.Format.SampleRate <= 0 || buf.Format.SampleRate > math.MaxInt32 {
|
||||
return 0, errors.New("nemo-speech-cpp: decoded audio has no usable sample rate")
|
||||
}
|
||||
return int32(buf.Format.SampleRate), nil
|
||||
}
|
||||
@@ -0,0 +1,502 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"strconv"
|
||||
"unsafe"
|
||||
|
||||
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
|
||||
"github.com/mudler/xlog"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
)
|
||||
|
||||
// diarSegmentsMaxAttempts bounds the count-then-fill retry.
|
||||
//
|
||||
// On a finished stream the count is stable and one attempt is always enough.
|
||||
// The bound exists because the RPC holds engineMu for its whole body, so a
|
||||
// runtime whose count kept growing would not merely spin, it would block the
|
||||
// unload behind it.
|
||||
const diarSegmentsMaxAttempts = 4
|
||||
|
||||
// maxDiarSegments caps the buffer collectSegments will allocate from a count
|
||||
// the C side reported.
|
||||
//
|
||||
// make() panics rather than erroring on a length it cannot satisfy, and a
|
||||
// panic in an RPC handler takes the backend process down, so an uninitialised
|
||||
// or corrupted size_t coming back across the ABI would kill the model rather
|
||||
// than fail the request. The ceiling turns that into a diagnosable error.
|
||||
//
|
||||
// It is set far above anything real: a segment spans at least one 80 ms frame,
|
||||
// so 2^22 segments is upwards of 93 hours of audio, and the buffer itself
|
||||
// would already be 100 MB at 24 bytes each.
|
||||
const maxDiarSegments = 1 << 22
|
||||
|
||||
// diarSegmenter is the result half of the diarization C API: the two-call
|
||||
// protocol nemo_speech_diar_segments documents.
|
||||
//
|
||||
// The two calls are the same C function with a different `out`, but they are
|
||||
// separate methods here because their contracts differ. countSegments passes
|
||||
// out=NULL, which the runtime answers by writing *count and returning OK
|
||||
// without touching a buffer. fillSegments passes a real buffer and gets
|
||||
// INVALID_ARGUMENT if it is too short, having written *count first, which is
|
||||
// what makes a growth retry possible at all.
|
||||
type diarSegmenter interface {
|
||||
// countSegments is the size query. It never fails for lack of a buffer.
|
||||
countSegments() (uint64, error)
|
||||
// fillSegments fills buf and returns the count the runtime reported. That
|
||||
// count is meaningful even alongside an error: on a short buffer the
|
||||
// runtime writes it before rejecting the call.
|
||||
fillSegments(buf []cDiarSegment) (uint64, error)
|
||||
}
|
||||
|
||||
// diarStream is one diarization job over the C API, narrowed to what the RPC
|
||||
// uses.
|
||||
//
|
||||
// It is an interface for the same reason asrSession is: no Sortformer GGUF is
|
||||
// small enough to keep in the tree, so without a seam at the ABI the loop on
|
||||
// top of it (the empty guard, chunking, finish-before-query, the growth retry)
|
||||
// would have no test at all. A fake here scripts what C returns; it does not
|
||||
// pretend to diarize anything.
|
||||
type diarStream interface {
|
||||
diarSegmenter
|
||||
push(pcm []float32, sampleRate int32) error
|
||||
finish() error
|
||||
close()
|
||||
}
|
||||
|
||||
// diarStreamOpener creates a job. n.openDiarStream is the C-backed one.
|
||||
//
|
||||
// The segmentation config is handed over at open time rather than per query
|
||||
// because it belongs to the whole job: every segments call on one stream must
|
||||
// use the same postprocessing or the segment ids would not be comparable
|
||||
// between calls.
|
||||
type diarStreamOpener func(cfg *cDiarSegmentationConfig) (diarStream, error)
|
||||
|
||||
// cDiarStream is the real diarStream, over one nemo_speech_diar_stream.
|
||||
type cDiarStream struct {
|
||||
handle uintptr
|
||||
cfg *cDiarSegmentationConfig
|
||||
}
|
||||
|
||||
// cfgPtr hands the segmentation config to C, or NULL when the request asked
|
||||
// for no postprocessing. NULL is not the same as a zeroed struct in spirit
|
||||
// even though src/asr/c_api.cpp treats them alike today: diar.h documents NULL
|
||||
// as "library defaults", so it is the one form that cannot be invalidated by a
|
||||
// future field whose sentinel is not zero.
|
||||
func (s *cDiarStream) cfgPtr() unsafe.Pointer {
|
||||
if s.cfg == nil {
|
||||
return nil
|
||||
}
|
||||
// #nosec G103 -- a plain *T to unsafe.Pointer conversion of a non-nil,
|
||||
// GC-traced field. cDiarSegmentationConfig is pure scalars (no uintptr
|
||||
// members to pin) and the stream owns it for its whole life, so the only
|
||||
// requirement is that it outlive the DiarSegments call, which it does.
|
||||
return unsafe.Pointer(s.cfg)
|
||||
}
|
||||
|
||||
func (s *cDiarStream) push(pcm []float32, sampleRate int32) error {
|
||||
// &pcm[0] panics on an empty slice before the C side ever sees the call.
|
||||
if len(pcm) == 0 {
|
||||
return nil
|
||||
}
|
||||
if st := DiarStreamPushF32(s.handle, &pcm[0], uint64(len(pcm)), sampleRate); st != 0 {
|
||||
return statusErrorf(st, "nemo-speech-cpp: diarization push: %s", ASRLastError())
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *cDiarStream) finish() error {
|
||||
if st := DiarStreamFinish(s.handle); st != 0 {
|
||||
return statusErrorf(st, "nemo-speech-cpp: diarization finish: %s", ASRLastError())
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *cDiarStream) close() { DiarStreamClose(s.handle) }
|
||||
|
||||
func (s *cDiarStream) countSegments() (uint64, error) {
|
||||
var count uint64
|
||||
// out=NULL and capacity=0: the size query. The runtime reads capacity only
|
||||
// once it has a buffer to check it against.
|
||||
if st := DiarSegments(s.handle, s.cfgPtr(), nil, 0, &count); st != 0 {
|
||||
return 0, statusErrorf(st,
|
||||
"nemo-speech-cpp: diarization segment count: %s", ASRLastError())
|
||||
}
|
||||
return count, nil
|
||||
}
|
||||
|
||||
func (s *cDiarStream) fillSegments(buf []cDiarSegment) (uint64, error) {
|
||||
if len(buf) == 0 {
|
||||
// A NULL out would silently turn this into a second size query, and the
|
||||
// caller would read it as "filled nothing" rather than "asked nothing".
|
||||
return 0, status.Error(codes.Internal,
|
||||
"nemo-speech-cpp: diarization segment fill needs a buffer")
|
||||
}
|
||||
var count uint64
|
||||
// #nosec G103 -- &buf[0] is guarded by the empty check above, and the
|
||||
// capacity handed over is exactly len(buf), so the runtime cannot write past
|
||||
// the caller's allocation. collectSegments sizes buf under maxDiarSegments
|
||||
// and rejects a reported count larger than it rather than slicing to it.
|
||||
st := DiarSegments(s.handle, s.cfgPtr(), unsafe.Pointer(&buf[0]), uint64(len(buf)), &count)
|
||||
if st != 0 {
|
||||
// count is returned alongside the error on purpose: a too-small buffer
|
||||
// is rejected only after the runtime has written the size it wanted.
|
||||
return count, statusErrorf(st,
|
||||
"nemo-speech-cpp: diarization segments: %s", ASRLastError())
|
||||
}
|
||||
return count, nil
|
||||
}
|
||||
|
||||
// diarGeometryDefault is the sentinel that means "keep the preset's value" for
|
||||
// every one of nemo_speech_diar_model_config's six frame counts.
|
||||
//
|
||||
// It has to be negative, not zero, and that is not a style choice.
|
||||
// src/asr/c_api.cpp:497-512 applies five of the six overrides when they are
|
||||
// > 0 but applies left_context_frames when it is >= 0, so a zero-valued config
|
||||
// reads as "unset" for five fields and as an explicit left context of zero for
|
||||
// the sixth. That silently changes the model's streaming geometry, and no
|
||||
// layout assertion can see it because the struct is the right shape either way.
|
||||
const diarGeometryDefault int32 = -1
|
||||
|
||||
// diarModelConfig builds the create-time config for the standalone diarizer.
|
||||
//
|
||||
// Extracted from loadDiarizer purely so the sentinels above can be asserted:
|
||||
// they are invisible to every other check in the tree, including the layout
|
||||
// assertions, so a spec pinning them is the only thing standing between a
|
||||
// dropped -1 and a quietly mis-configured model.
|
||||
//
|
||||
// modelPath is a C pointer from cstr, not a Go string, and the caller owns its
|
||||
// release. preset is deliberately left NULL, which diar.h reads as "streaming".
|
||||
// The "offline" preset is a different accuracy/latency tradeoff for long files
|
||||
// and is worth exposing, but not on an unverified guess: no Sortformer GGUF
|
||||
// exists here to measure the difference on.
|
||||
func diarModelConfig(modelPath uintptr, gpu int32) cDiarModelConfig {
|
||||
return cDiarModelConfig{
|
||||
Size: unsafe.Sizeof(cDiarModelConfig{}),
|
||||
ModelPath: modelPath,
|
||||
GPU: gpu,
|
||||
ChunkFrames: diarGeometryDefault,
|
||||
RightContextFrames: diarGeometryDefault,
|
||||
LeftContextFrames: diarGeometryDefault,
|
||||
FIFOFrames: diarGeometryDefault,
|
||||
SpkcacheFrames: diarGeometryDefault,
|
||||
UpdatePeriodFrames: diarGeometryDefault,
|
||||
}
|
||||
}
|
||||
|
||||
// loadDiarizer creates the standalone Sortformer diarizer.
|
||||
//
|
||||
// This must not take engineMu: Load is its only caller and already holds it.
|
||||
func (n *NemoSpeech) loadDiarizer(modelFile string) error {
|
||||
pathP, freePath := cstr(modelFile)
|
||||
defer freePath()
|
||||
|
||||
cfg := diarModelConfig(pathP, n.opts.gpu)
|
||||
|
||||
xlog.Info("nemo-speech-cpp: creating diarizer", "gpu", n.opts.gpu)
|
||||
|
||||
// #nosec G103 -- cfg is a local POD struct borrowed for this call only. Its
|
||||
// only uintptr member is ModelPath, the cstr allocation pinned by the
|
||||
// deferred freePath above (Preset is deliberately NULL), and
|
||||
// nemo_speech_diar_create deep-copies the path and retains nothing.
|
||||
if st := DiarCreate(unsafe.Pointer(&cfg), &n.diarizer); st != 0 {
|
||||
return statusErrorf(st, "nemo-speech-cpp: diarizer create: %s", ASRLastError())
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// openDiarStream starts a diarization job on the loaded model.
|
||||
//
|
||||
// The caller must hold engineMu.
|
||||
func (n *NemoSpeech) openDiarStream(cfg *cDiarSegmentationConfig) (diarStream, error) {
|
||||
var handle uintptr
|
||||
if st := DiarStreamOpen(n.diarizer, &handle); st != 0 {
|
||||
return nil, statusErrorf(st,
|
||||
"nemo-speech-cpp: diarization stream open: %s", ASRLastError())
|
||||
}
|
||||
return &cDiarStream{handle: handle, cfg: cfg}, nil
|
||||
}
|
||||
|
||||
// sizeofDiarSegmentationConfig is the size the runtime validates the config
|
||||
// against. It is a function so the specs can assert the value the config
|
||||
// actually carries rather than restate the number.
|
||||
func sizeofDiarSegmentationConfig() uintptr {
|
||||
return unsafe.Sizeof(cDiarSegmentationConfig{})
|
||||
}
|
||||
|
||||
// segmentationConfig maps the request's postprocessing knobs onto
|
||||
// nemo_speech_diar_segmentation_config, or returns nil when none were set.
|
||||
//
|
||||
// Only two of DiarizeRequest's tuning fields have a real equivalent here, and
|
||||
// both are exact rather than approximate: NeMo's ts_vad postprocessing is the
|
||||
// same algorithm the proto's wording describes.
|
||||
//
|
||||
// - min_duration_on ("discard segments shorter than this") is min_duration_sec
|
||||
// ("drop segments shorter than this"), which c_api.cpp assigns to
|
||||
// DiarSegmentationCfg.min_duration_on.
|
||||
// - min_duration_off ("merge gaps shorter than this") is min_gap_sec ("fill
|
||||
// silence gaps shorter than this"), assigned to min_duration_off.
|
||||
//
|
||||
// The names cross over between the proto and the C header, which is exactly the
|
||||
// kind of transposition a layout assertion cannot see, so each mapping is
|
||||
// pinned by its own spec.
|
||||
//
|
||||
// Nothing is written for a non-positive value: the runtime tests every field
|
||||
// with > 0 and keeps its default otherwise, so a zero here means "unset" on
|
||||
// both sides.
|
||||
func segmentationConfig(req *pb.DiarizeRequest) *cDiarSegmentationConfig {
|
||||
cfg := cDiarSegmentationConfig{Size: sizeofDiarSegmentationConfig()}
|
||||
|
||||
var set bool
|
||||
if v := req.GetMinDurationOn(); v > 0 {
|
||||
cfg.MinDurationSec = float64(v)
|
||||
set = true
|
||||
}
|
||||
if v := req.GetMinDurationOff(); v > 0 {
|
||||
cfg.MinGapSec = float64(v)
|
||||
set = true
|
||||
}
|
||||
if !set {
|
||||
return nil
|
||||
}
|
||||
return &cfg
|
||||
}
|
||||
|
||||
// unsupportedRequestFields names the DiarizeRequest fields this backend cannot
|
||||
// honour, so they are logged rather than silently dropped.
|
||||
//
|
||||
// Each is a deliberate omission, not a gap waiting to be filled:
|
||||
//
|
||||
// - num_speakers, min_speakers, max_speakers: Sortformer is end-to-end and
|
||||
// its speaker capacity is fixed by the checkpoint (v2: 4).
|
||||
// nemo_speech_diar_num_speakers reports that capacity, it does not set it,
|
||||
// and there is no config field for a target count.
|
||||
// - clustering_threshold: there is no clustering stage. The nearest knob is
|
||||
// the onset/offset probability hysteresis, which is a different quantity on
|
||||
// a different scale, so mapping one onto the other would invent an
|
||||
// equivalence the header does not have.
|
||||
// - include_text: this pipeline carries no ASR at all (diar.h: "no ASR
|
||||
// involved"). Word-level speaker tags on a transcript are the ASR surface's
|
||||
// job, through diar_model plus enable_speaker_diarization.
|
||||
// - threads: neither nemo_speech_diar_model_config nor the segmentation
|
||||
// config has a thread count.
|
||||
func unsupportedRequestFields(req *pb.DiarizeRequest) []string {
|
||||
var out []string
|
||||
if req.GetNumSpeakers() != 0 {
|
||||
out = append(out, "num_speakers")
|
||||
}
|
||||
if req.GetMinSpeakers() != 0 {
|
||||
out = append(out, "min_speakers")
|
||||
}
|
||||
if req.GetMaxSpeakers() != 0 {
|
||||
out = append(out, "max_speakers")
|
||||
}
|
||||
if req.GetClusteringThreshold() != 0 {
|
||||
out = append(out, "clustering_threshold")
|
||||
}
|
||||
if req.GetIncludeText() {
|
||||
out = append(out, "include_text")
|
||||
}
|
||||
if req.GetThreads() != 0 {
|
||||
out = append(out, "threads")
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// collectSegments runs the count-then-fill protocol and returns the segments.
|
||||
//
|
||||
// The growth retry is not defensive padding. nemo_speech_diar_segments writes
|
||||
// *count and only then rejects a buffer that is too small, so the size a
|
||||
// rejected call reports is the size to retry with; without the retry a stream
|
||||
// that gained a segment between the two calls would fail the whole request.
|
||||
// Truncating to the first count instead would be worse still, dropping turns
|
||||
// with nothing to show for it.
|
||||
func collectSegments(s diarSegmenter) ([]cDiarSegment, error) {
|
||||
want, err := s.countSegments()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
for range diarSegmentsMaxAttempts {
|
||||
if want == 0 {
|
||||
// No segments means no fill: the fill call needs a non-empty buffer
|
||||
// to be distinguishable from a second size query.
|
||||
return nil, nil
|
||||
}
|
||||
if want > maxDiarSegments {
|
||||
return nil, status.Errorf(codes.Internal,
|
||||
"nemo-speech-cpp: diarization reported %d segments, above the %d ceiling", want, maxDiarSegments)
|
||||
}
|
||||
|
||||
buf := make([]cDiarSegment, want)
|
||||
got, fillErr := s.fillSegments(buf)
|
||||
if fillErr == nil {
|
||||
if got > want {
|
||||
// The runtime cannot report this on success (it rejects a short
|
||||
// buffer instead), so it means the ABI is not what this code
|
||||
// thinks it is. Slicing to it would read past the allocation.
|
||||
return nil, status.Errorf(codes.Internal,
|
||||
"nemo-speech-cpp: diarization returned %d segments for a %d-segment buffer", got, want)
|
||||
}
|
||||
return buf[:got], nil
|
||||
}
|
||||
// A count that did not grow means the call failed for some other
|
||||
// reason, and retrying the same size would just fail the same way.
|
||||
if got <= want {
|
||||
return nil, fillErr
|
||||
}
|
||||
want = got
|
||||
}
|
||||
|
||||
return nil, status.Error(codes.Internal,
|
||||
"nemo-speech-cpp: diarization segment count kept growing, giving up")
|
||||
}
|
||||
|
||||
// toDiarizeSegments converts the runtime's segments to the wire form.
|
||||
//
|
||||
// No unit conversion happens here, and that is the point: nemo_speech_diar_segment
|
||||
// carries start_time and end_time in SECONDS already (diar.h), and
|
||||
// DiarizeSegment.start/end are seconds too. The frame indices the model works
|
||||
// in never reach this layer, so nemo_speech_diar_seconds_per_frame is not
|
||||
// involved. The narrowing to float32 is the proto's choice of type; at 80 ms
|
||||
// resolution it is lossless for any clip short enough to hold in memory.
|
||||
//
|
||||
// The speaker label is the runtime's 1-based tag rendered as a decimal string,
|
||||
// which is what wordsToSegments emits for the ASR path. The same speaker has to
|
||||
// read the same way whether the caller diarized a file or transcribed it.
|
||||
func toDiarizeSegments(in []cDiarSegment) []*pb.DiarizeSegment {
|
||||
if len(in) == 0 {
|
||||
return nil
|
||||
}
|
||||
out := make([]*pb.DiarizeSegment, 0, len(in))
|
||||
for i, s := range in {
|
||||
out = append(out, &pb.DiarizeSegment{
|
||||
Id: int32(i),
|
||||
Start: float32(s.StartTime),
|
||||
End: float32(s.EndTime),
|
||||
Speaker: strconv.Itoa(int(s.Speaker)),
|
||||
})
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// distinctSpeakers counts the speaker labels present in the segments.
|
||||
//
|
||||
// This is what DiarizeResponse.num_speakers is documented to hold, and it is
|
||||
// NOT nemo_speech_diar_num_speakers: that reports the checkpoint's capacity
|
||||
// (four for Sortformer v2), so a two-person interview would come back claiming
|
||||
// four speakers.
|
||||
func distinctSpeakers(segs []*pb.DiarizeSegment) int32 {
|
||||
seen := make(map[string]struct{}, len(segs))
|
||||
for _, s := range segs {
|
||||
seen[s.GetSpeaker()] = struct{}{}
|
||||
}
|
||||
// #nosec G115 -- seen holds at most one entry per segment, and collectSegments
|
||||
// refuses any count above maxDiarSegments (2^22), so this is orders of
|
||||
// magnitude below the int32 the proto field is.
|
||||
return int32(len(seen))
|
||||
}
|
||||
|
||||
// diarizePCM drives one whole clip through a diarization job.
|
||||
//
|
||||
// The caller must hold engineMu.
|
||||
func diarizePCM(open diarStreamOpener, pcm []float32, sampleRate int32, cfg *cDiarSegmentationConfig) (*pb.DiarizeResponse, error) {
|
||||
// Before the stream is opened, not inside the push: a silent or truncated
|
||||
// upload decodes to zero samples, &pcm[0] panics on that, and there is no
|
||||
// diarization to be had from it anyway.
|
||||
if len(pcm) == 0 {
|
||||
return nil, status.Error(codes.InvalidArgument, "nemo-speech-cpp: empty audio")
|
||||
}
|
||||
|
||||
stream, err := open(cfg)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer stream.close()
|
||||
|
||||
// Chunked rather than pushed whole so the runtime advances as it goes
|
||||
// instead of buffering the entire clip before the first chunk boundary.
|
||||
for _, chunk := range chunkPCM(pcm, streamChunkSamples) {
|
||||
if err := stream.push(chunk, sampleRate); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
// Before the query, always: finish is what labels the audio tail, so
|
||||
// segmenting first drops the last turn of every clip.
|
||||
if err := stream.finish(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
raw, err := collectSegments(stream)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
segs := toDiarizeSegments(raw)
|
||||
out := &pb.DiarizeResponse{
|
||||
Segments: segs,
|
||||
NumSpeakers: distinctSpeakers(segs),
|
||||
}
|
||||
// 0 is the proto's "unknown" and the C API's "already at the model rate",
|
||||
// so a rate that means the latter must not be divided by.
|
||||
if sampleRate > 0 {
|
||||
out.Duration = float32(len(pcm)) / float32(sampleRate)
|
||||
}
|
||||
// Language and the per-segment text stay empty: there is no ASR in this
|
||||
// pipeline to fill them, and the proto documents both as optional.
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// Diarize labels who spoke when in the audio at req.Dst.
|
||||
//
|
||||
// The whole body runs inside withEngine, so the family check and the C calls
|
||||
// that trust the handle happen under a single acquisition of engineMu. The
|
||||
// audio decode is in there too, for the reason documented on
|
||||
// AudioTranscription: the backend already serialises RPCs, so the wider hold
|
||||
// costs nothing, and the narrower one is the gap Free can land in.
|
||||
func (n *NemoSpeech) Diarize(req *pb.DiarizeRequest) (pb.DiarizeResponse, error) {
|
||||
var out *pb.DiarizeResponse
|
||||
if err := n.withEngine(familyDiarization, func() error {
|
||||
r, err := n.diarize(req)
|
||||
out = r
|
||||
return err
|
||||
}); err != nil {
|
||||
return pb.DiarizeResponse{}, err
|
||||
}
|
||||
if out == nil {
|
||||
return pb.DiarizeResponse{}, status.Error(codes.Internal,
|
||||
"nemo-speech-cpp: diarization produced no result")
|
||||
}
|
||||
|
||||
// Assembled field by field rather than dereferenced: the RPC returns the
|
||||
// message by value and the message embeds a mutex, so copying the struct is
|
||||
// a copylocks violation.
|
||||
return pb.DiarizeResponse{
|
||||
Segments: out.Segments,
|
||||
NumSpeakers: out.NumSpeakers,
|
||||
Duration: out.Duration,
|
||||
Language: out.Language,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// diarize is Diarize's body. The caller must hold engineMu.
|
||||
func (n *NemoSpeech) diarize(req *pb.DiarizeRequest) (*pb.DiarizeResponse, error) {
|
||||
if req.GetDst() == "" {
|
||||
return nil, status.Error(codes.InvalidArgument,
|
||||
"nemo-speech-cpp: DiarizeRequest.dst (audio path) is required")
|
||||
}
|
||||
// Logged rather than rejected: a client that asks for a speaker count still
|
||||
// wants the diarization it can have, and a request that names a field this
|
||||
// backend drops should say so somewhere the operator can find it.
|
||||
if dropped := unsupportedRequestFields(req); len(dropped) > 0 {
|
||||
xlog.Warn("nemo-speech-cpp: ignoring diarization request fields this model has no equivalent for",
|
||||
"fields", dropped)
|
||||
}
|
||||
|
||||
pcm, sampleRate, err := decodeAudioMono16k(req.GetDst())
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "nemo-speech-cpp: read audio: %v", err)
|
||||
}
|
||||
|
||||
return diarizePCM(n.openDiarStream, pcm, sampleRate, segmentationConfig(req))
|
||||
}
|
||||
@@ -0,0 +1,540 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"path/filepath"
|
||||
"unsafe"
|
||||
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
|
||||
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
|
||||
)
|
||||
|
||||
// fakeDiarStream scripts what the C API returns for one diarization job.
|
||||
//
|
||||
// There is no Sortformer GGUF in the tree, so this is the only way the loop on
|
||||
// top of the ABI (the empty guard, chunking, the count-then-fill protocol, the
|
||||
// buffer growth retry) gets tested at all. It fakes the C contract, not the
|
||||
// model: segs is whatever nemo_speech_diar_segments would have produced.
|
||||
type fakeDiarStream struct {
|
||||
segs []cDiarSegment
|
||||
|
||||
// countErr and fillErrs script failures. fillErrs is consumed one entry per
|
||||
// fillSegments call so a growth retry can be scripted.
|
||||
countErr error
|
||||
fillErrs []error
|
||||
// queryCount, when non-zero, is what the size query reports instead of
|
||||
// len(segs), so a runtime that under-reported can be scripted.
|
||||
queryCount uint64
|
||||
// growTo, when non-zero, is the count reported by the FIRST fillSegments
|
||||
// call, standing in for a runtime whose segment list outgrew the size query.
|
||||
growTo uint64
|
||||
|
||||
pushed [][]float32
|
||||
rates []int32
|
||||
finished int
|
||||
closed int
|
||||
counts int
|
||||
fills int
|
||||
// opened records the segmentation config the opener was handed.
|
||||
cfg *cDiarSegmentationConfig
|
||||
}
|
||||
|
||||
func (f *fakeDiarStream) push(pcm []float32, rate int32) error {
|
||||
f.pushed = append(f.pushed, pcm)
|
||||
f.rates = append(f.rates, rate)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (f *fakeDiarStream) finish() error {
|
||||
f.finished++
|
||||
return nil
|
||||
}
|
||||
|
||||
func (f *fakeDiarStream) close() { f.closed++ }
|
||||
|
||||
func (f *fakeDiarStream) countSegments() (uint64, error) {
|
||||
f.counts++
|
||||
if f.countErr != nil {
|
||||
return 0, f.countErr
|
||||
}
|
||||
if f.queryCount > 0 {
|
||||
return f.queryCount, nil
|
||||
}
|
||||
return uint64(len(f.segs)), nil
|
||||
}
|
||||
|
||||
func (f *fakeDiarStream) fillSegments(buf []cDiarSegment) (uint64, error) {
|
||||
f.fills++
|
||||
var err error
|
||||
if len(f.fillErrs) > 0 {
|
||||
err, f.fillErrs = f.fillErrs[0], f.fillErrs[1:]
|
||||
}
|
||||
if f.fills == 1 && f.growTo > 0 {
|
||||
// The runtime writes *count before it rejects a short buffer, so a
|
||||
// growth failure still reports the count the caller needs.
|
||||
return f.growTo, err
|
||||
}
|
||||
n := copy(buf, f.segs)
|
||||
return uint64(n), err
|
||||
}
|
||||
|
||||
// allPushed flattens what the fake received, so a spec can assert the audio
|
||||
// arrived intact regardless of how it was chunked.
|
||||
func (f *fakeDiarStream) allPushed() []float32 {
|
||||
var out []float32
|
||||
for _, c := range f.pushed {
|
||||
out = append(out, c...)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func (f *fakeDiarStream) opener() diarStreamOpener {
|
||||
return func(cfg *cDiarSegmentationConfig) (diarStream, error) {
|
||||
f.cfg = cfg
|
||||
return f, nil
|
||||
}
|
||||
}
|
||||
|
||||
var _ = Describe("Diarize", func() {
|
||||
It("refuses when the loaded model is not a diarization model", func() {
|
||||
n := &NemoSpeech{fam: familyASR}
|
||||
_, err := n.Diarize(&pb.DiarizeRequest{Dst: "/tmp/whatever.wav"})
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(status.Code(err)).To(Equal(codes.Unimplemented))
|
||||
})
|
||||
|
||||
It("refuses on a model that was never loaded", func() {
|
||||
n := &NemoSpeech{}
|
||||
_, err := n.Diarize(&pb.DiarizeRequest{Dst: "/tmp/whatever.wav"})
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(status.Code(err)).To(Equal(codes.Unimplemented))
|
||||
})
|
||||
|
||||
// The family gate has to run before anything reads the request, or a
|
||||
// misrouted request would be reported as a bad path rather than as a model
|
||||
// that cannot diarize.
|
||||
It("reports a missing audio path on a diarization model", func() {
|
||||
n := &NemoSpeech{fam: familyDiarization}
|
||||
_, err := n.Diarize(&pb.DiarizeRequest{})
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(status.Code(err)).To(Equal(codes.InvalidArgument))
|
||||
Expect(err.Error()).To(ContainSubstring("dst"))
|
||||
})
|
||||
|
||||
It("reports audio it cannot read", func() {
|
||||
n := &NemoSpeech{fam: familyDiarization}
|
||||
missing := filepath.Join(GinkgoT().TempDir(), "absent.wav")
|
||||
_, err := n.Diarize(&pb.DiarizeRequest{Dst: missing})
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(status.Code(err)).To(Equal(codes.InvalidArgument))
|
||||
Expect(err.Error()).To(ContainSubstring("read audio"))
|
||||
})
|
||||
|
||||
// A rejection must leave the mutex free, or the next request deadlocks
|
||||
// rather than fails.
|
||||
It("releases the engine lock on every rejection path", func() {
|
||||
n := &NemoSpeech{fam: familyDiarization}
|
||||
_, err := n.Diarize(&pb.DiarizeRequest{})
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(n.engineMu.TryLock()).To(BeTrue())
|
||||
n.engineMu.Unlock()
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("diarizePCM", func() {
|
||||
// Task 7 found that a purego-bound entry point reached with zero samples
|
||||
// panics on &pcm[0], so a silent clip must be rejected before the stream is
|
||||
// ever opened, not inside the push.
|
||||
It("rejects empty audio without opening a stream", func() {
|
||||
opened := false
|
||||
open := func(*cDiarSegmentationConfig) (diarStream, error) {
|
||||
opened = true
|
||||
return &fakeDiarStream{}, nil
|
||||
}
|
||||
_, err := diarizePCM(open, nil, 16000, nil)
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(status.Code(err)).To(Equal(codes.InvalidArgument))
|
||||
Expect(err.Error()).To(ContainSubstring("empty audio"))
|
||||
Expect(opened).To(BeFalse())
|
||||
})
|
||||
|
||||
It("pushes the whole clip, finishes, and closes the stream", func() {
|
||||
pcm := make([]float32, streamChunkSamples*2+7)
|
||||
for i := range pcm {
|
||||
pcm[i] = float32(i)
|
||||
}
|
||||
f := &fakeDiarStream{}
|
||||
_, err := diarizePCM(f.opener(), pcm, 16000, nil)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
|
||||
Expect(f.allPushed()).To(Equal(pcm))
|
||||
Expect(f.pushed).To(HaveLen(3), "the clip must be chunked, not pushed whole")
|
||||
Expect(f.rates).To(HaveEach(int32(16000)))
|
||||
Expect(f.finished).To(Equal(1))
|
||||
Expect(f.closed).To(Equal(1))
|
||||
})
|
||||
|
||||
// Segments must come from a finished stream: the tail of the audio is only
|
||||
// labelled by finish, so asking first silently drops the last turn.
|
||||
It("finishes the stream before it asks for segments", func() {
|
||||
f := &fakeDiarStream{segs: []cDiarSegment{{StartTime: 0, EndTime: 1, Speaker: 1}}}
|
||||
f.fillErrs = nil
|
||||
var finishedAtCount int
|
||||
wrapped := func(cfg *cDiarSegmentationConfig) (diarStream, error) {
|
||||
f.cfg = cfg
|
||||
return &countObserver{fakeDiarStream: f, seen: &finishedAtCount}, nil
|
||||
}
|
||||
_, err := diarizePCM(wrapped, []float32{1, 2, 3}, 16000, nil)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(finishedAtCount).To(Equal(1), "the size query ran before finish")
|
||||
})
|
||||
|
||||
It("converts the runtime's seconds straight through and numbers the segments", func() {
|
||||
f := &fakeDiarStream{segs: []cDiarSegment{
|
||||
{StartTime: 0, EndTime: 0.8, Speaker: 1},
|
||||
{StartTime: 0.8, EndTime: 2.0, Speaker: 2},
|
||||
}}
|
||||
res, err := diarizePCM(f.opener(), []float32{1, 2, 3}, 16000, nil)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(res.Segments).To(HaveLen(2))
|
||||
|
||||
Expect(res.Segments[0].GetId()).To(Equal(int32(0)))
|
||||
Expect(res.Segments[0].GetStart()).To(BeNumerically("~", 0.0, 1e-6))
|
||||
Expect(res.Segments[0].GetEnd()).To(BeNumerically("~", 0.8, 1e-6))
|
||||
Expect(res.Segments[0].GetSpeaker()).To(Equal("1"))
|
||||
|
||||
Expect(res.Segments[1].GetId()).To(Equal(int32(1)))
|
||||
Expect(res.Segments[1].GetStart()).To(BeNumerically("~", 0.8, 1e-6))
|
||||
Expect(res.Segments[1].GetEnd()).To(BeNumerically("~", 2.0, 1e-6))
|
||||
Expect(res.Segments[1].GetSpeaker()).To(Equal("2"))
|
||||
})
|
||||
|
||||
It("reports the clip duration in seconds", func() {
|
||||
f := &fakeDiarStream{}
|
||||
res, err := diarizePCM(f.opener(), make([]float32, 32000), 16000, nil)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(res.GetDuration()).To(BeNumerically("~", 2.0, 1e-6))
|
||||
})
|
||||
|
||||
// 0 is the proto's documented "unknown", and it is also the C API's "these
|
||||
// samples are already at the model rate", so a rate that cannot be trusted
|
||||
// must not be turned into a duration.
|
||||
It("reports no duration when the sample rate is unknown", func() {
|
||||
f := &fakeDiarStream{}
|
||||
res, err := diarizePCM(f.opener(), make([]float32, 32000), 0, nil)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(res.GetDuration()).To(BeZero())
|
||||
})
|
||||
|
||||
It("hands the segmentation config to the opener", func() {
|
||||
f := &fakeDiarStream{}
|
||||
cfg := &cDiarSegmentationConfig{MinDurationSec: 0.5}
|
||||
_, err := diarizePCM(f.opener(), []float32{1}, 16000, cfg)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(f.cfg).To(BeIdenticalTo(cfg))
|
||||
})
|
||||
|
||||
It("closes the stream when the segment query fails", func() {
|
||||
f := &fakeDiarStream{countErr: errors.New("boom")}
|
||||
_, err := diarizePCM(f.opener(), []float32{1}, 16000, nil)
|
||||
Expect(err).To(MatchError(ContainSubstring("boom")))
|
||||
Expect(f.closed).To(Equal(1))
|
||||
})
|
||||
|
||||
// The pipeline carries no ASR, so text and language stay empty whatever the
|
||||
// caller asked for.
|
||||
It("leaves the transcript fields empty", func() {
|
||||
f := &fakeDiarStream{segs: []cDiarSegment{{StartTime: 0, EndTime: 1, Speaker: 1}}}
|
||||
res, err := diarizePCM(f.opener(), []float32{1}, 16000, nil)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(res.GetLanguage()).To(BeEmpty())
|
||||
Expect(res.Segments[0].GetText()).To(BeEmpty())
|
||||
})
|
||||
})
|
||||
|
||||
// countObserver records how many size queries had run by the time finish was
|
||||
// called, so the ordering can be asserted without reaching into diarizePCM.
|
||||
type countObserver struct {
|
||||
*fakeDiarStream
|
||||
seen *int
|
||||
}
|
||||
|
||||
func (c *countObserver) finish() error {
|
||||
*c.seen = c.counts + 1 // finish must run before the first query
|
||||
return c.fakeDiarStream.finish()
|
||||
}
|
||||
|
||||
var _ = Describe("distinctSpeakers", func() {
|
||||
It("counts labels, not segments", func() {
|
||||
segs := []*pb.DiarizeSegment{
|
||||
{Speaker: "1"}, {Speaker: "2"}, {Speaker: "1"}, {Speaker: "2"}, {Speaker: "1"},
|
||||
}
|
||||
Expect(distinctSpeakers(segs)).To(Equal(int32(2)))
|
||||
})
|
||||
|
||||
It("counts a single-speaker recording as one", func() {
|
||||
segs := []*pb.DiarizeSegment{{Speaker: "1"}, {Speaker: "1"}, {Speaker: "1"}}
|
||||
Expect(distinctSpeakers(segs)).To(Equal(int32(1)))
|
||||
})
|
||||
|
||||
// Four segments over three labels, not three over three: with the segment
|
||||
// count and the label count equal, a `return len(segs)` would satisfy this
|
||||
// spec and it would assert nothing.
|
||||
It("counts every distinct label once", func() {
|
||||
segs := []*pb.DiarizeSegment{{Speaker: "1"}, {Speaker: "2"}, {Speaker: "3"}, {Speaker: "2"}}
|
||||
Expect(distinctSpeakers(segs)).To(Equal(int32(3)))
|
||||
})
|
||||
|
||||
It("is zero with no segments", func() {
|
||||
Expect(distinctSpeakers(nil)).To(BeZero())
|
||||
})
|
||||
|
||||
// The response field is documented as the count of distinct labels in
|
||||
// `segments`, which is not the model's capacity: Sortformer v2 can label
|
||||
// four speakers whatever the clip actually contains.
|
||||
It("reports what the segments contain, not the model capacity", func() {
|
||||
f := &fakeDiarStream{segs: []cDiarSegment{
|
||||
{StartTime: 0, EndTime: 1, Speaker: 1},
|
||||
{StartTime: 1, EndTime: 2, Speaker: 2},
|
||||
{StartTime: 2, EndTime: 3, Speaker: 1},
|
||||
}}
|
||||
res, err := diarizePCM(f.opener(), []float32{1}, 16000, nil)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(res.GetNumSpeakers()).To(Equal(int32(2)))
|
||||
})
|
||||
|
||||
It("reports no speakers when the runtime found no segments", func() {
|
||||
f := &fakeDiarStream{}
|
||||
res, err := diarizePCM(f.opener(), []float32{1}, 16000, nil)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(res.Segments).To(BeEmpty())
|
||||
Expect(res.GetNumSpeakers()).To(BeZero())
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("collectSegments", func() {
|
||||
It("skips the fill entirely when there is nothing to collect", func() {
|
||||
f := &fakeDiarStream{}
|
||||
segs, err := collectSegments(f)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(segs).To(BeEmpty())
|
||||
Expect(f.counts).To(Equal(1))
|
||||
Expect(f.fills).To(BeZero(), "a zero count must not be followed by a fill")
|
||||
})
|
||||
|
||||
It("sizes the buffer from the query and fills it", func() {
|
||||
f := &fakeDiarStream{segs: []cDiarSegment{
|
||||
{StartTime: 0, EndTime: 1, Speaker: 1},
|
||||
{StartTime: 1, EndTime: 2, Speaker: 2},
|
||||
}}
|
||||
segs, err := collectSegments(f)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(segs).To(HaveLen(2))
|
||||
Expect(segs[1].Speaker).To(Equal(int32(2)))
|
||||
Expect(f.counts).To(Equal(1))
|
||||
Expect(f.fills).To(Equal(1))
|
||||
})
|
||||
|
||||
// nemo_speech_diar_segments writes *count and only then rejects a buffer
|
||||
// that is too small, so the rejected call still reports the size to retry
|
||||
// with. Truncating instead would silently drop turns.
|
||||
It("grows the buffer and retries when the count outran the query", func() {
|
||||
f := &fakeDiarStream{
|
||||
segs: []cDiarSegment{
|
||||
{StartTime: 0, EndTime: 1, Speaker: 1},
|
||||
{StartTime: 1, EndTime: 2, Speaker: 2},
|
||||
{StartTime: 2, EndTime: 3, Speaker: 1},
|
||||
},
|
||||
// The query saw two, the fill found three and rejected the buffer.
|
||||
queryCount: 2,
|
||||
growTo: 3,
|
||||
fillErrs: []error{errors.New("capacity too small (need 3)")},
|
||||
}
|
||||
segs, err := collectSegments(f)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(segs).To(HaveLen(3))
|
||||
Expect(f.fills).To(Equal(2))
|
||||
})
|
||||
|
||||
It("propagates a failure that is not about capacity", func() {
|
||||
f := &fakeDiarStream{
|
||||
segs: []cDiarSegment{{StartTime: 0, EndTime: 1, Speaker: 1}},
|
||||
fillErrs: []error{errors.New("boom")},
|
||||
}
|
||||
_, err := collectSegments(f)
|
||||
Expect(err).To(MatchError(ContainSubstring("boom")))
|
||||
Expect(f.fills).To(Equal(1), "a non-capacity failure must not be retried")
|
||||
})
|
||||
|
||||
// make() panics on a length it cannot satisfy, and a panic in an RPC
|
||||
// handler kills the backend process. A count that could only come from an
|
||||
// uninitialised or corrupted size_t must fail the request instead.
|
||||
It("refuses an implausible count rather than trying to allocate it", func() {
|
||||
f := &fakeDiarStream{queryCount: maxDiarSegments + 1}
|
||||
_, err := collectSegments(f)
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(status.Code(err)).To(Equal(codes.Internal))
|
||||
Expect(err.Error()).To(ContainSubstring("ceiling"))
|
||||
Expect(f.fills).To(BeZero(), "nothing must be allocated or filled for a bad count")
|
||||
})
|
||||
|
||||
It("still accepts a count right at the ceiling", func() {
|
||||
// Only the guard is under test, so the fill is scripted to report zero
|
||||
// rather than actually materialising a hundred megabytes of segments.
|
||||
f := &fakeDiarStream{queryCount: maxDiarSegments}
|
||||
segs, err := collectSegments(f)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(segs).To(BeEmpty())
|
||||
Expect(f.fills).To(Equal(1))
|
||||
})
|
||||
|
||||
It("propagates a failed size query", func() {
|
||||
f := &fakeDiarStream{countErr: errors.New("no stream")}
|
||||
_, err := collectSegments(f)
|
||||
Expect(err).To(MatchError(ContainSubstring("no stream")))
|
||||
Expect(f.fills).To(BeZero())
|
||||
})
|
||||
|
||||
// A runtime whose count grew on every attempt would otherwise loop forever
|
||||
// holding engineMu, which blocks the unload too.
|
||||
It("gives up rather than retrying forever", func() {
|
||||
f := &fakeDiarStream{segs: []cDiarSegment{{StartTime: 0, EndTime: 1, Speaker: 1}}}
|
||||
g := &alwaysGrowing{fakeDiarStream: f}
|
||||
_, err := collectSegments(g)
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(status.Code(err)).To(Equal(codes.Internal))
|
||||
Expect(f.fills).To(Equal(diarSegmentsMaxAttempts))
|
||||
})
|
||||
})
|
||||
|
||||
// alwaysGrowing reports a bigger count on every fill, which is the pathological
|
||||
// case the attempt bound exists for.
|
||||
type alwaysGrowing struct {
|
||||
*fakeDiarStream
|
||||
n uint64
|
||||
}
|
||||
|
||||
func (a *alwaysGrowing) fillSegments([]cDiarSegment) (uint64, error) {
|
||||
a.n += 10
|
||||
a.fills++
|
||||
return a.n, errors.New("capacity too small")
|
||||
}
|
||||
|
||||
// The six frame counts are the one part of the create config that no other
|
||||
// check in the tree can see. The layout assertions pin the struct's shape, and
|
||||
// a wrong VALUE keeps that shape exactly, so without these specs deleting a
|
||||
// sentinel is invisible: c_api.cpp applies left_context_frames at >= 0, so a
|
||||
// dropped -1 there silently pins the model's left context to zero.
|
||||
var _ = Describe("diarModelConfig", func() {
|
||||
It("declares its own size so the runtime accepts the fields", func() {
|
||||
Expect(diarModelConfig(0, -1).Size).To(Equal(unsafe.Sizeof(cDiarModelConfig{})))
|
||||
})
|
||||
|
||||
It("carries the model path and the configured device", func() {
|
||||
cfg := diarModelConfig(0xDEADBEEF, 2)
|
||||
Expect(cfg.ModelPath).To(Equal(uintptr(0xDEADBEEF)))
|
||||
Expect(cfg.GPU).To(Equal(int32(2)))
|
||||
})
|
||||
|
||||
It("passes the CPU sentinel through untouched", func() {
|
||||
Expect(diarModelConfig(0, -1).GPU).To(Equal(int32(-1)))
|
||||
})
|
||||
|
||||
// Asserted field by field rather than as a whole struct so a failure names
|
||||
// the sentinel that went missing.
|
||||
It("leaves every frame-geometry override at the negative sentinel", func() {
|
||||
cfg := diarModelConfig(0, -1)
|
||||
Expect(cfg.ChunkFrames).To(Equal(int32(-1)), "chunk_frames")
|
||||
Expect(cfg.RightContextFrames).To(Equal(int32(-1)), "right_context_frames")
|
||||
Expect(cfg.FIFOFrames).To(Equal(int32(-1)), "fifo_frames")
|
||||
Expect(cfg.SpkcacheFrames).To(Equal(int32(-1)), "spkcache_frames")
|
||||
Expect(cfg.UpdatePeriodFrames).To(Equal(int32(-1)), "update_period_frames")
|
||||
|
||||
// Called out on its own because it is the only one of the six the
|
||||
// runtime applies at >= 0: zero here is a valid explicit left context,
|
||||
// not "unset", so this is the field a dropped sentinel actually breaks.
|
||||
Expect(cfg.LeftContextFrames).To(Equal(int32(-1)), "left_context_frames")
|
||||
Expect(cfg.LeftContextFrames).To(BeNumerically("<", 0),
|
||||
"left_context_frames is applied at >= 0, so a non-negative value pins the geometry")
|
||||
})
|
||||
|
||||
// The preset selects the streaming geometry wholesale, so it has to stay
|
||||
// NULL until there is a model to verify a different one against.
|
||||
It("leaves the preset unset", func() {
|
||||
Expect(diarModelConfig(0, -1).Preset).To(BeZero())
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("segmentationConfig", func() {
|
||||
// A request that set nothing must stay NULL on the C side: diar.h documents
|
||||
// NULL as "library defaults", and those defaults are NeMo's callhome-tuned
|
||||
// values for this checkpoint rather than zeros.
|
||||
It("is absent when the request asked for no postprocessing", func() {
|
||||
Expect(segmentationConfig(&pb.DiarizeRequest{})).To(BeNil())
|
||||
})
|
||||
|
||||
It("maps min_duration_on onto the minimum segment duration", func() {
|
||||
cfg := segmentationConfig(&pb.DiarizeRequest{MinDurationOn: 0.4})
|
||||
Expect(cfg).ToNot(BeNil())
|
||||
Expect(cfg.MinDurationSec).To(BeNumerically("~", 0.4, 1e-6))
|
||||
Expect(cfg.MinGapSec).To(BeZero())
|
||||
})
|
||||
|
||||
It("maps min_duration_off onto the gap fill", func() {
|
||||
cfg := segmentationConfig(&pb.DiarizeRequest{MinDurationOff: 0.25})
|
||||
Expect(cfg).ToNot(BeNil())
|
||||
Expect(cfg.MinGapSec).To(BeNumerically("~", 0.25, 1e-6))
|
||||
Expect(cfg.MinDurationSec).To(BeZero())
|
||||
})
|
||||
|
||||
It("declares its own size so the runtime accepts the fields", func() {
|
||||
cfg := segmentationConfig(&pb.DiarizeRequest{MinDurationOn: 0.4})
|
||||
Expect(cfg.Size).To(Equal(sizeofDiarSegmentationConfig()))
|
||||
})
|
||||
|
||||
// The onset/offset hysteresis is not a clustering threshold and Sortformer
|
||||
// has no clustering stage at all, so mapping one onto the other would be an
|
||||
// invented equivalence. It has to stay unset.
|
||||
It("ignores fields this pipeline has no equivalent for", func() {
|
||||
Expect(segmentationConfig(&pb.DiarizeRequest{
|
||||
NumSpeakers: 2,
|
||||
MinSpeakers: 1,
|
||||
MaxSpeakers: 4,
|
||||
ClusteringThreshold: 0.7,
|
||||
IncludeText: true,
|
||||
Threads: 8,
|
||||
})).To(BeNil())
|
||||
})
|
||||
|
||||
It("ignores non-positive values, which the runtime reads as unset", func() {
|
||||
Expect(segmentationConfig(&pb.DiarizeRequest{MinDurationOn: -1, MinDurationOff: 0})).To(BeNil())
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("unsupportedRequestFields", func() {
|
||||
It("is empty for a request this backend can honour in full", func() {
|
||||
Expect(unsupportedRequestFields(&pb.DiarizeRequest{
|
||||
Dst: "/tmp/a.wav",
|
||||
MinDurationOn: 0.4,
|
||||
MinDurationOff: 0.2,
|
||||
})).To(BeEmpty())
|
||||
})
|
||||
|
||||
It("names every field it had to drop", func() {
|
||||
Expect(unsupportedRequestFields(&pb.DiarizeRequest{
|
||||
NumSpeakers: 2,
|
||||
MinSpeakers: 1,
|
||||
MaxSpeakers: 4,
|
||||
ClusteringThreshold: 0.7,
|
||||
IncludeText: true,
|
||||
Threads: 8,
|
||||
})).To(ConsistOf(
|
||||
"num_speakers", "min_speakers", "max_speakers",
|
||||
"clustering_threshold", "include_text", "threads",
|
||||
))
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,116 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
|
||||
gguf "github.com/gpustack/gguf-parser-go"
|
||||
)
|
||||
|
||||
// auxOnlyArchitectures are converted NeMo components that attach to a primary
|
||||
// model but are never loadable on their own. Pointing a model config at one is
|
||||
// a configuration mistake worth naming explicitly.
|
||||
var auxOnlyArchitectures = map[string]string{
|
||||
"nemo-nano-codec": "a TTS codec, set it with the codec_model option on a magpietts model",
|
||||
"vad": "a VAD model, set it with the vad_model option on an asr model",
|
||||
"pnc": "a punctuation model, set it with the pnc_model option on an asr model",
|
||||
}
|
||||
|
||||
// familyFor maps a GGUF general.architecture value to a model family.
|
||||
//
|
||||
// Unknown architectures resolve to NMT rather than an error: NMT GGUFs come
|
||||
// from llama.cpp's converter and carry an ordinary LLM architecture, so there
|
||||
// is no NeMo-specific string to match. The user selected this backend
|
||||
// explicitly, which is the signal that the model is meant for it.
|
||||
func familyFor(arch string) (family, error) {
|
||||
if reason, ok := auxOnlyArchitectures[arch]; ok {
|
||||
return familyUnknown, fmt.Errorf(
|
||||
"nemo-speech-cpp: %q is %s, not a model that can be loaded directly", arch, reason)
|
||||
}
|
||||
switch arch {
|
||||
case "asr":
|
||||
return familyASR, nil
|
||||
case "sortformer":
|
||||
return familyDiarization, nil
|
||||
case "magpietts":
|
||||
return familyTTS, nil
|
||||
}
|
||||
return familyNMT, nil
|
||||
}
|
||||
|
||||
// ggufArchitecture reads general.architecture from a GGUF file.
|
||||
func ggufArchitecture(path string) (string, error) {
|
||||
f, err := gguf.ParseGGUFFile(path, gguf.UseMMap(), gguf.SkipLargeMetadata())
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("nemo-speech-cpp: parse gguf %q: %w", path, err)
|
||||
}
|
||||
kv, found := f.Header.MetadataKV.Index([]string{"general.architecture"})
|
||||
if found == 0 {
|
||||
return "", fmt.Errorf("nemo-speech-cpp: %q has no general.architecture key", path)
|
||||
}
|
||||
arch := kv["general.architecture"]
|
||||
// ValueString panics on a mistyped key, and a hand-written or half-converted
|
||||
// GGUF is exactly where that happens. This function is the load-time guard;
|
||||
// it reports, it does not take the process down.
|
||||
if arch.ValueType != gguf.GGUFMetadataValueTypeString {
|
||||
return "", fmt.Errorf(
|
||||
"nemo-speech-cpp: %q has a non-string general.architecture (type %v)", path, arch.ValueType)
|
||||
}
|
||||
return arch.ValueString(), nil
|
||||
}
|
||||
|
||||
// discoverTTSAssets fills in codecModel and tokenizerDir when they were not set
|
||||
// explicitly, by scanning the primary GGUF's own directory.
|
||||
//
|
||||
// A missing asset is a hard error rather than a warning: the runtime would
|
||||
// otherwise load and emit garbage audio, which surfaces far from the cause.
|
||||
func discoverTTSAssets(primaryGGUF string, o *loadOptions) error {
|
||||
dir := filepath.Dir(primaryGGUF)
|
||||
|
||||
if o.codecModel == "" {
|
||||
entries, err := os.ReadDir(dir)
|
||||
if err != nil {
|
||||
return fmt.Errorf("nemo-speech-cpp: scan %q for a codec model: %w", dir, err)
|
||||
}
|
||||
for _, e := range entries {
|
||||
if e.IsDir() {
|
||||
continue
|
||||
}
|
||||
name := e.Name()
|
||||
candidate := filepath.Join(dir, name)
|
||||
// Skip the primary model itself: a file called nanocodec-magpie.gguf
|
||||
// would otherwise be selected as its own codec. Compare basenames,
|
||||
// because candidate is Cleaned by filepath.Join while primaryGGUF
|
||||
// arrives as the caller wrote it, so "/models//magpie.gguf" would
|
||||
// slip past a whole-path equality.
|
||||
if name == filepath.Base(primaryGGUF) {
|
||||
continue
|
||||
}
|
||||
if strings.Contains(strings.ToLower(name), "nanocodec") ||
|
||||
strings.Contains(strings.ToLower(name), "nano-codec") {
|
||||
o.codecModel = candidate
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
if o.codecModel == "" {
|
||||
return fmt.Errorf(
|
||||
"nemo-speech-cpp: no NanoCodec GGUF found next to %q, set the codec_model option",
|
||||
primaryGGUF)
|
||||
}
|
||||
|
||||
if o.tokenizerDir == "" {
|
||||
candidate := filepath.Join(dir, "extracted")
|
||||
if st, err := os.Stat(candidate); err == nil && st.IsDir() {
|
||||
o.tokenizerDir = candidate
|
||||
}
|
||||
}
|
||||
if o.tokenizerDir == "" {
|
||||
return fmt.Errorf(
|
||||
"nemo-speech-cpp: no tokenizer directory found next to %q, set the tokenizer_dir option",
|
||||
primaryGGUF)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,168 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"os"
|
||||
"path/filepath"
|
||||
|
||||
gguf "github.com/gpustack/gguf-parser-go"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
var _ = Describe("familyFor", func() {
|
||||
It("maps the NeMo architectures to their families", func() {
|
||||
for arch, want := range map[string]family{
|
||||
"asr": familyASR,
|
||||
"sortformer": familyDiarization,
|
||||
"magpietts": familyTTS,
|
||||
} {
|
||||
got, err := familyFor(arch)
|
||||
Expect(err).ToNot(HaveOccurred(), "arch %q", arch)
|
||||
Expect(got).To(Equal(want), "arch %q", arch)
|
||||
}
|
||||
})
|
||||
|
||||
It("treats an unknown architecture as NMT", func() {
|
||||
// NMT GGUFs are produced by llama.cpp's converter, so they carry an LLM
|
||||
// architecture such as qwen3 rather than a NeMo-specific string.
|
||||
got, err := familyFor("qwen3")
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(got).To(Equal(familyNMT))
|
||||
})
|
||||
|
||||
It("rejects an auxiliary-only architecture as a primary model", func() {
|
||||
for _, arch := range []string{"nemo-nano-codec", "vad", "pnc"} {
|
||||
_, err := familyFor(arch)
|
||||
Expect(err).To(HaveOccurred(), "arch %q", arch)
|
||||
Expect(err.Error()).To(ContainSubstring(arch))
|
||||
}
|
||||
})
|
||||
})
|
||||
|
||||
// ggufWithArchValue builds a minimal GGUF v3 carrying general.architecture as
|
||||
// its single metadata entry, with the caller's value type and encoded value.
|
||||
func ggufWithArchValue(valueType gguf.GGUFMetadataValueType, value []byte) []byte {
|
||||
const key = "general.architecture"
|
||||
|
||||
var b []byte
|
||||
b = append(b, 'G', 'G', 'U', 'F')
|
||||
b = binary.LittleEndian.AppendUint32(b, 3) // version
|
||||
b = binary.LittleEndian.AppendUint64(b, 0) // tensor count
|
||||
b = binary.LittleEndian.AppendUint64(b, 1) // metadata kv count
|
||||
b = binary.LittleEndian.AppendUint64(b, uint64(len(key)))
|
||||
b = append(b, key...)
|
||||
b = binary.LittleEndian.AppendUint32(b, uint32(valueType))
|
||||
return append(b, value...)
|
||||
}
|
||||
|
||||
// writeGGUFWithUint32Arch writes a minimal GGUF v3 whose single metadata entry
|
||||
// is general.architecture typed UINT32 rather than STRING. Handwritten and
|
||||
// half-converted files really do carry mistyped keys, and the parser hands them
|
||||
// back rather than rejecting them.
|
||||
func writeGGUFWithUint32Arch(path string) {
|
||||
b := ggufWithArchValue(gguf.GGUFMetadataValueTypeUint32, binary.LittleEndian.AppendUint32(nil, 7))
|
||||
ExpectWithOffset(1, os.WriteFile(path, b, 0o600)).To(Succeed())
|
||||
}
|
||||
|
||||
// writeGGUFWithArch writes a minimal GGUF v3 that parses cleanly and reports
|
||||
// arch as its general.architecture. It is the only way to reach the code past
|
||||
// ggufArchitecture in a test, since there are no real NeMo GGUFs to point at.
|
||||
func writeGGUFWithArch(path, arch string) {
|
||||
v := binary.LittleEndian.AppendUint64(nil, uint64(len(arch)))
|
||||
v = append(v, arch...)
|
||||
ExpectWithOffset(1, os.WriteFile(path, ggufWithArchValue(gguf.GGUFMetadataValueTypeString, v), 0o600)).To(Succeed())
|
||||
}
|
||||
|
||||
var _ = Describe("ggufArchitecture", func() {
|
||||
It("returns an error rather than panicking on a file that is not a GGUF", func() {
|
||||
p := filepath.Join(GinkgoT().TempDir(), "not-a-model.gguf")
|
||||
Expect(os.WriteFile(p, []byte("definitely not a gguf header"), 0o600)).To(Succeed())
|
||||
|
||||
_, err := ggufArchitecture(p)
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(err.Error()).To(ContainSubstring(p))
|
||||
})
|
||||
|
||||
It("returns an error rather than panicking when general.architecture is not a string", func() {
|
||||
p := filepath.Join(GinkgoT().TempDir(), "mistyped-arch.gguf")
|
||||
writeGGUFWithUint32Arch(p)
|
||||
|
||||
arch, err := ggufArchitecture(p)
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(arch).To(BeEmpty())
|
||||
Expect(err.Error()).To(ContainSubstring("general.architecture"))
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("discoverTTSAssets", func() {
|
||||
var dir string
|
||||
|
||||
BeforeEach(func() {
|
||||
dir = GinkgoT().TempDir()
|
||||
})
|
||||
|
||||
write := func(name string) string {
|
||||
p := filepath.Join(dir, name)
|
||||
Expect(os.WriteFile(p, []byte("x"), 0o600)).To(Succeed())
|
||||
return p
|
||||
}
|
||||
|
||||
It("finds a sibling nanocodec gguf and extracted dir", func() {
|
||||
primary := write("magpie.f16.gguf")
|
||||
codec := write("nemo-nano-codec-22khz.f16.gguf")
|
||||
Expect(os.Mkdir(filepath.Join(dir, "extracted"), 0o755)).To(Succeed())
|
||||
|
||||
o := loadOptions{}
|
||||
Expect(discoverTTSAssets(primary, &o)).To(Succeed())
|
||||
Expect(o.codecModel).To(Equal(codec))
|
||||
Expect(o.tokenizerDir).To(Equal(filepath.Join(dir, "extracted")))
|
||||
})
|
||||
|
||||
It("never selects the primary gguf as its own codec", func() {
|
||||
// A file named so it would match a naive *.gguf scan.
|
||||
primary := write("nanocodec-magpie.gguf")
|
||||
Expect(os.Mkdir(filepath.Join(dir, "extracted"), 0o755)).To(Succeed())
|
||||
|
||||
o := loadOptions{}
|
||||
err := discoverTTSAssets(primary, &o)
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(err.Error()).To(ContainSubstring("codec_model"))
|
||||
})
|
||||
|
||||
It("never selects the primary gguf as its own codec through an uncleaned path", func() {
|
||||
// LocalAI joins the model directory and the model name itself, so a
|
||||
// trailing separator on ModelPath produces a doubled slash here. The
|
||||
// self-codec guard has to survive that.
|
||||
write("nanocodec-magpie.gguf")
|
||||
primary := dir + "//nanocodec-magpie.gguf"
|
||||
Expect(os.Mkdir(filepath.Join(dir, "extracted"), 0o755)).To(Succeed())
|
||||
|
||||
o := loadOptions{}
|
||||
err := discoverTTSAssets(primary, &o)
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(o.codecModel).To(BeEmpty())
|
||||
Expect(err.Error()).To(ContainSubstring("codec_model"))
|
||||
})
|
||||
|
||||
It("does not overwrite explicitly configured paths", func() {
|
||||
primary := write("magpie.f16.gguf")
|
||||
write("nemo-nano-codec.gguf")
|
||||
Expect(os.Mkdir(filepath.Join(dir, "extracted"), 0o755)).To(Succeed())
|
||||
|
||||
o := loadOptions{codecModel: "/explicit/codec.gguf", tokenizerDir: "/explicit/tok"}
|
||||
Expect(discoverTTSAssets(primary, &o)).To(Succeed())
|
||||
Expect(o.codecModel).To(Equal("/explicit/codec.gguf"))
|
||||
Expect(o.tokenizerDir).To(Equal("/explicit/tok"))
|
||||
})
|
||||
|
||||
It("names the missing option key when the tokenizer dir cannot be found", func() {
|
||||
primary := write("magpie.f16.gguf")
|
||||
write("nemo-nano-codec.gguf")
|
||||
|
||||
o := loadOptions{}
|
||||
err := discoverTTSAssets(primary, &o)
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(err.Error()).To(ContainSubstring("tokenizer_dir"))
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,76 @@
|
||||
package main
|
||||
|
||||
// Started internally by LocalAI, one gRPC server per loaded model.
|
||||
//
|
||||
// Binds NVIDIA NeMo-Speech.cpp through purego. The runtime splits its C ABI
|
||||
// across three shared objects: asr (which also exports the diarization
|
||||
// symbols), tts, and nmt. Library names can be overridden with
|
||||
// NEMO_SPEECH_ASR_LIBRARY / _TTS_LIBRARY / _NMT_LIBRARY, mirroring the
|
||||
// PARAKEET_LIBRARY convention in the sibling backends.
|
||||
//
|
||||
// The naming is asymmetric on purpose: upstream links a dedicated
|
||||
// libnemo_speech_asr_c / libnemo_speech_nmt_c around a private C++ core, but
|
||||
// compiles the TTS c_api straight into libnemo_speech_tts and only aliases the
|
||||
// nemo_speech_tts_c CMake target, so there is no libnemo_speech_tts_c on disk.
|
||||
import (
|
||||
"flag"
|
||||
"fmt"
|
||||
"os"
|
||||
"runtime"
|
||||
|
||||
"github.com/ebitengine/purego"
|
||||
grpc "github.com/mudler/LocalAI/pkg/grpc"
|
||||
)
|
||||
|
||||
var addr = flag.String("addr", "localhost:50051", "the address to connect to")
|
||||
|
||||
// libSuffix is the platform's shared-object extension.
|
||||
func libSuffix() string {
|
||||
if runtime.GOOS == "darwin" {
|
||||
return ".dylib"
|
||||
}
|
||||
return ".so"
|
||||
}
|
||||
|
||||
// libraryName resolves an override env var, falling back to the platform name.
|
||||
func libraryName(envVar, base string) string {
|
||||
if v := os.Getenv(envVar); v != "" {
|
||||
return v
|
||||
}
|
||||
return base + libSuffix()
|
||||
}
|
||||
|
||||
func main() {
|
||||
flag.Parse()
|
||||
|
||||
if err := openLibraries(); err != nil {
|
||||
panic(err)
|
||||
}
|
||||
|
||||
if err := grpc.StartServer(*addr, &NemoSpeech{}); err != nil {
|
||||
panic(err)
|
||||
}
|
||||
}
|
||||
|
||||
// openLibraries dlopens the three C ABI shared objects. All three are opened
|
||||
// eagerly so a packaging mistake fails at startup with a clear message rather
|
||||
// than at first inference of one particular family.
|
||||
func openLibraries() error {
|
||||
for _, l := range []struct {
|
||||
env string
|
||||
base string
|
||||
dst *uintptr
|
||||
}{
|
||||
{"NEMO_SPEECH_ASR_LIBRARY", "libnemo_speech_asr_c", &asrLib},
|
||||
{"NEMO_SPEECH_TTS_LIBRARY", "libnemo_speech_tts", &ttsLib},
|
||||
{"NEMO_SPEECH_NMT_LIBRARY", "libnemo_speech_nmt_c", &nmtLib},
|
||||
} {
|
||||
name := libraryName(l.env, l.base)
|
||||
h, err := purego.Dlopen(name, purego.RTLD_NOW|purego.RTLD_GLOBAL)
|
||||
if err != nil {
|
||||
return fmt.Errorf("nemo-speech-cpp: dlopen %q: %w", name, err)
|
||||
}
|
||||
*l.dst = h
|
||||
}
|
||||
return registerSymbols()
|
||||
}
|
||||
@@ -0,0 +1,273 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"runtime"
|
||||
"sync"
|
||||
"unsafe"
|
||||
|
||||
"github.com/mudler/LocalAI/pkg/grpc/base"
|
||||
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
|
||||
"github.com/mudler/xlog"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
)
|
||||
|
||||
// family is the model family selected at load time from the GGUF architecture.
|
||||
type family int
|
||||
|
||||
const (
|
||||
familyUnknown family = iota
|
||||
familyASR
|
||||
familyDiarization
|
||||
familyTTS
|
||||
familyNMT
|
||||
)
|
||||
|
||||
func (f family) String() string {
|
||||
switch f {
|
||||
case familyASR:
|
||||
return "asr"
|
||||
case familyDiarization:
|
||||
return "diarization"
|
||||
case familyTTS:
|
||||
return "tts"
|
||||
case familyNMT:
|
||||
return "nmt"
|
||||
}
|
||||
return "unknown"
|
||||
}
|
||||
|
||||
// NemoSpeech is one loaded model. Exactly one of the handles is non-zero,
|
||||
// matching fam.
|
||||
type NemoSpeech struct {
|
||||
base.SingleThread
|
||||
|
||||
fam family
|
||||
opts loadOptions
|
||||
|
||||
// engineMu guards fam and the handles, and serializes calls into the C
|
||||
// runtime for this model. Its participants today are withEngine and Free;
|
||||
// the per-family RPCs in Tasks 6 to 9 join it by routing through withEngine.
|
||||
engineMu sync.Mutex
|
||||
|
||||
// synth and nmt are shortened rather than spelled out: synthesizer and
|
||||
// translator are the names of the two RPC-side interfaces those handles are
|
||||
// wrapped in (tts.go, nmt.go), and a field sharing a name with an interface in
|
||||
// the same package makes every construction site read as a conversion.
|
||||
recognizer uintptr
|
||||
diarizer uintptr
|
||||
synth uintptr
|
||||
nmt uintptr
|
||||
}
|
||||
|
||||
// cstr allocates a NUL-terminated C string and returns its pointer plus a
|
||||
// release function. The empty string maps to a null pointer because the C API
|
||||
// treats NULL and "" as equivalent for every optional field.
|
||||
//
|
||||
// The address leaves the Go type system as a uintptr, which the collector does
|
||||
// not trace, so the bytes are pinned for as long as C may read them. Pinning is
|
||||
// the only mechanism with a documented guarantee here: the config structs hold
|
||||
// raw addresses, and an unpinned Go allocation is free to be collected (and, in
|
||||
// principle, moved) the moment its last traced reference dies.
|
||||
//
|
||||
// The returned pointer is for C only, and the direction is one-way. Converting
|
||||
// it back to an unsafe.Pointer to read the bytes from Go is checked by checkptr
|
||||
// (which -race turns on) and kills the process with
|
||||
//
|
||||
// fatal error: checkptr: pointer arithmetic result points to invalid allocation
|
||||
//
|
||||
// as soon as the address lands inside a Go allocation, which is exactly what
|
||||
// this produces. C reading it is fine because C is not instrumented; Go reading
|
||||
// it back is not.
|
||||
//
|
||||
// The caller MUST defer the release function immediately, in the same statement
|
||||
// that takes the pointer. Dropping it leaks the pin, which the runtime reports
|
||||
// at the next collection as:
|
||||
//
|
||||
// runtime.Pinner: found leaking pinned pointer; forgot to call Unpin()?
|
||||
//
|
||||
// That is loud and wrong-looking on purpose: the alternative failure mode is C
|
||||
// reading freed memory, which shows up as rare corruption with no trace back
|
||||
// to here.
|
||||
func cstr(s string) (uintptr, func()) {
|
||||
if s == "" {
|
||||
return 0, func() {}
|
||||
}
|
||||
b := append([]byte(s), 0)
|
||||
pin := new(runtime.Pinner)
|
||||
pin.Pin(&b[0])
|
||||
// #nosec G103 -- b is non-empty (s != "" above) and &b[0] is pinned on the
|
||||
// previous line, so the address C receives cannot be collected or moved
|
||||
// until the returned release runs. One-way by construction: the doc comment
|
||||
// above forbids converting this uintptr back, which is what keeps checkptr
|
||||
// (and therefore -race) out of it.
|
||||
return uintptr(unsafe.Pointer(&b[0])), func() {
|
||||
if pin == nil {
|
||||
return
|
||||
}
|
||||
pin.Unpin()
|
||||
pin = nil
|
||||
}
|
||||
}
|
||||
|
||||
// There is deliberately no inverse of cstr in this package. Every C entry point
|
||||
// that returns a string is bound in abi.go with a Go `string` return, which
|
||||
// purego converts from the char* itself, so a hand-rolled reader would have no
|
||||
// production caller and would exist only as an unsafe helper waiting to be
|
||||
// pointed at the wrong kind of address. Reach for purego's conversion instead;
|
||||
// if a future symbol genuinely needs the raw char* (to tell NULL from ""), bind
|
||||
// it as uintptr at that call site, where the ownership can be reasoned about.
|
||||
|
||||
// requireFamily gates an RPC on the family selected at load time. Returning
|
||||
// Unimplemented rather than a nil dereference means a misconfigured model YAML
|
||||
// produces a message a user can act on.
|
||||
//
|
||||
// Callers must already hold engineMu: Free writes n.fam under it, so an
|
||||
// unlocked read here is a data race. Use withEngine rather than calling this
|
||||
// directly.
|
||||
func (n *NemoSpeech) requireFamily(want family) error {
|
||||
if n.fam != want {
|
||||
return status.Errorf(codes.Unimplemented,
|
||||
"nemo-speech-cpp: this model was loaded as %s, not %s", n.fam, want)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// withEngine runs fn holding engineMu, having first checked the family.
|
||||
//
|
||||
// Every RPC must go through this rather than calling requireFamily on its own.
|
||||
// pkg/grpc/server.go takes the backend lock around each RPC but calls Free
|
||||
// without it, so a teardown can land mid-request. Checking the family and then
|
||||
// making the C calls that trust it under two separate acquisitions leaves a
|
||||
// window in which Free destroys the handle, and the request goes on to use a
|
||||
// zeroed one.
|
||||
func (n *NemoSpeech) withEngine(want family, fn func() error) error {
|
||||
n.engineMu.Lock()
|
||||
defer n.engineMu.Unlock()
|
||||
|
||||
if err := n.requireFamily(want); err != nil {
|
||||
return err
|
||||
}
|
||||
return fn()
|
||||
}
|
||||
|
||||
func (n *NemoSpeech) Load(opts *pb.ModelOptions) error {
|
||||
modelFile := opts.GetModelFile()
|
||||
if modelFile == "" {
|
||||
return errors.New("nemo-speech-cpp: ModelFile is required")
|
||||
}
|
||||
|
||||
// Free writes fam and the handles under engineMu and runs without the
|
||||
// backend lock that serialises the RPCs (pkg/grpc/server.go), so the
|
||||
// load-side writes to those same fields need the same protection: without
|
||||
// it this is the write-side half of the race withEngine closed on the read
|
||||
// side. n.opts is in here too, since the loaders read it.
|
||||
//
|
||||
// The loaders called below must NOT take engineMu themselves; sync.Mutex is
|
||||
// not reentrant and this is why.
|
||||
n.engineMu.Lock()
|
||||
defer n.engineMu.Unlock()
|
||||
|
||||
n.opts = parseOptions(opts.GetOptions(), opts.GetModelPath())
|
||||
|
||||
arch, err := ggufArchitecture(modelFile)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
fam, err := familyFor(arch)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
xlog.Info("nemo-speech-cpp: loading model", "arch", arch, "family", fam.String())
|
||||
|
||||
// fam is committed only once the family-specific loader has succeeded.
|
||||
// requireFamily is the gate every RPC goes through, so a half-loaded model
|
||||
// that kept its family would route requests at a handle that was never
|
||||
// created.
|
||||
switch fam {
|
||||
case familyASR:
|
||||
err = n.loadASR(modelFile)
|
||||
case familyDiarization:
|
||||
err = n.loadDiarizer(modelFile)
|
||||
case familyTTS:
|
||||
if err = discoverTTSAssets(modelFile, &n.opts); err == nil {
|
||||
err = n.loadTTS(modelFile)
|
||||
}
|
||||
case familyNMT:
|
||||
err = n.loadNMT(modelFile)
|
||||
default:
|
||||
err = fmt.Errorf("nemo-speech-cpp: unhandled family for architecture %q", arch)
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
n.fam = fam
|
||||
return nil
|
||||
}
|
||||
|
||||
// Free destroys the runtime handle created at load time.
|
||||
//
|
||||
// base.SingleThread.Free is a no-op that derived backends are expected to
|
||||
// override, and every family here owns C memory that only its own destroy
|
||||
// entry point can release, so without this an unloaded model leaks a whole
|
||||
// acoustic model. Clearing fam as well means an RPC that races the unload is
|
||||
// refused by the gate rather than handed a dangling handle, but that only holds
|
||||
// for callers that took engineMu, which today means callers that went through
|
||||
// withEngine.
|
||||
func (n *NemoSpeech) Free() error {
|
||||
n.engineMu.Lock()
|
||||
defer n.engineMu.Unlock()
|
||||
|
||||
// Guarded on the handle, not on fam: a load that failed part way through
|
||||
// leaves fam unset, and the destroy functions are nil pointers until
|
||||
// openLibraries has bound them.
|
||||
// Each is tested independently rather than switched on: the one-handle
|
||||
// invariant is an invariant, and if it ever broke, a switch would silently
|
||||
// leak the others.
|
||||
if n.recognizer != 0 {
|
||||
ASRDestroy(n.recognizer)
|
||||
n.recognizer = 0
|
||||
}
|
||||
if n.diarizer != 0 {
|
||||
DiarDestroy(n.diarizer)
|
||||
n.diarizer = 0
|
||||
}
|
||||
if n.synth != 0 {
|
||||
TTSDestroy(n.synth)
|
||||
n.synth = 0
|
||||
}
|
||||
if n.nmt != 0 {
|
||||
NMTDestroy(n.nmt)
|
||||
n.nmt = 0
|
||||
}
|
||||
n.fam = familyUnknown
|
||||
return nil
|
||||
}
|
||||
|
||||
// The loaders are one per family: loadASR in asr.go, loadDiarizer in diar.go,
|
||||
// loadTTS in tts.go and loadNMT in nmt.go. Each populates its config structs
|
||||
// from n.opts and stores the handle in the matching field.
|
||||
//
|
||||
// Locking protocol, in both directions:
|
||||
//
|
||||
// - Every RPC must hold engineMu across its family check AND its C calls,
|
||||
// which means wrapping its body in withEngine. Free runs without the
|
||||
// backend lock (pkg/grpc/server.go:1019), so anything that checks the
|
||||
// family and then releases the lock before calling C can have the handle
|
||||
// destroyed underneath it. asr.go's AudioTranscription is the worked
|
||||
// example: even the audio decode sits inside the closure, because the
|
||||
// backend already serialises RPCs through base.SingleThread and so the
|
||||
// wider hold costs nothing.
|
||||
// - A loader must NOT take engineMu. Load holds it across the whole switch,
|
||||
// and sync.Mutex is not reentrant, so locking in a loader deadlocks.
|
||||
//
|
||||
// One consequence the streaming RPCs have to plan around: a stream whose body
|
||||
// is wrapped in withEngine holds engineMu for the WHOLE stream, so Free blocks
|
||||
// until the stream ends rather than tearing the handle out from under it. That
|
||||
// is the behaviour we want (a half-closed stream over a destroyed recognizer
|
||||
// has no good outcome), but it means an unload waits on a client that has
|
||||
// stopped sending, so a streaming loop must have its own way out: honour the
|
||||
// request context and stop on it, rather than blocking forever on the next
|
||||
// chunk.
|
||||
@@ -0,0 +1,13 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
func TestNemoSpeech(t *testing.T) {
|
||||
RegisterFailHandler(Fail)
|
||||
RunSpecs(t, "nemo-speech-cpp Backend Suite")
|
||||
}
|
||||
@@ -0,0 +1,260 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"sync"
|
||||
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
|
||||
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
|
||||
)
|
||||
|
||||
var _ = Describe("requireFamily", func() {
|
||||
It("accepts the loaded family", func() {
|
||||
n := &NemoSpeech{fam: familyASR}
|
||||
Expect(n.requireFamily(familyASR)).To(Succeed())
|
||||
})
|
||||
|
||||
It("rejects a mismatched family with Unimplemented and names both", func() {
|
||||
n := &NemoSpeech{fam: familyTTS}
|
||||
err := n.requireFamily(familyASR)
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(status.Code(err)).To(Equal(codes.Unimplemented))
|
||||
Expect(err.Error()).To(ContainSubstring("tts"))
|
||||
Expect(err.Error()).To(ContainSubstring("asr"))
|
||||
})
|
||||
|
||||
It("rejects an unloaded model", func() {
|
||||
n := &NemoSpeech{}
|
||||
err := n.requireFamily(familyASR)
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(status.Code(err)).To(Equal(codes.Unimplemented))
|
||||
})
|
||||
|
||||
It("rejects every family when the model is unloaded", func() {
|
||||
n := &NemoSpeech{}
|
||||
for _, f := range []family{familyASR, familyDiarization, familyTTS, familyNMT} {
|
||||
Expect(n.requireFamily(f)).To(HaveOccurred(), "family %s must be gated on an unloaded model", f)
|
||||
}
|
||||
})
|
||||
})
|
||||
|
||||
// The brief's round-trip spec (cstr then a reader) cannot exist: cstr pins a Go
|
||||
// allocation, and converting a uintptr back into a pointer to Go memory is a
|
||||
// checkptr violation that aborts the process under -race. So cstr is asserted
|
||||
// on what is observable without dereferencing its result.
|
||||
var _ = Describe("cstr", func() {
|
||||
It("returns a non-null pointer for a non-empty string", func() {
|
||||
p, free := cstr("hello")
|
||||
defer free()
|
||||
Expect(p).ToNot(BeZero())
|
||||
})
|
||||
|
||||
It("returns a null pointer for the empty string", func() {
|
||||
// The C API documents NULL and "" as equivalent for optional fields, and
|
||||
// passing NULL avoids allocating for every unset option.
|
||||
p, free := cstr("")
|
||||
defer free()
|
||||
Expect(p).To(BeZero())
|
||||
})
|
||||
|
||||
// The pin has to hold for the whole create call, which spans at least one
|
||||
// safepoint. A collection must therefore neither move nor invalidate the
|
||||
// address that C was handed.
|
||||
It("keeps the pointer stable across a garbage collection", func() {
|
||||
p, free := cstr("/models/nemo/parakeet.gguf")
|
||||
defer free()
|
||||
before := p
|
||||
runtime.GC()
|
||||
runtime.GC()
|
||||
Expect(p).To(Equal(before))
|
||||
})
|
||||
|
||||
It("survives releasing more than once", func() {
|
||||
_, free := cstr("twice")
|
||||
free()
|
||||
Expect(free).ToNot(Panic())
|
||||
})
|
||||
|
||||
// A dropped release leaks the pin, and the runtime turns that into a process
|
||||
// abort at some later collection. Nothing can catch it, so this only pins the
|
||||
// contract in prose: release in the same statement that takes the pointer.
|
||||
It("releases without panicking when used as documented", func() {
|
||||
Expect(func() {
|
||||
p, free := cstr("released")
|
||||
defer free()
|
||||
_ = p
|
||||
}).ToNot(Panic())
|
||||
})
|
||||
})
|
||||
|
||||
// pkg/grpc/server.go:1019 calls Free without taking the backend lock every
|
||||
// other RPC holds, so a teardown really can land while a request is in flight.
|
||||
// The family check and the C calls that trust it therefore have to happen under
|
||||
// engineMu together, or Free can destroy the handle in the gap between them.
|
||||
var _ = Describe("engine locking", func() {
|
||||
It("serialises a teardown against an in-flight request", func() {
|
||||
n := &NemoSpeech{fam: familyASR}
|
||||
|
||||
var wg sync.WaitGroup
|
||||
wg.Add(2)
|
||||
go func() {
|
||||
defer GinkgoRecover()
|
||||
defer wg.Done()
|
||||
for i := 0; i < 2000; i++ {
|
||||
// Errors are expected once the teardown wins the race; what must
|
||||
// not happen is an unsynchronised read of the family.
|
||||
_ = n.withEngine(familyASR, func() error { return nil })
|
||||
}
|
||||
}()
|
||||
go func() {
|
||||
defer GinkgoRecover()
|
||||
defer wg.Done()
|
||||
for i := 0; i < 2000; i++ {
|
||||
Expect(n.Free()).To(Succeed())
|
||||
}
|
||||
}()
|
||||
wg.Wait()
|
||||
})
|
||||
|
||||
It("refuses the body when the family does not match, and still unlocks", func() {
|
||||
n := &NemoSpeech{fam: familyTTS}
|
||||
called := false
|
||||
err := n.withEngine(familyASR, func() error {
|
||||
called = true
|
||||
return nil
|
||||
})
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(status.Code(err)).To(Equal(codes.Unimplemented))
|
||||
Expect(called).To(BeFalse())
|
||||
|
||||
// A lock leaked on the rejection path would deadlock the next request
|
||||
// rather than fail it, so prove the mutex is free afterwards.
|
||||
Expect(n.engineMu.TryLock()).To(BeTrue())
|
||||
n.engineMu.Unlock()
|
||||
})
|
||||
|
||||
It("propagates the body's error and still unlocks", func() {
|
||||
n := &NemoSpeech{fam: familyASR}
|
||||
boom := errors.New("boom")
|
||||
Expect(n.withEngine(familyASR, func() error { return boom })).To(MatchError(boom))
|
||||
Expect(n.engineMu.TryLock()).To(BeTrue())
|
||||
n.engineMu.Unlock()
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("Free", func() {
|
||||
// The destroy entry points are nil function values until openLibraries has
|
||||
// bound them, so an unloaded model must not reach them. LocalAI frees every
|
||||
// backend it shuts down, including one whose Load failed.
|
||||
It("is a no-op on a model that was never loaded", func() {
|
||||
n := &NemoSpeech{}
|
||||
Expect(n.Free()).To(Succeed())
|
||||
})
|
||||
|
||||
It("is idempotent", func() {
|
||||
n := &NemoSpeech{}
|
||||
Expect(n.Free()).To(Succeed())
|
||||
Expect(n.Free()).To(Succeed())
|
||||
})
|
||||
|
||||
It("does not reach the runtime for a load that failed part way through", func() {
|
||||
n := &NemoSpeech{}
|
||||
path := filepath.Join(GinkgoT().TempDir(), "broken.gguf")
|
||||
Expect(os.WriteFile(path, []byte("broken"), 0o600)).To(Succeed())
|
||||
Expect(n.Load(&pb.ModelOptions{ModelFile: path})).ToNot(Succeed())
|
||||
Expect(n.Free()).To(Succeed())
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("Load", func() {
|
||||
It("rejects an empty model file", func() {
|
||||
n := &NemoSpeech{}
|
||||
err := n.Load(&pb.ModelOptions{})
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(err.Error()).To(ContainSubstring("ModelFile"))
|
||||
})
|
||||
|
||||
It("reports a model file that does not exist", func() {
|
||||
n := &NemoSpeech{}
|
||||
missing := filepath.Join(GinkgoT().TempDir(), "absent.gguf")
|
||||
err := n.Load(&pb.ModelOptions{ModelFile: missing})
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(err.Error()).To(ContainSubstring("absent.gguf"))
|
||||
})
|
||||
|
||||
It("reports a file that is not a GGUF", func() {
|
||||
n := &NemoSpeech{}
|
||||
path := filepath.Join(GinkgoT().TempDir(), "notagguf.gguf")
|
||||
Expect(os.WriteFile(path, []byte("this is not a gguf file at all"), 0o600)).To(Succeed())
|
||||
err := n.Load(&pb.ModelOptions{ModelFile: path})
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(err.Error()).To(ContainSubstring("nemo-speech-cpp"))
|
||||
})
|
||||
|
||||
// A failed load must not leave a family selected, or the RPC gate would wave
|
||||
// requests through to a nil handle.
|
||||
It("leaves no family selected when the load fails", func() {
|
||||
n := &NemoSpeech{}
|
||||
path := filepath.Join(GinkgoT().TempDir(), "broken.gguf")
|
||||
Expect(os.WriteFile(path, []byte("broken"), 0o600)).To(Succeed())
|
||||
Expect(n.Load(&pb.ModelOptions{ModelFile: path})).ToNot(Succeed())
|
||||
Expect(n.fam).To(Equal(familyUnknown))
|
||||
Expect(n.requireFamily(familyASR)).To(HaveOccurred())
|
||||
})
|
||||
|
||||
// The load path picks a family and only then runs that family's loader, so
|
||||
// there is a window where the family is known and the load still fails.
|
||||
// Committing n.fam before the loader runs would leave the RPC gate open on a
|
||||
// handle that was never created, and pkg/grpc/server.go keeps serving the
|
||||
// instance after a failed LoadModel, so the next request really would reach
|
||||
// it. TTS is the only family whose loader can fail before touching C.
|
||||
It("does not select the family until that family's loader has succeeded", func() {
|
||||
dir := GinkgoT().TempDir()
|
||||
path := filepath.Join(dir, "magpie.f16.gguf")
|
||||
writeGGUFWithArch(path, "magpietts")
|
||||
|
||||
// Self-guard: if the handwritten GGUF ever stops parsing, Load would fail
|
||||
// at ggufArchitecture instead, before a family is ever chosen, and the
|
||||
// assertions below would pass without exercising the ordering at all.
|
||||
Expect(ggufArchitecture(path)).To(Equal("magpietts"))
|
||||
|
||||
// No sibling codec in the directory, so discoverTTSAssets fails after
|
||||
// familyFor has already resolved familyTTS.
|
||||
n := &NemoSpeech{}
|
||||
err := n.Load(&pb.ModelOptions{ModelFile: path})
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(err.Error()).To(ContainSubstring("codec_model"))
|
||||
|
||||
Expect(n.fam).To(Equal(familyUnknown))
|
||||
Expect(n.requireFamily(familyTTS)).To(HaveOccurred())
|
||||
})
|
||||
|
||||
It("closes the family gate again after a free", func() {
|
||||
n := &NemoSpeech{fam: familyASR}
|
||||
Expect(n.Free()).To(Succeed())
|
||||
Expect(n.fam).To(Equal(familyUnknown))
|
||||
Expect(n.requireFamily(familyASR)).To(HaveOccurred())
|
||||
})
|
||||
|
||||
It("parses the model options before it touches the model file", func() {
|
||||
// The options are what tell a TTS load where its codec lives, so they have
|
||||
// to be in place before any family-specific loader runs.
|
||||
n := &NemoSpeech{}
|
||||
path := filepath.Join(GinkgoT().TempDir(), "broken.gguf")
|
||||
Expect(os.WriteFile(path, []byte("broken"), 0o600)).To(Succeed())
|
||||
Expect(n.Load(&pb.ModelOptions{
|
||||
ModelFile: path,
|
||||
ModelPath: "/models",
|
||||
Options: []string{"gpu:2", "codec_model:codec.gguf"},
|
||||
})).ToNot(Succeed())
|
||||
Expect(n.opts.gpu).To(Equal(int32(2)))
|
||||
Expect(n.opts.codecModel).To(Equal("/models/codec.gguf"))
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,368 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"regexp"
|
||||
"runtime"
|
||||
"strings"
|
||||
"unsafe"
|
||||
|
||||
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
|
||||
"github.com/mudler/xlog"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
)
|
||||
|
||||
// pairDirective matches a leading "[src->tgt] " override.
|
||||
//
|
||||
// Each side is an unbounded run of two-letter segments, not one or two of them.
|
||||
// Either side may also be omitted, which keeps the model-level default for it:
|
||||
// resolve_tag accepts a READY pair tag in one field with the other empty
|
||||
// (src/nmt/langpairs.cc:167-172), so "[->en-de]" names a pair for one request.
|
||||
//
|
||||
// The two rules together are what force the unbounded run. A regional code on
|
||||
// its own is only two segments (pt-br, zh-cn, es-us) and would parse under a
|
||||
// stricter pattern; it is the SINGLE-FIELD form of a regional pair that runs to
|
||||
// three (en-zh-cn, en-zh-tw, en-es-us, en-pt-br, pt-br-en, zh-tw-en). And the
|
||||
// failure is not a mis-split: a pattern too short to cover the tag does not
|
||||
// match the directive at all, so the whole bracket survives into the text and
|
||||
// is handed to the model as something to translate.
|
||||
//
|
||||
// The codes are not normalised or validated here. normalize_language_code
|
||||
// lowercases and folds BCP-47 down to a supported base, and is_supported has the
|
||||
// authoritative table; duplicating either would be a second source of truth that
|
||||
// drifts on the next pin bump.
|
||||
var pairDirective = regexp.MustCompile(`^\[\s*([a-zA-Z]{2}(?:-[a-zA-Z]{2})*)?\s*->\s*([a-zA-Z]{2}(?:-[a-zA-Z]{2})*)?\s*\]\s*`)
|
||||
|
||||
// translator is the NMT half of the C API, narrowed to what Predict uses.
|
||||
//
|
||||
// It is an interface for the same reason synthesizer and diarStream are: no
|
||||
// Riva-Translate GGUF is small enough to keep in the tree, so the layer above
|
||||
// the ABI (pair resolution, validation, the text array, the single-chunk stream)
|
||||
// would otherwise have no test at all. A fake here scripts what the C API
|
||||
// returns; it does not pretend to translate anything.
|
||||
type translator interface {
|
||||
// translate returns one translation per input text, in order.
|
||||
translate(texts []string, source, target string) ([]string, error)
|
||||
}
|
||||
|
||||
// cTranslator is the real translator, over one nemo_speech_nmt_translator.
|
||||
type cTranslator struct {
|
||||
handle uintptr
|
||||
}
|
||||
|
||||
// nmtTexts builds the `const char* const* texts` argument and returns it with
|
||||
// the release the caller MUST defer.
|
||||
//
|
||||
// Two levels need pinning, not one. cstr pins each string's bytes, but the array
|
||||
// carrying their addresses is a separate Go allocation holding uintptrs: the
|
||||
// collector neither traces through it nor is obliged to leave it where it is,
|
||||
// and C dereferences it for the whole call. Pinning only the strings would leave
|
||||
// the array itself free to move out from under the runtime.
|
||||
//
|
||||
// An empty element is refused rather than passed on. cstr maps "" to NULL and
|
||||
// src/nmt/c_api.cpp maps a NULL element back to "" (str_or_empty), so a blank
|
||||
// text would come back as a confident translation of nothing rather than an
|
||||
// error.
|
||||
func nmtTexts(texts []string) ([]uintptr, func(), error) {
|
||||
pin := new(runtime.Pinner)
|
||||
// The pin is released first so that the array stops being pinned before the
|
||||
// strings it points at do.
|
||||
frees := []func(){pin.Unpin}
|
||||
release := func() {
|
||||
for _, f := range frees {
|
||||
f()
|
||||
}
|
||||
}
|
||||
|
||||
if len(texts) == 0 {
|
||||
return nil, release, status.Error(codes.InvalidArgument,
|
||||
"nemo-speech-cpp: nothing to translate")
|
||||
}
|
||||
|
||||
ptrs := make([]uintptr, len(texts))
|
||||
for i, t := range texts {
|
||||
if t == "" {
|
||||
return nil, release, status.Error(codes.InvalidArgument,
|
||||
"nemo-speech-cpp: nothing to translate")
|
||||
}
|
||||
p, free := cstr(t)
|
||||
frees = append(frees, free)
|
||||
ptrs[i] = p
|
||||
}
|
||||
pin.Pin(&ptrs[0])
|
||||
return ptrs, release, nil
|
||||
}
|
||||
|
||||
func (t *cTranslator) translate(texts []string, source, target string) ([]string, error) {
|
||||
ptrs, release, err := nmtTexts(texts)
|
||||
if err != nil {
|
||||
release()
|
||||
return nil, err
|
||||
}
|
||||
defer release()
|
||||
|
||||
// source and target cross as Go strings: purego NUL-terminates and copies
|
||||
// them itself for the duration of the call, and c_api.cpp deep-copies both
|
||||
// into std::string before doing anything with them.
|
||||
var result uintptr
|
||||
if st := NMTTranslate(t.handle, &ptrs[0], uint64(len(ptrs)), source, target, &result); st != 0 {
|
||||
// An unsupported language pair arrives here as INVALID_ARGUMENT
|
||||
// (src/nmt/translator.cpp throws std::invalid_argument, which
|
||||
// src/nmt/c_api.cpp's guard maps to it), which statusErrorf turns into
|
||||
// the caller-facing code rather than Internal.
|
||||
return nil, statusErrorf(st, "nemo-speech-cpp: translate: %s", NMTLastError())
|
||||
}
|
||||
defer NMTResultDestroy(result)
|
||||
|
||||
count := NMTResultCount(result)
|
||||
out := make([]string, 0, count)
|
||||
for i := uint64(0); i < count; i++ {
|
||||
out = append(out, NMTResultText(result, i))
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// nmtTranslatorConfig builds the create-time config.
|
||||
//
|
||||
// Extracted from loadNMT so its four adjacent pointer fields can be asserted
|
||||
// against distinct sentinels. Backend, Model, Generation and Pool are all
|
||||
// uintptr and all sit next to each other, so transposing two of them changes
|
||||
// neither the struct's size nor any field's offset: the layout assertions in
|
||||
// abi_test.go are blind to it, and what it produces at runtime is the backend
|
||||
// config being read as the model config.
|
||||
//
|
||||
// Generation and Pool stay NULL, which nmt.h documents as "library defaults":
|
||||
// max_new_tokens (256) and contexts (1) are create-time settings this backend
|
||||
// has no option to fill them from, and PredictOptions carries no per-request
|
||||
// equivalent that a create-time config could honour anyway.
|
||||
//
|
||||
// backend and model are pinned addresses, not Go pointers, and the caller owns
|
||||
// the pins.
|
||||
func nmtTranslatorConfig(backend, model uintptr) cNMTTranslatorConfig {
|
||||
return cNMTTranslatorConfig{
|
||||
Size: unsafe.Sizeof(cNMTTranslatorConfig{}),
|
||||
Backend: backend,
|
||||
Model: model,
|
||||
}
|
||||
}
|
||||
|
||||
// loadNMT creates the Riva-Translate translator.
|
||||
//
|
||||
// This must not take engineMu: Load is its only caller and already holds it.
|
||||
func (n *NemoSpeech) loadNMT(modelFile string) error {
|
||||
// nemo_speech_nmt_create deep-copies the path into a std::string
|
||||
// (src/nmt/c_api.cpp to_config, via str_or_empty) and retains no pointer
|
||||
// afterwards, so pinning for the duration of the create call is both
|
||||
// necessary and sufficient.
|
||||
var pinner runtime.Pinner
|
||||
defer pinner.Unpin()
|
||||
|
||||
pathP, freePath := cstr(modelFile)
|
||||
defer freePath()
|
||||
|
||||
// NCtx is left at 0, which to_config reads as "keep the default" (it applies
|
||||
// the field only when > 0) and which the runtime resolves to 1024 tokens.
|
||||
// That is sized for the sentence-length input Riva-Translate is built for,
|
||||
// and raising it costs one n_ctx-sized KV cache per pooled context, so it
|
||||
// wants a deliberate option rather than a guess made here.
|
||||
model := cNMTModelConfig{Size: unsafe.Sizeof(cNMTModelConfig{}), Path: pathP}
|
||||
// BackendConfig.gpu defaults to 0 in C++ (device 0), not to CPU, and
|
||||
// to_config assigns it unconditionally, so the option's own -1 default is
|
||||
// what keeps an unconfigured model on the CPU.
|
||||
backend := cNMTBackendConfig{Size: unsafe.Sizeof(cNMTBackendConfig{}), GPU: n.opts.gpu}
|
||||
|
||||
cfg := nmtTranslatorConfig(pinPtr(&pinner, &backend), pinPtr(&pinner, &model))
|
||||
|
||||
xlog.Info("nemo-speech-cpp: creating translator",
|
||||
"gpu", n.opts.gpu,
|
||||
"source_language", n.opts.sourceLanguage,
|
||||
"target_language", n.opts.targetLanguage)
|
||||
|
||||
// #nosec G103 -- cfg is a local POD struct borrowed for this call only. Its
|
||||
// Backend and Model members are pinPtr addresses held by the pinner unpinned
|
||||
// on return, Model.Path is the cstr allocation freed by the defer above, and
|
||||
// nemo_speech_nmt_create deep-copies everything it reads.
|
||||
if st := NMTCreate(unsafe.Pointer(&cfg), &n.nmt); st != 0 {
|
||||
return statusErrorf(st, "nemo-speech-cpp: nmt create: %s", NMTLastError())
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// languagePair resolves the languages for one request and returns the text to
|
||||
// translate.
|
||||
//
|
||||
// nemo_speech_nmt_translate takes explicit source and target languages and has
|
||||
// no free-form generation entry point at all, so there is no prompt in the LLM
|
||||
// sense to carry an instruction. The pair therefore comes from the model
|
||||
// options, and a leading "[src->tgt]" directive is the only per-request control
|
||||
// Predict can offer.
|
||||
func (n *NemoSpeech) languagePair(prompt string) (source, target, text string) {
|
||||
source, target = n.opts.sourceLanguage, n.opts.targetLanguage
|
||||
|
||||
m := pairDirective.FindStringSubmatch(prompt)
|
||||
if m == nil {
|
||||
return source, target, strings.TrimSpace(prompt)
|
||||
}
|
||||
// An omitted side keeps the model-level default rather than blanking it.
|
||||
if m[1] != "" {
|
||||
source = m[1]
|
||||
}
|
||||
if m[2] != "" {
|
||||
target = m[2]
|
||||
}
|
||||
// The directive must not survive into the text: the runtime wraps it in a
|
||||
// chat template (src/nmt/langpairs.cc build_prompt), so anything left here is
|
||||
// translated along with the sentence.
|
||||
return source, target, strings.TrimSpace(prompt[len(m[0]):])
|
||||
}
|
||||
|
||||
// unsupportedPredictFields names the PredictOptions fields a caller may have set
|
||||
// that this C API has no way to honour, so they are logged rather than silently
|
||||
// dropped.
|
||||
//
|
||||
// The list is deliberately narrow. Everything nemo_speech_nmt_translate accepts
|
||||
// is in its five arguments: a translator, the texts, and two language codes.
|
||||
// Everything else in PredictOptions is therefore unsupported, and naming all of
|
||||
// it would log on every single request, because LocalAI fills the sampling
|
||||
// defaults in from the model config whether or not the user asked for them.
|
||||
//
|
||||
// So the sampling and decoding knobs (temperature, top_p, top_k, min_p, seed,
|
||||
// tokens, repeat/frequency/presence penalties, mirostat, tfz, typical_p,
|
||||
// stop_prompts, prompt caching, rope scaling, n_draft, logit_bias) are ignored
|
||||
// silently: there is no field for any of them on either side of the ABI.
|
||||
// max_new_tokens and n_ctx exist but are CREATE-time settings on the translator,
|
||||
// not per-request ones, so PredictOptions.Tokens has nowhere to go either.
|
||||
//
|
||||
// What is named here is the structural asks: requests that only make sense
|
||||
// against a general language model, where honouring them partially would be
|
||||
// worse than saying nothing at all.
|
||||
func unsupportedPredictFields(opts *pb.PredictOptions) []string {
|
||||
var out []string
|
||||
if opts.GetGrammar() != "" {
|
||||
out = append(out, "grammar")
|
||||
}
|
||||
if opts.GetTools() != "" {
|
||||
out = append(out, "tools")
|
||||
}
|
||||
if len(opts.GetImages()) > 0 {
|
||||
out = append(out, "images")
|
||||
}
|
||||
if len(opts.GetVideos()) > 0 {
|
||||
out = append(out, "videos")
|
||||
}
|
||||
if len(opts.GetAudios()) > 0 {
|
||||
out = append(out, "audios")
|
||||
}
|
||||
if opts.GetNegativePrompt() != "" {
|
||||
out = append(out, "negative_prompt")
|
||||
}
|
||||
if opts.GetLogprobs() > 0 {
|
||||
out = append(out, "logprobs")
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// translateText runs one translation and returns it.
|
||||
//
|
||||
// The two rejections happen before anything crosses the ABI. An empty text would
|
||||
// otherwise reach the runtime as a NULL element (see nmtTexts), and a missing
|
||||
// target would come back as "unsupported language pair: -> ", which names
|
||||
// neither the option the operator has to set nor the request that failed.
|
||||
func translateText(t translator, source, target, text string) (string, error) {
|
||||
if text == "" {
|
||||
return "", status.Error(codes.InvalidArgument,
|
||||
"nemo-speech-cpp: PredictOptions.prompt is required, it is the text to translate")
|
||||
}
|
||||
if target == "" {
|
||||
return "", status.Error(codes.InvalidArgument,
|
||||
"nemo-speech-cpp: no target language: set the target_language model option, "+
|
||||
"or prefix the prompt with a [src->tgt] directive")
|
||||
}
|
||||
|
||||
out, err := t.translate([]string{text}, source, target)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
// One text in, one translation out. A call that returned OK with none is a
|
||||
// runtime bug, and the empty string it would hand back reaches the user as a
|
||||
// successful but blank completion with nothing anywhere to say why.
|
||||
if len(out) == 0 {
|
||||
return "", status.Error(codes.Internal, "nemo-speech-cpp: translation produced no result")
|
||||
}
|
||||
return out[0], nil
|
||||
}
|
||||
|
||||
// streamTranslation runs one translation and puts the whole of it on out as a
|
||||
// single chunk.
|
||||
//
|
||||
// That is a limit of the C API and not a shortcut taken here.
|
||||
// nemo_speech_nmt_translate has no token callback and no incremental result: it
|
||||
// returns once the decode has finished, with the completed text. There is
|
||||
// nothing finer to stream, and splitting the finished string into fake chunks
|
||||
// would imitate progress that never happened.
|
||||
//
|
||||
// out is not closed here. PredictStream owns it, and closing it in one of two
|
||||
// places depending on how far the request got is how a stream ends up
|
||||
// half-closed.
|
||||
func streamTranslation(t translator, source, target, text string, out chan<- string) error {
|
||||
translated, err := translateText(t, source, target, text)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
out <- translated
|
||||
return nil
|
||||
}
|
||||
|
||||
// resolveRequest is the shared front half of both RPCs: it names what it is
|
||||
// dropping and works out the pair and the text.
|
||||
func (n *NemoSpeech) resolveRequest(opts *pb.PredictOptions) (source, target, text string) {
|
||||
// Logged rather than rejected, for the reason the diarization path logs its
|
||||
// own dropped fields: a caller that asked for something extra still wants the
|
||||
// translation it can have, and a request naming a field this backend drops
|
||||
// should say so where an operator can find it.
|
||||
if dropped := unsupportedPredictFields(opts); len(dropped) > 0 {
|
||||
xlog.Warn("nemo-speech-cpp: ignoring request fields this model has no equivalent for",
|
||||
"fields", dropped)
|
||||
}
|
||||
return n.languagePair(opts.GetPrompt())
|
||||
}
|
||||
|
||||
// Predict translates PredictOptions.Prompt.
|
||||
//
|
||||
// The whole body runs inside withEngine, so the family check and the C calls
|
||||
// that trust the handle happen under a single acquisition of engineMu. See the
|
||||
// handoff notes at the bottom of nemospeech.go: Free runs without the backend
|
||||
// lock, so anything that checks the family and then releases the lock before
|
||||
// calling C can have the handle destroyed underneath it.
|
||||
func (n *NemoSpeech) Predict(opts *pb.PredictOptions) (string, error) {
|
||||
var out string
|
||||
if err := n.withEngine(familyNMT, func() error {
|
||||
source, target, text := n.resolveRequest(opts)
|
||||
s, err := translateText(&cTranslator{handle: n.nmt}, source, target, text)
|
||||
out = s
|
||||
return err
|
||||
}); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// PredictStream translates PredictOptions.Prompt and emits the result on
|
||||
// results.
|
||||
//
|
||||
// results is closed on EVERY path, including the family rejection and a
|
||||
// validation failure, and the close is deferred outside withEngine so that a
|
||||
// rejected family still closes it. This is the LEGACY streaming contract, which
|
||||
// is the opposite of PredictStreamRich's: pkg/grpc/server.go:529 calls this and
|
||||
// then blocks on a drain goroutine that only finishes when the channel closes,
|
||||
// so a channel left open does not fail the request, it hangs the RPC and, with
|
||||
// the backend lock still held, every request queued behind it. The rich variant
|
||||
// is the one whose channel the host closes; this one is not.
|
||||
func (n *NemoSpeech) PredictStream(opts *pb.PredictOptions, results chan string) error {
|
||||
defer close(results)
|
||||
|
||||
return n.withEngine(familyNMT, func() error {
|
||||
source, target, text := n.resolveRequest(opts)
|
||||
return streamTranslation(&cTranslator{handle: n.nmt}, source, target, text, results)
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,416 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"unsafe"
|
||||
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
|
||||
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
|
||||
)
|
||||
|
||||
// fakeTranslator scripts the C API's answer and records what it was asked, so
|
||||
// the layer above the ABI (pair resolution, validation, the single-element text
|
||||
// array) has a test at all. No Riva-Translate GGUF is small enough to keep in
|
||||
// the tree, and this pretends to translate nothing.
|
||||
type fakeTranslator struct {
|
||||
texts []string
|
||||
source, target string
|
||||
calls int
|
||||
|
||||
out []string
|
||||
err error
|
||||
}
|
||||
|
||||
func (f *fakeTranslator) translate(texts []string, source, target string) ([]string, error) {
|
||||
f.calls++
|
||||
f.texts = texts
|
||||
f.source = source
|
||||
f.target = target
|
||||
return f.out, f.err
|
||||
}
|
||||
|
||||
// collectStrings drains ch until it closes and hands back everything it saw.
|
||||
// The host does the same, so a channel this backend forgets to close hangs the
|
||||
// RPC rather than failing it.
|
||||
func collectStrings(ch chan string) chan []string {
|
||||
done := make(chan []string, 1)
|
||||
go func() {
|
||||
var got []string
|
||||
for s := range ch {
|
||||
got = append(got, s)
|
||||
}
|
||||
done <- got
|
||||
}()
|
||||
return done
|
||||
}
|
||||
|
||||
var _ = Describe("languagePair", func() {
|
||||
It("uses the configured pair and returns the prompt unchanged", func() {
|
||||
n := &NemoSpeech{opts: loadOptions{sourceLanguage: "en", targetLanguage: "de"}}
|
||||
src, tgt, text := n.languagePair("hello world")
|
||||
Expect(src).To(Equal("en"))
|
||||
Expect(tgt).To(Equal("de"))
|
||||
Expect(text).To(Equal("hello world"))
|
||||
})
|
||||
|
||||
// nemo_speech_nmt_translate takes explicit languages and has no prompt path,
|
||||
// so an inline directive is the only way a caller can pick a pair per request.
|
||||
It("honours an inline pair directive and strips it from the text", func() {
|
||||
n := &NemoSpeech{opts: loadOptions{sourceLanguage: "en", targetLanguage: "de"}}
|
||||
src, tgt, text := n.languagePair("[en->fr] hello world")
|
||||
Expect(src).To(Equal("en"))
|
||||
Expect(tgt).To(Equal("fr"))
|
||||
Expect(text).To(Equal("hello world"))
|
||||
})
|
||||
|
||||
// The directive has to be gone from what reaches the model: the runtime
|
||||
// wraps the text in a chat template (src/nmt/langpairs.cc build_prompt), so
|
||||
// a leftover "[en->fr]" would be translated along with the sentence.
|
||||
It("leaves no trace of the directive in the translated text", func() {
|
||||
n := &NemoSpeech{opts: loadOptions{targetLanguage: "de"}}
|
||||
_, _, text := n.languagePair("[en->fr] hello world")
|
||||
Expect(text).ToNot(ContainSubstring("["))
|
||||
Expect(text).ToNot(ContainSubstring("->"))
|
||||
Expect(text).ToNot(ContainSubstring("fr"))
|
||||
Expect(text).To(Equal("hello world"))
|
||||
})
|
||||
|
||||
It("leaves an unparseable directive in the text", func() {
|
||||
n := &NemoSpeech{opts: loadOptions{sourceLanguage: "en", targetLanguage: "de"}}
|
||||
src, tgt, text := n.languagePair("[not a directive] hi")
|
||||
Expect(src).To(Equal("en"))
|
||||
Expect(tgt).To(Equal("de"))
|
||||
Expect(text).To(Equal("[not a directive] hi"))
|
||||
})
|
||||
|
||||
It("trims surrounding whitespace from the text", func() {
|
||||
n := &NemoSpeech{opts: loadOptions{sourceLanguage: "en", targetLanguage: "de"}}
|
||||
_, _, text := n.languagePair(" hello ")
|
||||
Expect(text).To(Equal("hello"))
|
||||
})
|
||||
|
||||
// The model's own tags carry region subtags (src/nmt/langpairs.cc: en-zh-cn,
|
||||
// pt-br, es-us), so a directive that only accepted bare two-letter codes
|
||||
// could not name half the pairs the runtime supports.
|
||||
It("accepts a regional code on either side", func() {
|
||||
n := &NemoSpeech{}
|
||||
src, tgt, text := n.languagePair("[pt-br->en] ola")
|
||||
Expect(src).To(Equal("pt-br"))
|
||||
Expect(tgt).To(Equal("en"))
|
||||
Expect(text).To(Equal("ola"))
|
||||
|
||||
src, tgt, _ = n.languagePair("[en->zh-cn] hi")
|
||||
Expect(src).To(Equal("en"))
|
||||
Expect(tgt).To(Equal("zh-cn"))
|
||||
})
|
||||
|
||||
// resolve_tag accepts a ready pair tag in one field with the other empty, so
|
||||
// a directive that names only one side must keep the configured value for the
|
||||
// other rather than blanking it.
|
||||
It("keeps the configured code for a side the directive omits", func() {
|
||||
n := &NemoSpeech{opts: loadOptions{sourceLanguage: "en", targetLanguage: "de"}}
|
||||
src, tgt, text := n.languagePair("[->fr] hello")
|
||||
Expect(src).To(Equal("en"))
|
||||
Expect(tgt).To(Equal("fr"))
|
||||
Expect(text).To(Equal("hello"))
|
||||
|
||||
src, tgt, _ = n.languagePair("[fr->] hello")
|
||||
Expect(src).To(Equal("fr"))
|
||||
Expect(tgt).To(Equal("de"))
|
||||
})
|
||||
|
||||
// resolve_tag (src/nmt/langpairs.cc:167-172) accepts a READY pair tag in one
|
||||
// field with the other empty, and the model's own tags run to three segments
|
||||
// (en-zh-cn, en-zh-tw, en-es-us, en-pt-br). That single-field three-segment
|
||||
// form is the case a two-segment pattern cannot express: it does not merely
|
||||
// mis-split the tag, it fails to match the directive at all, so the whole
|
||||
// bracket survives into the text and is handed to the model as something to
|
||||
// translate.
|
||||
//
|
||||
// Two-segment codes like pt-br and zh-cn are NOT this case; they parse either
|
||||
// way.
|
||||
It("accepts a three-segment pair tag given in one side of the directive", func() {
|
||||
n := &NemoSpeech{opts: loadOptions{targetLanguage: "de"}}
|
||||
src, tgt, text := n.languagePair("[->en-zh-cn] hi")
|
||||
Expect(src).To(BeEmpty())
|
||||
Expect(tgt).To(Equal("en-zh-cn"))
|
||||
Expect(text).To(Equal("hi"))
|
||||
|
||||
src, tgt, text = n.languagePair("[pt-br-en->] hola")
|
||||
Expect(src).To(Equal("pt-br-en"))
|
||||
Expect(tgt).To(Equal("de"))
|
||||
Expect(text).To(Equal("hola"))
|
||||
})
|
||||
|
||||
It("does not treat a bracketed sentence as a directive", func() {
|
||||
n := &NemoSpeech{opts: loadOptions{targetLanguage: "de"}}
|
||||
_, _, text := n.languagePair("[see figure 1] the cat sat")
|
||||
Expect(text).To(Equal("[see figure 1] the cat sat"))
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("nmtTranslatorConfig", func() {
|
||||
// Backend, Model, Generation and Pool are four adjacent same-typed pointers.
|
||||
// Transposing two of them changes neither the struct's size nor any field's
|
||||
// offset, so the layout assertions in abi_test.go cannot see it, and the
|
||||
// failure it produces is the runtime reading the backend config as the model
|
||||
// config. Distinct sentinels are the only thing that catches it.
|
||||
It("wires each pointer into its own field", func() {
|
||||
cfg := nmtTranslatorConfig(0xB, 0xD)
|
||||
Expect(cfg.Backend).To(Equal(uintptr(0xB)))
|
||||
Expect(cfg.Model).To(Equal(uintptr(0xD)))
|
||||
})
|
||||
|
||||
// NULL is what nmt.h documents as "library defaults" for a subsystem config,
|
||||
// and this backend has no option to fill either of them from.
|
||||
It("leaves the generation and pool configs null", func() {
|
||||
cfg := nmtTranslatorConfig(0xB, 0xD)
|
||||
Expect(cfg.Generation).To(BeZero())
|
||||
Expect(cfg.Pool).To(BeZero())
|
||||
})
|
||||
|
||||
// The runtime decides a field is present with HAS_FIELD, which tests the
|
||||
// caller's size against offsetof + sizeof (src/nmt/c_api.cpp), so a config
|
||||
// sent with Size 0 has every field ignored and the model loads from a path
|
||||
// it was never given.
|
||||
It("declares its own size", func() {
|
||||
Expect(nmtTranslatorConfig(0xB, 0xD).Size).To(Equal(unsafe.Sizeof(cNMTTranslatorConfig{})))
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("nmtTexts", func() {
|
||||
It("produces one non-null pointer per text", func() {
|
||||
ptrs, release, err := nmtTexts([]string{"one", "two", "three"})
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
defer release()
|
||||
|
||||
Expect(ptrs).To(HaveLen(3))
|
||||
for i, p := range ptrs {
|
||||
Expect(p).ToNot(BeZero(), "texts[%d] must not be NULL", i)
|
||||
}
|
||||
// Distinct addresses: one buffer reused for every element would make the
|
||||
// runtime translate the last text three times.
|
||||
Expect(ptrs[0]).ToNot(Equal(ptrs[1]))
|
||||
Expect(ptrs[1]).ToNot(Equal(ptrs[2]))
|
||||
})
|
||||
|
||||
// cstr maps "" to NULL and src/nmt/c_api.cpp maps a NULL element back to "",
|
||||
// so a blank text would be answered with a translation of nothing instead of
|
||||
// an error.
|
||||
It("refuses an empty element", func() {
|
||||
_, release, err := nmtTexts([]string{"one", ""})
|
||||
Expect(release).ToNot(BeNil())
|
||||
release()
|
||||
Expect(status.Code(err)).To(Equal(codes.InvalidArgument))
|
||||
})
|
||||
|
||||
It("refuses an empty batch", func() {
|
||||
_, release, err := nmtTexts(nil)
|
||||
Expect(release).ToNot(BeNil())
|
||||
release()
|
||||
Expect(status.Code(err)).To(Equal(codes.InvalidArgument))
|
||||
})
|
||||
|
||||
It("survives releasing more than once", func() {
|
||||
_, release, err := nmtTexts([]string{"once"})
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
release()
|
||||
Expect(release).ToNot(Panic())
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("translateText", func() {
|
||||
It("passes the resolved pair and the text through to the runtime", func() {
|
||||
f := &fakeTranslator{out: []string{"hallo welt"}}
|
||||
got, err := translateText(f, "en", "de", "hello world")
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(got).To(Equal("hallo welt"))
|
||||
Expect(f.texts).To(Equal([]string{"hello world"}))
|
||||
Expect(f.source).To(Equal("en"))
|
||||
Expect(f.target).To(Equal("de"))
|
||||
})
|
||||
|
||||
// A single-pair model is configured with target_language alone, and
|
||||
// resolve_tag accepts a ready tag in one field with the other empty.
|
||||
It("allows an empty source language", func() {
|
||||
f := &fakeTranslator{out: []string{"ciao"}}
|
||||
_, err := translateText(f, "", "en-it", "hi")
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(f.source).To(BeEmpty())
|
||||
})
|
||||
|
||||
It("rejects a missing target language and names the option to set", func() {
|
||||
f := &fakeTranslator{}
|
||||
_, err := translateText(f, "en", "", "hello")
|
||||
Expect(status.Code(err)).To(Equal(codes.InvalidArgument))
|
||||
Expect(err.Error()).To(ContainSubstring("target_language"))
|
||||
Expect(f.calls).To(BeZero())
|
||||
})
|
||||
|
||||
It("rejects an empty text without calling the runtime", func() {
|
||||
f := &fakeTranslator{}
|
||||
_, err := translateText(f, "en", "de", "")
|
||||
Expect(status.Code(err)).To(Equal(codes.InvalidArgument))
|
||||
Expect(f.calls).To(BeZero())
|
||||
})
|
||||
|
||||
It("propagates a runtime failure", func() {
|
||||
boom := errors.New("boom")
|
||||
_, err := translateText(&fakeTranslator{err: boom}, "en", "de", "hello")
|
||||
Expect(err).To(MatchError(boom))
|
||||
})
|
||||
|
||||
// A call that returned OK with no translations is a runtime bug, and the
|
||||
// empty string it would hand back reaches the user as a successful but blank
|
||||
// completion with nothing anywhere to say why.
|
||||
It("refuses a result that carries no translation", func() {
|
||||
_, err := translateText(&fakeTranslator{}, "en", "de", "hello")
|
||||
Expect(status.Code(err)).To(Equal(codes.Internal))
|
||||
})
|
||||
|
||||
It("takes the first translation when the runtime returns several", func() {
|
||||
f := &fakeTranslator{out: []string{"first", "second"}}
|
||||
got, err := translateText(f, "en", "de", "hello")
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(got).To(Equal("first"))
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("Predict", func() {
|
||||
It("refuses a model loaded as another family", func() {
|
||||
n := &NemoSpeech{fam: familyASR}
|
||||
out, err := n.Predict(&pb.PredictOptions{Prompt: "hello"})
|
||||
Expect(status.Code(err)).To(Equal(codes.Unimplemented))
|
||||
Expect(out).To(BeEmpty())
|
||||
|
||||
// A lock leaked on the rejection path deadlocks the next request rather
|
||||
// than failing it.
|
||||
Expect(n.engineMu.TryLock()).To(BeTrue())
|
||||
n.engineMu.Unlock()
|
||||
})
|
||||
|
||||
It("refuses an unloaded model", func() {
|
||||
n := &NemoSpeech{}
|
||||
_, err := n.Predict(&pb.PredictOptions{Prompt: "hello"})
|
||||
Expect(status.Code(err)).To(Equal(codes.Unimplemented))
|
||||
})
|
||||
|
||||
// The validation has to happen before anything crosses the ABI: nothing is
|
||||
// loaded here, so a guard placed after the C call would panic on a nil
|
||||
// function value instead of failing the request.
|
||||
It("rejects an empty prompt before it reaches the runtime", func() {
|
||||
n := &NemoSpeech{fam: familyNMT, opts: loadOptions{targetLanguage: "de"}}
|
||||
var err error
|
||||
Expect(func() {
|
||||
_, err = n.Predict(&pb.PredictOptions{})
|
||||
}).ToNot(Panic())
|
||||
Expect(status.Code(err)).To(Equal(codes.InvalidArgument))
|
||||
})
|
||||
|
||||
It("rejects a request with no target language, before it reaches the runtime", func() {
|
||||
n := &NemoSpeech{fam: familyNMT}
|
||||
var err error
|
||||
Expect(func() {
|
||||
_, err = n.Predict(&pb.PredictOptions{Prompt: "hello"})
|
||||
}).ToNot(Panic())
|
||||
Expect(status.Code(err)).To(Equal(codes.InvalidArgument))
|
||||
Expect(err.Error()).To(ContainSubstring("target_language"))
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("PredictStream", func() {
|
||||
// pkg/grpc/server.go drains this channel from a goroutine and then blocks on
|
||||
// that goroutine finishing, so a channel left open does not fail the request,
|
||||
// it hangs the RPC and every request queued behind the backend lock.
|
||||
It("closes the channel when the family does not match", func() {
|
||||
n := &NemoSpeech{fam: familyTTS}
|
||||
ch := make(chan string)
|
||||
done := collectStrings(ch)
|
||||
|
||||
err := n.PredictStream(&pb.PredictOptions{Prompt: "hello"}, ch)
|
||||
Expect(status.Code(err)).To(Equal(codes.Unimplemented))
|
||||
Expect(<-done).To(BeEmpty())
|
||||
Expect(n.engineMu.TryLock()).To(BeTrue())
|
||||
n.engineMu.Unlock()
|
||||
})
|
||||
|
||||
It("closes the channel when the request is rejected", func() {
|
||||
n := &NemoSpeech{fam: familyNMT}
|
||||
ch := make(chan string)
|
||||
done := collectStrings(ch)
|
||||
|
||||
err := n.PredictStream(&pb.PredictOptions{Prompt: "hello"}, ch)
|
||||
Expect(status.Code(err)).To(Equal(codes.InvalidArgument))
|
||||
Expect(<-done).To(BeEmpty())
|
||||
})
|
||||
|
||||
It("closes the channel on an unloaded model", func() {
|
||||
n := &NemoSpeech{}
|
||||
ch := make(chan string)
|
||||
done := collectStrings(ch)
|
||||
|
||||
Expect(n.PredictStream(&pb.PredictOptions{Prompt: "hi"}, ch)).ToNot(Succeed())
|
||||
Expect(<-done).To(BeEmpty())
|
||||
})
|
||||
|
||||
// The C API has no token callback, so the whole translation is one chunk.
|
||||
// The seam is the only place that can be asserted without a model.
|
||||
It("emits the whole translation as a single chunk", func() {
|
||||
f := &fakeTranslator{out: []string{"hallo welt"}}
|
||||
ch := make(chan string)
|
||||
done := collectStrings(ch)
|
||||
|
||||
Expect(streamTranslation(f, "en", "de", "hello world", ch)).To(Succeed())
|
||||
close(ch)
|
||||
Expect(<-done).To(Equal([]string{"hallo welt"}))
|
||||
})
|
||||
|
||||
It("emits nothing when the translation fails", func() {
|
||||
f := &fakeTranslator{err: errors.New("boom")}
|
||||
ch := make(chan string)
|
||||
done := collectStrings(ch)
|
||||
|
||||
Expect(streamTranslation(f, "en", "de", "hello", ch)).ToNot(Succeed())
|
||||
close(ch)
|
||||
Expect(<-done).To(BeEmpty())
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("unsupportedPredictFields", func() {
|
||||
It("names nothing for a plain translation request", func() {
|
||||
Expect(unsupportedPredictFields(&pb.PredictOptions{Prompt: "hello"})).To(BeEmpty())
|
||||
})
|
||||
|
||||
// The sampling knobs are deliberately absent from this list: LocalAI fills
|
||||
// them in from the model config on every request, so warning about them
|
||||
// would log on every translation and say nothing.
|
||||
It("stays quiet about sampling parameters the runtime has no field for", func() {
|
||||
Expect(unsupportedPredictFields(&pb.PredictOptions{
|
||||
Prompt: "hello",
|
||||
Temperature: 0.7,
|
||||
TopP: 0.9,
|
||||
TopK: 40,
|
||||
Seed: 42,
|
||||
Tokens: 256,
|
||||
})).To(BeEmpty())
|
||||
})
|
||||
|
||||
It("names the asks the C API cannot serve at all", func() {
|
||||
got := unsupportedPredictFields(&pb.PredictOptions{
|
||||
Prompt: "hello",
|
||||
Grammar: "root ::= x",
|
||||
Tools: `[{"type":"function"}]`,
|
||||
Images: []string{"a.png"},
|
||||
Videos: []string{"a.mp4"},
|
||||
Audios: []string{"a.wav"},
|
||||
NegativePrompt: "no",
|
||||
Logprobs: 3,
|
||||
})
|
||||
Expect(got).To(ConsistOf("grammar", "tools", "images", "videos", "audios",
|
||||
"negative_prompt", "logprobs"))
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,100 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/mudler/xlog"
|
||||
)
|
||||
|
||||
// loadOptions holds the parsed model-level options. Path fields are resolved
|
||||
// against ModelOptions.ModelPath at parse time so every consumer sees an
|
||||
// absolute path.
|
||||
type loadOptions struct {
|
||||
// ASR
|
||||
vadModel string
|
||||
pncModel string
|
||||
diarModel string
|
||||
itnDir string
|
||||
languageCode string
|
||||
|
||||
// TTS
|
||||
codecModel string
|
||||
tokenizerDir string
|
||||
tnDir string
|
||||
|
||||
// NMT
|
||||
sourceLanguage string
|
||||
targetLanguage string
|
||||
|
||||
// gpu is the device index passed to the runtime's backend config.
|
||||
// -1 selects CPU, matching the C API's own sentinel.
|
||||
gpu int32
|
||||
}
|
||||
|
||||
// splitOption splits on the FIRST colon so values may themselves contain one.
|
||||
func splitOption(o string) (key, value string, ok bool) {
|
||||
i := strings.Index(o, ":")
|
||||
if i < 0 {
|
||||
return "", "", false
|
||||
}
|
||||
return strings.TrimSpace(o[:i]), strings.TrimSpace(o[i+1:]), true
|
||||
}
|
||||
|
||||
// resolve makes a relative asset path absolute against the models directory.
|
||||
// Empty stays empty so callers can test for "unset".
|
||||
func resolve(base, p string) string {
|
||||
if p == "" || filepath.IsAbs(p) {
|
||||
return p
|
||||
}
|
||||
return filepath.Join(base, p)
|
||||
}
|
||||
|
||||
// parseOptions reads the backend "key:value" option slice. Unknown keys are
|
||||
// ignored rather than rejected, so a config written for a newer backend still
|
||||
// loads on an older one.
|
||||
func parseOptions(opts []string, modelPath string) loadOptions {
|
||||
o := loadOptions{gpu: -1}
|
||||
for _, oo := range opts {
|
||||
key, value, ok := splitOption(oo)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
switch key {
|
||||
case "vad_model":
|
||||
o.vadModel = resolve(modelPath, value)
|
||||
case "pnc_model":
|
||||
o.pncModel = resolve(modelPath, value)
|
||||
case "diar_model":
|
||||
o.diarModel = resolve(modelPath, value)
|
||||
case "itn_dir":
|
||||
o.itnDir = resolve(modelPath, value)
|
||||
case "language_code":
|
||||
o.languageCode = value
|
||||
case "codec_model":
|
||||
o.codecModel = resolve(modelPath, value)
|
||||
case "tokenizer_dir":
|
||||
o.tokenizerDir = resolve(modelPath, value)
|
||||
case "tn_dir":
|
||||
o.tnDir = resolve(modelPath, value)
|
||||
case "source_language":
|
||||
o.sourceLanguage = value
|
||||
case "target_language":
|
||||
o.targetLanguage = value
|
||||
case "gpu":
|
||||
// An unknown key is ignored for forward compatibility, but a known key
|
||||
// with an unparseable value is a typo, and this one fails expensively:
|
||||
// the model still loads and still produces correct output, just on CPU
|
||||
// and far slower, with nothing anywhere to say why.
|
||||
n, err := strconv.ParseInt(value, 10, 32)
|
||||
if err != nil {
|
||||
xlog.Warn("nemo-speech-cpp: ignoring unparseable option value, falling back to CPU",
|
||||
"key", key, "value", value)
|
||||
continue
|
||||
}
|
||||
o.gpu = int32(n)
|
||||
}
|
||||
}
|
||||
return o
|
||||
}
|
||||
@@ -0,0 +1,70 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
var _ = Describe("parseOptions", func() {
|
||||
It("parses every known key", func() {
|
||||
o := parseOptions([]string{
|
||||
"vad_model:silero.gguf",
|
||||
"pnc_model:pnc.gguf",
|
||||
"diar_model:sortformer.gguf",
|
||||
"itn_dir:tn_configs",
|
||||
"language_code:es-ES",
|
||||
"codec_model:nanocodec.gguf",
|
||||
"tokenizer_dir:extracted",
|
||||
"tn_dir:tn",
|
||||
"source_language:en",
|
||||
"target_language:de",
|
||||
}, "/models")
|
||||
|
||||
Expect(o.vadModel).To(Equal("/models/silero.gguf"))
|
||||
Expect(o.pncModel).To(Equal("/models/pnc.gguf"))
|
||||
Expect(o.diarModel).To(Equal("/models/sortformer.gguf"))
|
||||
Expect(o.itnDir).To(Equal("/models/tn_configs"))
|
||||
Expect(o.languageCode).To(Equal("es-ES"))
|
||||
Expect(o.codecModel).To(Equal("/models/nanocodec.gguf"))
|
||||
Expect(o.tokenizerDir).To(Equal("/models/extracted"))
|
||||
Expect(o.tnDir).To(Equal("/models/tn"))
|
||||
Expect(o.sourceLanguage).To(Equal("en"))
|
||||
Expect(o.targetLanguage).To(Equal("de"))
|
||||
})
|
||||
|
||||
It("leaves absolute paths untouched", func() {
|
||||
o := parseOptions([]string{"vad_model:/abs/silero.gguf"}, "/models")
|
||||
Expect(o.vadModel).To(Equal("/abs/silero.gguf"))
|
||||
})
|
||||
|
||||
It("ignores unknown keys and entries without a separator", func() {
|
||||
o := parseOptions([]string{"nonsense", "unknown_key:value"}, "/models")
|
||||
Expect(o).To(Equal(loadOptions{gpu: -1}))
|
||||
})
|
||||
|
||||
It("trims whitespace around keys and values", func() {
|
||||
o := parseOptions([]string{" language_code : en-US "}, "/models")
|
||||
Expect(o.languageCode).To(Equal("en-US"))
|
||||
})
|
||||
|
||||
It("keeps a value containing a colon intact", func() {
|
||||
// URIs must survive the split on the FIRST colon.
|
||||
o := parseOptions([]string{"tokenizer_dir:/a/b:c"}, "/models")
|
||||
Expect(o.tokenizerDir).To(Equal("/a/b:c"))
|
||||
})
|
||||
|
||||
It("leaves an empty value empty so callers can detect unset", func() {
|
||||
o := parseOptions([]string{"vad_model:"}, "/models")
|
||||
Expect(o.vadModel).To(BeEmpty())
|
||||
})
|
||||
|
||||
It("defaults gpu to -1 meaning CPU", func() {
|
||||
o := parseOptions(nil, "/models")
|
||||
Expect(o.gpu).To(Equal(int32(-1)))
|
||||
})
|
||||
|
||||
It("parses an explicit gpu index", func() {
|
||||
o := parseOptions([]string{"gpu:0"}, "/models")
|
||||
Expect(o.gpu).To(Equal(int32(0)))
|
||||
})
|
||||
})
|
||||
Executable
+161
@@ -0,0 +1,161 @@
|
||||
#!/bin/bash
|
||||
#
|
||||
# Bundle the nemo-speech-cpp-grpc binary, the five nemo_speech shared objects,
|
||||
# the text-normalization stack on a WITH_NORM build, the core runtime libs
|
||||
# (libc/libstdc++/libgomp + ld.so) and the GPU runtime for the active BUILD_TYPE
|
||||
# so the package is self-contained. Mirrors backend/go/whisper/package.sh;
|
||||
# run.sh routes the (CGO_ENABLED=0) binary through lib/ld.so so the packaged
|
||||
# libc is used instead of the host's.
|
||||
#
|
||||
# Five, not three: ASR and NMT each ship a thin _c ABI shim plus the
|
||||
# implementation DSO it depends on, while TTS ships one object with no _c
|
||||
# suffix at all.
|
||||
|
||||
set -e
|
||||
|
||||
CURDIR=$(dirname "$(realpath "$0")")
|
||||
REPO_ROOT="${CURDIR}/../../.."
|
||||
|
||||
mkdir -p "$CURDIR/package/lib"
|
||||
|
||||
cp -avf "$CURDIR/nemo-speech-cpp-grpc" "$CURDIR/package/"
|
||||
cp -avf "$CURDIR/run.sh" "$CURDIR/package/"
|
||||
|
||||
# The runtime ships three C ABI shared objects, not one. All three are
|
||||
# required: main.go dlopens them eagerly, so a package missing any of them
|
||||
# fails at startup. ASR and NMT expose the ABI through a dedicated _c library;
|
||||
# TTS compiles its c_api into libnemo_speech_tts itself and has no _c variant,
|
||||
# hence the asymmetric list. purego.Dlopen resolves them via the
|
||||
# NEMO_SPEECH_*_LIBRARY paths that run.sh points at lib/.
|
||||
#
|
||||
# libnemo_speech_asr and libnemo_speech_nmt are in the list because the matching
|
||||
# _c shims carry a DT_NEEDED on them: dlopen of the shim fails without the
|
||||
# implementation DSO alongside it.
|
||||
for lib in libnemo_speech_asr_c libnemo_speech_asr libnemo_speech_tts libnemo_speech_nmt_c libnemo_speech_nmt; do
|
||||
cp -avf "$CURDIR"/${lib}.so* "$CURDIR/package/lib/" 2>/dev/null || true
|
||||
cp -avf "$CURDIR"/${lib}*.dylib "$CURDIR/package/lib/" 2>/dev/null || true
|
||||
if ! ls "$CURDIR"/package/lib/${lib}.* >/dev/null 2>&1; then
|
||||
echo "ERROR: ${lib} shared library not found in $CURDIR, run 'make' first" >&2
|
||||
exit 1
|
||||
fi
|
||||
done
|
||||
|
||||
# Text normalization (WITH_NORM=ON, Linux only) links Sparrowhawk and OpenFST
|
||||
# into libnemo_speech_asr.so. Those live in a project-local prefix that the
|
||||
# Makefile stages here, so anything staged that is not a nemo_speech object is
|
||||
# part of that stack. Absent on a WITH_NORM=OFF build, which is why this is a
|
||||
# glob that tolerates no matches rather than a required list.
|
||||
shopt -s nullglob
|
||||
for so in "$CURDIR"/*.so "$CURDIR"/*.so.* "$CURDIR"/*.dylib; do
|
||||
case "$(basename "$so")" in
|
||||
libnemo_speech_*) continue ;;
|
||||
esac
|
||||
cp -avf "$so" "$CURDIR/package/lib/"
|
||||
done
|
||||
shopt -u nullglob
|
||||
|
||||
# Detect architecture and copy the core runtime libs the shared objects link
|
||||
# against, plus the matching dynamic loader as lib/ld.so.
|
||||
source "$CURDIR/../../../scripts/build/package-system-libs.sh" "$CURDIR/package/lib" ""
|
||||
|
||||
# Dependency-closure guard.
|
||||
#
|
||||
# The lists above are maintained by hand, and the WITH_NORM build in particular
|
||||
# pulls in transitive dependencies nobody enumerated: Sparrowhawk drags in
|
||||
# protobuf, re2 and absl, none of which package-system-libs.sh provides. Rather
|
||||
# than hard-code that set, walk the DT_NEEDED entries of everything staged and
|
||||
# copy whatever is still unresolved. On a WITH_NORM=OFF build the closure is
|
||||
# already complete, so this copies nothing.
|
||||
#
|
||||
# Skipped deliberately: the core runtime set that package-system-libs.sh owns,
|
||||
# and the GPU stack that package-gpu-libs.sh owns.
|
||||
shopt -s nullglob
|
||||
staged_libs=("$CURDIR"/package/lib/*.so*)
|
||||
shopt -u nullglob
|
||||
|
||||
if [ "$(uname)" != "Darwin" ] && [ "${#staged_libs[@]}" -gt 0 ]; then
|
||||
# No silent skip. If the closure cannot be checked, the package cannot be
|
||||
# shown to be complete, and shipping an unverified one is the failure this
|
||||
# guard exists to prevent.
|
||||
if command -v readelf >/dev/null 2>&1; then
|
||||
read_needed() { readelf -d "$1" 2>/dev/null | sed -n 's/.*(NEEDED).*\[\(.*\)\]/\1/p'; }
|
||||
elif command -v objdump >/dev/null 2>&1; then
|
||||
read_needed() { objdump -p "$1" 2>/dev/null | awk '$1 == "NEEDED" { print $2 }'; }
|
||||
else
|
||||
echo "ERROR: neither readelf nor objdump is available, so the dependency" >&2
|
||||
echo " closure of ${#staged_libs[@]} staged libraries cannot be verified." >&2
|
||||
echo " Install binutils in the build image; refusing to ship an" >&2
|
||||
echo " unverified package." >&2
|
||||
exit 1
|
||||
fi
|
||||
|
||||
is_provided() {
|
||||
case "$1" in
|
||||
ld-linux*|libc.so.6|libstdc++.so.6|libgcc_s.so.1|libm.so.6|libgomp.so.1) return 0 ;;
|
||||
libdl.so.2|librt.so.1|libpthread.so.0) return 0 ;;
|
||||
libcuda*|libcudart*|libcublas*|libcublasLt*|libnvrtc*|libnvidia*) return 0 ;;
|
||||
libamdhip*|libhsa*|librocm*|libze_*|libOpenCL*|libvulkan*) return 0 ;;
|
||||
esac
|
||||
[ -e "$CURDIR/package/lib/$1" ]
|
||||
}
|
||||
|
||||
# Walk until the staged set stops growing. The glob below expands once per
|
||||
# pass, so each pass advances the closure by exactly one dependency level;
|
||||
# a copied library can itself pull in new dependencies.
|
||||
#
|
||||
# CLOSURE_MAX_PASSES is a runaway guard, not a depth limit. Exhausting it
|
||||
# means the walk never converged and the package is therefore incomplete,
|
||||
# which has to fail the build: a fixed pass count that just falls out of the
|
||||
# loop would silently ship a package missing its deepest libraries, and
|
||||
# libnemo_speech_asr -> sparrowhawk -> protobuf -> absl already runs several
|
||||
# levels deep.
|
||||
CLOSURE_MAX_PASSES="${CLOSURE_MAX_PASSES:-64}"
|
||||
converged=0
|
||||
for (( pass=1; pass<=CLOSURE_MAX_PASSES; pass++ )); do
|
||||
missing=0
|
||||
for so in "$CURDIR"/package/lib/*.so*; do
|
||||
[ -f "$so" ] || continue
|
||||
for need in $(read_needed "$so"); do
|
||||
# Written as an if rather than "is_provided && continue" so a
|
||||
# false return cannot trip set -e via the AND-list exit status.
|
||||
if is_provided "$need"; then
|
||||
continue
|
||||
fi
|
||||
# Resolve against the staging dir first, then the system loader.
|
||||
src="$(LD_LIBRARY_PATH="$CURDIR:$CURDIR/package/lib:${LD_LIBRARY_PATH:-}" \
|
||||
ldd "$so" 2>/dev/null | awk -v n="$need" '$1 == n { print $3 }' | head -1)"
|
||||
if [ -z "$src" ] || [ ! -e "$src" ]; then
|
||||
echo "ERROR: $(basename "$so") needs $need and it could not be resolved." >&2
|
||||
echo " The packaged backend would fail to dlopen at runtime." >&2
|
||||
exit 1
|
||||
fi
|
||||
cp -aLvf "$src" "$CURDIR/package/lib/$need"
|
||||
missing=1
|
||||
done
|
||||
done
|
||||
if [ "$missing" -eq 0 ]; then
|
||||
converged=1
|
||||
break
|
||||
fi
|
||||
done
|
||||
|
||||
if [ "$converged" -ne 1 ]; then
|
||||
echo "ERROR: the dependency closure was still growing after" >&2
|
||||
echo " $CLOSURE_MAX_PASSES passes, so the package is incomplete and" >&2
|
||||
echo " would fail to dlopen at runtime. Refusing to ship it." >&2
|
||||
exit 1
|
||||
fi
|
||||
fi
|
||||
|
||||
# Package GPU libraries (CUDA/ROCm/Intel/Vulkan loader + ICDs + drivers)
|
||||
# based on BUILD_TYPE so the backend can reach the GPU without the runtime
|
||||
# base image shipping those drivers.
|
||||
GPU_LIB_SCRIPT="${REPO_ROOT}/scripts/build/package-gpu-libs.sh"
|
||||
if [ -f "$GPU_LIB_SCRIPT" ]; then
|
||||
echo "Packaging GPU libraries for BUILD_TYPE=${BUILD_TYPE:-cpu}..."
|
||||
source "$GPU_LIB_SCRIPT" "$CURDIR/package/lib"
|
||||
package_gpu_libs
|
||||
fi
|
||||
|
||||
echo "Packaging completed successfully"
|
||||
ls -liah "$CURDIR/package/" "$CURDIR/package/lib/"
|
||||
Executable
+28
@@ -0,0 +1,28 @@
|
||||
#!/bin/bash
|
||||
set -e
|
||||
|
||||
CURDIR=$(dirname "$(realpath "$0")")
|
||||
|
||||
# The runtime splits its C ABI across three shared objects, so each gets its
|
||||
# own override variable. main.go reads exactly these names.
|
||||
if [ "$(uname)" = "Darwin" ]; then
|
||||
export DYLD_LIBRARY_PATH="$CURDIR/lib:"$CURDIR":${DYLD_LIBRARY_PATH:-}"
|
||||
export NEMO_SPEECH_ASR_LIBRARY="$CURDIR/lib/libnemo_speech_asr_c.dylib"
|
||||
export NEMO_SPEECH_TTS_LIBRARY="$CURDIR/lib/libnemo_speech_tts.dylib"
|
||||
export NEMO_SPEECH_NMT_LIBRARY="$CURDIR/lib/libnemo_speech_nmt_c.dylib"
|
||||
else
|
||||
export LD_LIBRARY_PATH="$CURDIR/lib:"$CURDIR":${LD_LIBRARY_PATH:-}"
|
||||
export NEMO_SPEECH_ASR_LIBRARY="$CURDIR/lib/libnemo_speech_asr_c.so"
|
||||
export NEMO_SPEECH_TTS_LIBRARY="$CURDIR/lib/libnemo_speech_tts.so"
|
||||
export NEMO_SPEECH_NMT_LIBRARY="$CURDIR/lib/libnemo_speech_nmt_c.so"
|
||||
fi
|
||||
|
||||
# If a self-contained ld.so was packaged, route through it so the
|
||||
# packaged libc / libstdc++ are used instead of the host's (matches the
|
||||
# whisper backend's runtime layout). Linux only.
|
||||
if [ -f "$CURDIR/lib/ld.so" ]; then
|
||||
echo "Using lib/ld.so"
|
||||
exec "$CURDIR/lib/ld.so" "$CURDIR/nemo-speech-cpp-grpc" "$@"
|
||||
fi
|
||||
|
||||
exec "$CURDIR/nemo-speech-cpp-grpc" "$@"
|
||||
@@ -0,0 +1,102 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
)
|
||||
|
||||
// The status values every C entry point in this backend returns.
|
||||
//
|
||||
// There is no single C enum to mirror. asr.h:43-49, tts.h:38-44 and nmt.h:40-45
|
||||
// each declare their own, and diar.h has none of its own at all: it includes
|
||||
// asr.h and types every diarization function as nemo_speech_asr_status
|
||||
// (diar.h:18). The names the three surfaces share carry the same numbers:
|
||||
//
|
||||
// value asr.h tts.h nmt.h
|
||||
// 0 NEMO_SPEECH_ASR_OK NEMO_SPEECH_TTS_OK NEMO_SPEECH_NMT_OK
|
||||
// 1 NEMO_SPEECH_ASR_ERROR_INVALID_... NEMO_SPEECH_TTS_ERROR_INVALID_... NEMO_SPEECH_NMT_ERROR_INVALID_...
|
||||
// 2 NEMO_SPEECH_ASR_ERROR_OUT_OF_MEM. NEMO_SPEECH_TTS_ERROR_OUT_OF_MEM. NEMO_SPEECH_NMT_ERROR_OUT_OF_MEM.
|
||||
// 3 NEMO_SPEECH_ASR_ERROR_RUNTIME NEMO_SPEECH_TTS_ERROR_RUNTIME NEMO_SPEECH_NMT_ERROR_RUNTIME
|
||||
// 4 NEMO_SPEECH_ASR_ERROR_CANCELLED NEMO_SPEECH_TTS_ERROR_CANCELLED (not declared)
|
||||
//
|
||||
// The one divergence is 4, and it is an absence rather than a disagreement. ASR
|
||||
// and TTS both drive a consumer callback that can ask for the work to stop, and
|
||||
// cancellation is what they report when it does; nemo_speech_nmt_translate takes
|
||||
// no callback and returns only when the decode has finished, so the NMT surface
|
||||
// has no cancellation to name. That is why one table can serve all three: 4 is
|
||||
// not some other NMT status that would be mislabelled, it is a value the NMT
|
||||
// surface never produces.
|
||||
//
|
||||
// Recheck this table after an upstream pin bump. A status added to one header
|
||||
// and not the others is exactly the shape of change that would break the single
|
||||
// mapping, and nothing in the build or the linker can see it: purego binds by
|
||||
// name, and the return value is a bare int32 on the Go side.
|
||||
const (
|
||||
statusOK int32 = 0
|
||||
statusInvalidArgument int32 = 1
|
||||
statusOutOfMemory int32 = 2
|
||||
statusRuntime int32 = 3
|
||||
statusCancelled int32 = 4
|
||||
)
|
||||
|
||||
// statusCode maps a C status onto the gRPC code the caller should be told.
|
||||
//
|
||||
// What the mapping is really carrying is whose mistake the failure was.
|
||||
// INVALID_ARGUMENT is what every guard in src/{asr,tts,nmt}/c_api.cpp returns
|
||||
// for a std::invalid_argument from the runtime, and the things that throw it are
|
||||
// requests: an unknown voice_name (src/tts/synthesizer.cpp), an unsupported
|
||||
// language pair (src/nmt/translator.cpp), an out-of-range sample rate. Reporting
|
||||
// those as Internal turns a 400 into a 500 and sends the user hunting for a
|
||||
// broken model or a broken backend instead of fixing the request.
|
||||
//
|
||||
// OUT_OF_MEMORY is a resource limit rather than a defect, which is what
|
||||
// ResourceExhausted means, and it is the one failure a client can sensibly
|
||||
// retry later or retry smaller. CANCELLED is the consumer having stopped
|
||||
// listening, which is not a failure of this backend at all: the streaming sinks
|
||||
// return false once their client is gone (see ttsDeliverPCM), and the runtime
|
||||
// turns that into status 4.
|
||||
//
|
||||
// RUNTIME, and anything a future pin adds that this table has not been taught,
|
||||
// stay Internal. An unrecognised status is precisely the case where the backend
|
||||
// does not know whose fault it was, and Internal is the honest answer.
|
||||
func statusCode(st int32) codes.Code {
|
||||
switch st {
|
||||
case statusOK:
|
||||
return codes.OK
|
||||
case statusInvalidArgument:
|
||||
return codes.InvalidArgument
|
||||
case statusOutOfMemory:
|
||||
return codes.ResourceExhausted
|
||||
case statusCancelled:
|
||||
return codes.Canceled
|
||||
case statusRuntime:
|
||||
// Named rather than folded into the default so this switch reads as the
|
||||
// whole enum. A status the table has never heard of is a different thing
|
||||
// from a runtime error even though both answer Internal, and a reader
|
||||
// checking the mapping against the headers should not have to work out
|
||||
// which arm RUNTIME lands in.
|
||||
return codes.Internal
|
||||
default:
|
||||
return codes.Internal
|
||||
}
|
||||
}
|
||||
|
||||
// statusErrorf builds the gRPC error for a failed C call.
|
||||
//
|
||||
// Every C call site in this backend goes through this rather than through
|
||||
// status.Errorf directly, and that is the whole point of it existing: the
|
||||
// mapping used to be written out at exactly one of sixteen call sites, so the
|
||||
// same backend answered an unsupported language pair with InvalidArgument and an
|
||||
// unknown TTS voice, which is the same class of caller mistake against the same
|
||||
// process, with Internal.
|
||||
//
|
||||
// The OK guard is not defensive noise. status.Errorf(codes.OK, ...) returns a
|
||||
// nil error, so a call site that built its error without first checking the
|
||||
// status would report a hard C failure as a successful request with no
|
||||
// diagnostic anywhere. Returning Internal instead keeps that mistake loud.
|
||||
func statusErrorf(st int32, format string, args ...any) error {
|
||||
if st == statusOK {
|
||||
return status.Errorf(codes.Internal, format, args...)
|
||||
}
|
||||
return status.Errorf(statusCode(st), format, args...)
|
||||
}
|
||||
@@ -0,0 +1,109 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"os"
|
||||
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
|
||||
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
|
||||
)
|
||||
|
||||
var _ = Describe("C status mapping", func() {
|
||||
// The whole enum, so a status that quietly moves to a different code is
|
||||
// visible here rather than in a bug report about an HTTP 500. The names on
|
||||
// the left are transcribed from asr.h:43-49, tts.h:38-44 and nmt.h:40-45;
|
||||
// see the table in status.go for how the three surfaces line up.
|
||||
DescribeTable("maps each declared C status onto a gRPC code",
|
||||
func(st int32, want codes.Code) {
|
||||
Expect(statusCode(st)).To(Equal(want))
|
||||
},
|
||||
Entry("OK", statusOK, codes.OK),
|
||||
Entry("INVALID_ARGUMENT", statusInvalidArgument, codes.InvalidArgument),
|
||||
Entry("OUT_OF_MEMORY", statusOutOfMemory, codes.ResourceExhausted),
|
||||
Entry("RUNTIME", statusRuntime, codes.Internal),
|
||||
Entry("CANCELLED", statusCancelled, codes.Canceled),
|
||||
)
|
||||
|
||||
// A pin bump that adds a status this table has never been taught must not
|
||||
// guess. Internal is the honest answer when the backend does not know whose
|
||||
// mistake the failure was.
|
||||
DescribeTable("reports an unknown status as Internal",
|
||||
func(st int32) {
|
||||
Expect(statusCode(st)).To(Equal(codes.Internal))
|
||||
},
|
||||
Entry("one past the last declared value", int32(5)),
|
||||
Entry("far past it", int32(99)),
|
||||
Entry("negative", int32(-1)),
|
||||
)
|
||||
|
||||
It("carries the mapped code and the formatted message into the error", func() {
|
||||
err := statusErrorf(statusInvalidArgument, "nemo-speech-cpp: %s: %d", "synthesize", 7)
|
||||
Expect(status.Code(err)).To(Equal(codes.InvalidArgument))
|
||||
Expect(err.Error()).To(ContainSubstring("nemo-speech-cpp: synthesize: 7"))
|
||||
})
|
||||
|
||||
// status.Errorf(codes.OK, ...) returns nil, so a call site that built its
|
||||
// error without checking the status first would turn a hard C failure into a
|
||||
// silent success with no diagnostic anywhere.
|
||||
It("never returns nil, not even for OK", func() {
|
||||
err := statusErrorf(statusOK, "nemo-speech-cpp: should not happen")
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(status.Code(err)).To(Equal(codes.Internal))
|
||||
})
|
||||
})
|
||||
|
||||
// These drive real C statuses out of the real shared objects, one per family,
|
||||
// rather than asserting the Go mapping against itself.
|
||||
//
|
||||
// A NULL handle is the one failure every surface can be provoked into without a
|
||||
// model: nemo_speech_asr_recognize_f32 and nemo_speech_nmt_translate check the
|
||||
// handle up front, nemo_speech_tts_synthesize_text does the same, and the
|
||||
// diarization stream entry points throw std::invalid_argument for a dead stream,
|
||||
// which src/asr/c_api.cpp's guard maps to the same status. All four are the
|
||||
// caller's mistake, and the point of the exercise is that all four now come back
|
||||
// as InvalidArgument instead of Internal.
|
||||
var _ = Describe("C status mapping at the call sites", func() {
|
||||
BeforeEach(func() {
|
||||
if !librariesPresent() {
|
||||
if requireLibs() {
|
||||
cwd, _ := os.Getwd()
|
||||
Fail("NEMO_SPEECH_REQUIRE_LIBS=1 but the shared libraries are not in " + cwd +
|
||||
": these specs are the ABI defence and must not be skipped." +
|
||||
" Run make -C backend/go/nemo-speech-cpp stage-libs")
|
||||
}
|
||||
Skip("shared libraries not built, run make in backend/go/nemo-speech-cpp")
|
||||
}
|
||||
Expect(openLibraries()).To(Succeed())
|
||||
})
|
||||
|
||||
It("reports an ASR INVALID_ARGUMENT as InvalidArgument", func() {
|
||||
opts := ASRRecognitionOptionsDef()
|
||||
// Non-empty PCM on purpose: recognizeF32 rejects an empty slice itself,
|
||||
// which would prove nothing about what the C side returned.
|
||||
_, err := recognizeF32(0, &opts, []float32{0, 0, 0}, 16000)
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(status.Code(err)).To(Equal(codes.InvalidArgument))
|
||||
})
|
||||
|
||||
It("reports a diarization INVALID_ARGUMENT as InvalidArgument", func() {
|
||||
err := (&cDiarStream{}).finish()
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(status.Code(err)).To(Equal(codes.InvalidArgument))
|
||||
})
|
||||
|
||||
It("reports a TTS INVALID_ARGUMENT as InvalidArgument", func() {
|
||||
s := &cSynthesizer{}
|
||||
err := s.synthesize(&pb.TTSRequest{Text: "hello"}, "en", func([]byte) bool { return true })
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(status.Code(err)).To(Equal(codes.InvalidArgument))
|
||||
})
|
||||
|
||||
It("reports an NMT INVALID_ARGUMENT as InvalidArgument", func() {
|
||||
_, err := (&cTranslator{}).translate([]string{"hello"}, "en", "de")
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(status.Code(err)).To(Equal(codes.InvalidArgument))
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,600 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"math"
|
||||
"os"
|
||||
"runtime"
|
||||
"strconv"
|
||||
"sync"
|
||||
"unsafe"
|
||||
|
||||
"github.com/ebitengine/purego"
|
||||
laudio "github.com/mudler/LocalAI/pkg/audio"
|
||||
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
|
||||
"github.com/mudler/xlog"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
)
|
||||
|
||||
// The backend-preference enum from include/nemo_speech/tts.h. The C type is an
|
||||
// enum, which this toolchain lays out as int32, so the values are written here
|
||||
// rather than inferred.
|
||||
const (
|
||||
ttsBackendAuto int32 = 0
|
||||
ttsBackendCPU int32 = 1
|
||||
)
|
||||
|
||||
// maxWAVDataBytes is the largest PCM payload a RIFF WAV can describe.
|
||||
//
|
||||
// Both size fields in the header are uint32, so a longer payload would not
|
||||
// merely be unusual, it would wrap and produce a file whose header disagrees
|
||||
// with its contents. At 22.05 kHz mono 16-bit that ceiling is about 27 hours of
|
||||
// speech, so nothing real is being refused.
|
||||
//
|
||||
// Typed int64 rather than left untyped so the comparison below is the same one
|
||||
// on every architecture: an untyped constant this large does not fit in a
|
||||
// 32-bit int and would not compile there at all.
|
||||
const maxWAVDataBytes int64 = math.MaxUint32 - laudio.WAVHeaderSize
|
||||
|
||||
// wavStreamingSize is the placeholder both size fields carry while the total
|
||||
// length is still unknown. Players read it as "stream until the socket closes".
|
||||
const wavStreamingSize = 0xFFFFFFFF
|
||||
|
||||
// ttsSink receives one PCM chunk, already copied into Go memory.
|
||||
//
|
||||
// It returns false to cancel the synthesis in progress: that is the C
|
||||
// callback's only way to stop work early, and the runtime turns it into
|
||||
// NEMO_SPEECH_TTS_ERROR_CANCELLED.
|
||||
type ttsSink func(pcm []byte) bool
|
||||
|
||||
// ttsSinkTable maps the user_data value handed to C back to the Go sink the
|
||||
// chunk belongs to.
|
||||
//
|
||||
// A single "current sink" pointer would be enough for one model, since every
|
||||
// RPC holds that model's engineMu for its whole body. It is not enough for the
|
||||
// process: engineMu is per-NemoSpeech, one backend process can hold several
|
||||
// loaded models, and the callback below is shared by all of them, so two TTS
|
||||
// models synthesizing at once would overwrite each other's sink. The id is what
|
||||
// keeps them apart.
|
||||
//
|
||||
// The id is an integer and never a Go pointer. user_data crosses into C as a
|
||||
// void*, which the collector does not trace, so a Go pointer parked there would
|
||||
// have exactly the lifetime problem cstr documents.
|
||||
type ttsSinkTable struct {
|
||||
mu sync.Mutex
|
||||
next uintptr
|
||||
sinks map[uintptr]ttsSink
|
||||
}
|
||||
|
||||
var ttsSinks = &ttsSinkTable{sinks: map[uintptr]ttsSink{}}
|
||||
|
||||
// register adds sink and returns its id together with the release the caller
|
||||
// MUST defer. Ids start at 1 so a zeroed or stale user_data cannot resolve to
|
||||
// somebody else's sink.
|
||||
func (t *ttsSinkTable) register(sink ttsSink) (uintptr, func()) {
|
||||
t.mu.Lock()
|
||||
defer t.mu.Unlock()
|
||||
|
||||
t.next++
|
||||
id := t.next
|
||||
t.sinks[id] = sink
|
||||
return id, func() {
|
||||
t.mu.Lock()
|
||||
defer t.mu.Unlock()
|
||||
delete(t.sinks, id)
|
||||
}
|
||||
}
|
||||
|
||||
// lookup returns the sink for id, or nil once it has been released.
|
||||
func (t *ttsSinkTable) lookup(id uintptr) ttsSink {
|
||||
t.mu.Lock()
|
||||
defer t.mu.Unlock()
|
||||
|
||||
return t.sinks[id]
|
||||
}
|
||||
|
||||
var (
|
||||
ttsCallbackOnce sync.Once
|
||||
ttsCallbackFn uintptr
|
||||
)
|
||||
|
||||
// ttsPCMCallback returns the C function pointer the runtime drives PCM through,
|
||||
// compiling it on first use.
|
||||
//
|
||||
// Exactly one is ever created per process, and that is a hard requirement
|
||||
// rather than a tidiness argument. purego.NewCallback writes into a fixed table
|
||||
// of maxCB = 2000 entries (purego/syscall_sysv.go) and never releases an entry,
|
||||
// so a callback compiled per request panics the whole backend process with
|
||||
// "purego: the maximum number of callbacks has been reached" on the 2001st
|
||||
// synthesis. Per model load is not safe either: a server that swaps models
|
||||
// reaches the same ceiling, just later and even less predictably. Routing every
|
||||
// synthesis through one callback plus a user_data id is what keeps the count at
|
||||
// one for the life of the process.
|
||||
func ttsPCMCallback() uintptr {
|
||||
ttsCallbackOnce.Do(func() { ttsCallbackFn = purego.NewCallback(ttsDeliverPCM) })
|
||||
return ttsCallbackFn
|
||||
}
|
||||
|
||||
// ttsDeliverPCM is the body of that callback: nemo_speech_tts_pcm_callback,
|
||||
// which the runtime invokes on its own thread for each chunk it produces.
|
||||
//
|
||||
// The bytes are copied rather than aliased. The pointer addresses a std::string
|
||||
// the runtime owns and reuses for the next chunk (src/tts/c_api.cpp
|
||||
// make_callback), so a slice over it would be rewritten under the consumer as
|
||||
// soon as this returns.
|
||||
func ttsDeliverPCM(pcm unsafe.Pointer, nBytes uint64, userData uintptr) bool {
|
||||
// The table's lock is released before the sink runs, which matters because
|
||||
// TTSStream's sink blocks on a channel send until its client drains it.
|
||||
// Holding the lock across that would stall every other model's callback
|
||||
// behind one slow consumer.
|
||||
sink := ttsSinks.lookup(userData)
|
||||
if sink == nil {
|
||||
// The request that registered this sink has already returned, so there
|
||||
// is nowhere to put the audio. false cancels rather than letting the
|
||||
// runtime synthesize to completion into a consumer that stopped
|
||||
// listening.
|
||||
return false
|
||||
}
|
||||
// c_api.cpp filters empty chunks before calling us, so this is belt and
|
||||
// braces: unsafe.Slice on a null pointer is what it protects against.
|
||||
if pcm == nil || nBytes == 0 {
|
||||
return true
|
||||
}
|
||||
|
||||
buf := make([]byte, nBytes)
|
||||
// #nosec G103 -- pcm and nBytes are the C-owned buffer and its length from
|
||||
// one callback invocation, both null/zero-checked above. The slice is read
|
||||
// only, its length is the length the runtime declared for that buffer, and it
|
||||
// is copied into Go memory here and never retained past this return.
|
||||
copy(buf, unsafe.Slice((*byte)(pcm), nBytes))
|
||||
return sink(buf)
|
||||
}
|
||||
|
||||
// synthesizer is the TTS half of the C API, narrowed to what the two RPCs use.
|
||||
//
|
||||
// It is an interface for the same reason asrSession and diarStream are: no
|
||||
// MagpieTTS GGUF is small enough to keep in the tree, so the logic layered on
|
||||
// top of the ABI (validation, the WAV framing, chunk ordering) would otherwise
|
||||
// have no test at all. The seam is at the ABI, not at the model: a fake here
|
||||
// scripts what the C API emits, it does not pretend to synthesize anything.
|
||||
type synthesizer interface {
|
||||
// sampleRate is the rate the PCM chunks arrive at.
|
||||
sampleRate() int32
|
||||
// synthesize maps req onto the runtime's per-request options and runs one
|
||||
// synthesis, handing each chunk to sink as it is produced.
|
||||
synthesize(req *pb.TTSRequest, defaultLanguage string, sink ttsSink) error
|
||||
}
|
||||
|
||||
// cSynthesizer is the real synthesizer, over one nemo_speech_tts_synthesizer.
|
||||
type cSynthesizer struct {
|
||||
handle uintptr
|
||||
}
|
||||
|
||||
func (s *cSynthesizer) sampleRate() int32 { return TTSSampleRate(s.handle) }
|
||||
|
||||
func (s *cSynthesizer) synthesize(req *pb.TTSRequest, defaultLanguage string, sink ttsSink) error {
|
||||
// Started from the runtime's own defaults, not from a zero struct: every
|
||||
// numeric field here is sentinel-sensitive (speaker/seed < 0, steps/top_k
|
||||
// <= 0 all mean "use the synthesizer's value"), and a zeroed struct would
|
||||
// read as speaker 0, seed 0 and zero decoding steps.
|
||||
opts := TTSSynthesisOptionsDefault()
|
||||
|
||||
// A per-request language wins over the model-level default; both may be
|
||||
// empty, which the runtime resolves to the synthesizer's own default.
|
||||
language := req.GetLanguage()
|
||||
if language == "" {
|
||||
language = defaultLanguage
|
||||
}
|
||||
langP, freeLang := cstr(language)
|
||||
defer freeLang()
|
||||
opts.LanguageCode = langP
|
||||
|
||||
speaker, voiceName := resolveSpeaker(req.GetVoice())
|
||||
opts.Speaker = speaker
|
||||
voiceP, freeVoice := cstr(voiceName)
|
||||
defer freeVoice()
|
||||
opts.VoiceName = voiceP
|
||||
|
||||
applySynthesisParams(&opts, req.GetParams())
|
||||
|
||||
id, release := ttsSinks.register(sink)
|
||||
defer release()
|
||||
|
||||
// stats_out is NULL: nemo_speech_tts_synthesis_stats is 300-odd bytes of
|
||||
// timing detail with nowhere to go on either RPC, and the C API documents
|
||||
// NULL as the way to decline it.
|
||||
// #nosec G103 -- opts is a local POD struct borrowed for this call only. Its
|
||||
// two uintptr members (LanguageCode, VoiceName) are cstr allocations pinned
|
||||
// by the defers above, and this entry point is synchronous, so it returns
|
||||
// before those pins are released even though the callbacks run off-thread.
|
||||
st := TTSSynthesizeText(s.handle, unsafe.Pointer(&opts), req.GetText(), ttsPCMCallback(), id, nil)
|
||||
if st != 0 {
|
||||
// An unknown voice_name arrives here as INVALID_ARGUMENT
|
||||
// (src/tts/synthesizer.cpp throws std::invalid_argument, which
|
||||
// src/tts/c_api.cpp's guard maps to it), and a consumer that stopped
|
||||
// reading arrives as CANCELLED. Neither is this backend's failure, so
|
||||
// neither goes out as Internal.
|
||||
return statusErrorf(st, "nemo-speech-cpp: synthesize: %s", TTSLastError())
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// resolveSpeaker splits a request's voice into the two fields the C API has for
|
||||
// it: a speaker index and a voice name.
|
||||
//
|
||||
// nemo_speech_tts_synthesis_options.voice_name is documented as ignored
|
||||
// whenever speaker >= 0, and src/tts/synthesizer.cpp only calls resolve_speaker
|
||||
// when options.speaker is negative, so the two are alternatives and never a
|
||||
// pair. A named voice must therefore leave the index at -1 or the name is
|
||||
// silently dropped.
|
||||
//
|
||||
// The numeric split cannot change what the runtime picks: resolve_speaker parses
|
||||
// a numeric voice_name itself, so anything this function passes through as a
|
||||
// name and that happens to be a number lands on the same speaker anyway. What it
|
||||
// must not do is let a NEGATIVE number through as an index. "-1" is not a
|
||||
// speaker, it is the sentinel for "use the default", and treating it as an index
|
||||
// would turn a request naming an invalid voice into one that quietly synthesizes
|
||||
// in the default voice instead of being rejected.
|
||||
func resolveSpeaker(voice string) (int32, string) {
|
||||
if voice == "" {
|
||||
return -1, ""
|
||||
}
|
||||
if idx, err := strconv.ParseInt(voice, 10, 32); err == nil && idx >= 0 {
|
||||
return int32(idx), ""
|
||||
}
|
||||
return -1, voice
|
||||
}
|
||||
|
||||
// applySynthesisParams maps TTSRequest.params onto the runtime's per-request
|
||||
// options.
|
||||
//
|
||||
// Only the five knobs nemo_speech_tts_synthesis_options actually has are read.
|
||||
// An unset or unparseable value leaves the field alone rather than resetting it:
|
||||
// the struct arrives carrying the runtime's defaults, and params is documented
|
||||
// as "unset leaves the backend's configured defaults".
|
||||
//
|
||||
// The sentinels are the reason each write is guarded rather than unconditional.
|
||||
// src/tts/magpietts/runtime.cpp takes the request's seed only when it is >= 0
|
||||
// and its steps and top_k only when they are > 0, so writing a parsed 0 or a
|
||||
// negative would not merely be ignored, it would erase the option's meaning for
|
||||
// a caller who passed "0" expecting something.
|
||||
//
|
||||
// temperature and cfg_scale each need their override flag set as well. The
|
||||
// runtime reads the float only when the flag is true and otherwise falls back to
|
||||
// the synthesizer's config, so a temperature written without its flag is
|
||||
// silently discarded.
|
||||
func applySynthesisParams(o *cTTSSynthesisOptions, params map[string]string) {
|
||||
if len(params) == 0 {
|
||||
return
|
||||
}
|
||||
if v, ok := parseInt32Param(params["seed"]); ok && v >= 0 {
|
||||
o.Seed = v
|
||||
}
|
||||
if v, ok := parseInt32Param(params["steps"]); ok && v > 0 {
|
||||
o.Steps = v
|
||||
}
|
||||
if v, ok := parseInt32Param(params["top_k"]); ok && v > 0 {
|
||||
o.TopK = v
|
||||
}
|
||||
if v, ok := parseFloat32Param(params["temperature"]); ok {
|
||||
o.Temperature = v
|
||||
o.OverrideTemperature = true
|
||||
}
|
||||
if v, ok := parseFloat32Param(params["cfg_scale"]); ok {
|
||||
o.CFGScale = v
|
||||
o.OverrideCFGScale = true
|
||||
}
|
||||
}
|
||||
|
||||
// parseInt32Param reads one params entry. ok is false for an absent or
|
||||
// unparseable value, which the caller reads as "leave the default".
|
||||
func parseInt32Param(v string) (int32, bool) {
|
||||
if v == "" {
|
||||
return 0, false
|
||||
}
|
||||
n, err := strconv.ParseInt(v, 10, 32)
|
||||
if err != nil {
|
||||
xlog.Warn("nemo-speech-cpp: ignoring unparseable TTS parameter", "value", v)
|
||||
return 0, false
|
||||
}
|
||||
return int32(n), true
|
||||
}
|
||||
|
||||
func parseFloat32Param(v string) (float32, bool) {
|
||||
if v == "" {
|
||||
return 0, false
|
||||
}
|
||||
f, err := strconv.ParseFloat(v, 32)
|
||||
if err != nil {
|
||||
xlog.Warn("nemo-speech-cpp: ignoring unparseable TTS parameter", "value", v)
|
||||
return 0, false
|
||||
}
|
||||
return float32(f), true
|
||||
}
|
||||
|
||||
// ttsModelConfig builds the create-time model config.
|
||||
//
|
||||
// Extracted from loadTTS and asserted field by field because three adjacent
|
||||
// members of nemo_speech_tts_model_config are same-typed paths. Swapping two of
|
||||
// them changes neither the struct's size nor any field's offset, so the layout
|
||||
// assertions in abi_test.go cannot see it, and the failure it produces is the
|
||||
// runtime loading the codec as the acoustic model.
|
||||
//
|
||||
// Every argument is a C pointer from cstr, not a Go string, and the caller owns
|
||||
// the releases. tnDir may be null: text_normalizer_model_dir is optional and an
|
||||
// empty one leaves the text unchanged.
|
||||
func ttsModelConfig(magpieModel, codecModel, tokenizerDir, tnDir uintptr) cTTSModelConfig {
|
||||
return cTTSModelConfig{
|
||||
Size: unsafe.Sizeof(cTTSModelConfig{}),
|
||||
MagpieModel: magpieModel,
|
||||
CodecModel: codecModel,
|
||||
TokenizerModelDir: tokenizerDir,
|
||||
TextNormalizerModelDir: tnDir,
|
||||
}
|
||||
}
|
||||
|
||||
// ttsRuntimeBackend maps the backend's gpu option onto the TTS runtime's
|
||||
// backend preference.
|
||||
//
|
||||
// nemo_speech_tts_runtime_config has no device index at all, only a three-way
|
||||
// AUTO/CPU/CUDA preference, so a gpu option naming a particular device cannot be
|
||||
// honoured and AUTO is the honest answer for it. A negative gpu is different: it
|
||||
// is the option's documented "CPU" across this whole backend (asr.h: "-1 = CPU")
|
||||
// and it is also the default, so it has to pin the preference rather than leave
|
||||
// the runtime free to pick CUDA.
|
||||
func ttsRuntimeBackend(gpu int32) int32 {
|
||||
if gpu < 0 {
|
||||
return ttsBackendCPU
|
||||
}
|
||||
return ttsBackendAuto
|
||||
}
|
||||
|
||||
// loadTTS creates the MagpieTTS synthesizer.
|
||||
//
|
||||
// It runs after discoverTTSAssets, so codecModel and tokenizerDir are already
|
||||
// resolved and non-empty; tnDir stays optional.
|
||||
//
|
||||
// This must not take engineMu: Load is its only caller and already holds it.
|
||||
func (n *NemoSpeech) loadTTS(modelFile string) error {
|
||||
// nemo_speech_tts_create deep-copies every const char* into a std::string
|
||||
// (src/tts/c_api.cpp, via str_or_empty) and keeps no pointer afterwards, so
|
||||
// pinning across the create call is both necessary and sufficient.
|
||||
var pinner runtime.Pinner
|
||||
defer pinner.Unpin()
|
||||
|
||||
magpieP, freeMagpie := cstr(modelFile)
|
||||
defer freeMagpie()
|
||||
codecP, freeCodec := cstr(n.opts.codecModel)
|
||||
defer freeCodec()
|
||||
tokenizerP, freeTokenizer := cstr(n.opts.tokenizerDir)
|
||||
defer freeTokenizer()
|
||||
tnP, freeTN := cstr(n.opts.tnDir)
|
||||
defer freeTN()
|
||||
|
||||
model := ttsModelConfig(magpieP, codecP, tokenizerP, tnP)
|
||||
|
||||
rt := TTSRuntimeConfigDefault()
|
||||
backend := ttsRuntimeBackend(n.opts.gpu)
|
||||
rt.LTBackend = backend
|
||||
rt.SamplingBackend = backend
|
||||
// The codec is a separate graph with its own placement, so a CPU-only
|
||||
// request has to say so here too or it would still try to run on the GPU.
|
||||
rt.CodecCPU = backend == ttsBackendCPU
|
||||
|
||||
langP, freeLang := cstr(n.opts.languageCode)
|
||||
defer freeLang()
|
||||
|
||||
cfg := cTTSSynthesizerConfig{
|
||||
Size: unsafe.Sizeof(cTTSSynthesizerConfig{}),
|
||||
Model: pinPtr(&pinner, &model),
|
||||
Runtime: pinPtr(&pinner, &rt),
|
||||
DefaultLanguageCode: langP,
|
||||
}
|
||||
|
||||
xlog.Info("nemo-speech-cpp: creating synthesizer",
|
||||
"gpu", n.opts.gpu,
|
||||
"codec", n.opts.codecModel,
|
||||
"tokenizer", n.opts.tokenizerDir,
|
||||
"text_normalizer", n.opts.tnDir != "")
|
||||
|
||||
// Compiled before the handle exists so that a full callback table fails the
|
||||
// load, where the operator can see it, rather than the first synthesis.
|
||||
ttsPCMCallback()
|
||||
|
||||
// #nosec G103 -- cfg is a local POD struct borrowed for this call only. Model
|
||||
// and Runtime are pinPtr addresses held by the pinner unpinned on return, the
|
||||
// paths they carry are cstr allocations freed by the defers above, and
|
||||
// nemo_speech_tts_create deep-copies every string it reads.
|
||||
if st := TTSCreate(unsafe.Pointer(&cfg), &n.synth); st != 0 {
|
||||
return statusErrorf(st, "nemo-speech-cpp: tts create: %s", TTSLastError())
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// validateTTSRequest rejects what the runtime would reject, before anything
|
||||
// crosses the ABI, and names the fields this backend drops.
|
||||
//
|
||||
// Empty text is checked here rather than left to the C side for the error code:
|
||||
// src/tts/synthesizer.cpp throws "text is required", which arrives as a status
|
||||
// this layer would otherwise report as Internal, and an empty prompt is a client
|
||||
// mistake, not a backend failure.
|
||||
//
|
||||
// instructions is logged rather than rejected, for the reason the diarization
|
||||
// path logs its own dropped fields: a caller that asked for an expressive style
|
||||
// still wants the audio it can have, and a request naming something this backend
|
||||
// silently ignores should say so where an operator can find it. There is nothing
|
||||
// to map it onto, because MagpieTTS conditions on a speaker, not on a prose
|
||||
// style description: nemo_speech_tts_synthesis_options has speaker and
|
||||
// voice_name and no free-text field at all.
|
||||
func validateTTSRequest(req *pb.TTSRequest) error {
|
||||
if req.GetText() == "" {
|
||||
return status.Error(codes.InvalidArgument, "nemo-speech-cpp: TTSRequest.text is required")
|
||||
}
|
||||
if req.GetInstructions() != "" {
|
||||
xlog.Warn("nemo-speech-cpp: ignoring TTSRequest.instructions, this model has no equivalent")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// outputSampleRate reads the rate the synthesizer emits at.
|
||||
//
|
||||
// A non-positive rate is refused rather than passed on. nemo_speech_tts_sample_rate
|
||||
// answers 0 for a null handle, and a WAV header carrying 0 is not a slightly
|
||||
// wrong file, it is one no player can decode and one whose duration is
|
||||
// undefined.
|
||||
func outputSampleRate(s synthesizer) (uint32, error) {
|
||||
rate := s.sampleRate()
|
||||
if rate <= 0 {
|
||||
return 0, status.Error(codes.Internal,
|
||||
"nemo-speech-cpp: the synthesizer reported no sample rate")
|
||||
}
|
||||
return uint32(rate), nil
|
||||
}
|
||||
|
||||
// wavFile frames PCM as a complete WAV: a header with real sizes, then the
|
||||
// samples.
|
||||
//
|
||||
// pcm is little-endian signed 16-bit mono, which is what the runtime's callback
|
||||
// delivers, and is exactly what pkg/audio's header describes, so nothing is
|
||||
// converted on the way through.
|
||||
func wavFile(pcm []byte, sampleRate uint32) ([]byte, error) {
|
||||
if int64(len(pcm)) > maxWAVDataBytes {
|
||||
return nil, status.Errorf(codes.Internal,
|
||||
"nemo-speech-cpp: synthesis produced %d bytes, more than a WAV header can describe", len(pcm))
|
||||
}
|
||||
|
||||
var buf bytes.Buffer
|
||||
// #nosec G115 -- len(pcm) is checked against maxWAVDataBytes (MaxUint32 minus
|
||||
// the header) immediately above, so the narrowing to uint32 cannot wrap.
|
||||
h := laudio.NewWAVHeaderWithRate(uint32(len(pcm)), sampleRate)
|
||||
if err := h.Write(&buf); err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "nemo-speech-cpp: write WAV header: %v", err)
|
||||
}
|
||||
buf.Write(pcm)
|
||||
return buf.Bytes(), nil
|
||||
}
|
||||
|
||||
// streamingWAVHeader is the first chunk of a streamed synthesis: the same header
|
||||
// with both sizes left unknown, since the total length is not known until the
|
||||
// synthesis ends.
|
||||
//
|
||||
// NewWAVHeaderWithRate derives ChunkSize from the payload length, so the RIFF
|
||||
// size has to be overwritten as well: 36 + 0xFFFFFFFF wraps to 35, which is a
|
||||
// smaller number than the header itself.
|
||||
func streamingWAVHeader(sampleRate uint32) []byte {
|
||||
h := laudio.NewWAVHeaderWithRate(wavStreamingSize, sampleRate)
|
||||
h.ChunkSize = wavStreamingSize
|
||||
|
||||
var buf bytes.Buffer
|
||||
// Write only fails on the writer, and bytes.Buffer does not fail.
|
||||
_ = h.Write(&buf)
|
||||
return buf.Bytes()
|
||||
}
|
||||
|
||||
// synthesizeWAV runs one synthesis and writes the whole result to dst.
|
||||
func synthesizeWAV(s synthesizer, req *pb.TTSRequest, defaultLanguage string) error {
|
||||
if err := validateTTSRequest(req); err != nil {
|
||||
return err
|
||||
}
|
||||
if req.GetDst() == "" {
|
||||
return status.Error(codes.InvalidArgument,
|
||||
"nemo-speech-cpp: TTSRequest.dst (output path) is required")
|
||||
}
|
||||
|
||||
// Read before the synthesis rather than after: it is what the header is
|
||||
// built from, and failing on a bad handle here costs nothing, where failing
|
||||
// after costs the whole synthesis.
|
||||
rate, err := outputSampleRate(s)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
var pcm []byte
|
||||
err = s.synthesize(req, defaultLanguage, func(chunk []byte) bool {
|
||||
pcm = append(pcm, chunk...)
|
||||
return true
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
// A synthesis that returned OK having emitted nothing is a runtime bug, but
|
||||
// the file it would produce is a valid empty WAV, which reaches the user as
|
||||
// silence with no error anywhere.
|
||||
if len(pcm) == 0 {
|
||||
return status.Error(codes.Internal, "nemo-speech-cpp: synthesis produced no audio")
|
||||
}
|
||||
|
||||
out, err := wavFile(pcm, rate)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := os.WriteFile(req.GetDst(), out, 0o600); err != nil {
|
||||
return status.Errorf(codes.Internal, "nemo-speech-cpp: write %q: %v", req.GetDst(), err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// streamWAV runs one synthesis and emits a WAV header followed by each PCM
|
||||
// chunk as the runtime produces it.
|
||||
//
|
||||
// The header is the backend's job, not the caller's: pkg/grpc/server.go only
|
||||
// ever sets Reply.Audio on this path and never Reply.Message, and
|
||||
// core/backend/tts.go's own header branch is keyed on Message, so a backend that
|
||||
// emitted bare PCM would stream something no client could decode. sherpa-onnx
|
||||
// and magpie-tts-cpp both do the same.
|
||||
//
|
||||
// out is not closed here. TTSStream owns it, and closing it in one of two places
|
||||
// depending on how far the request got is how a stream ends up half-closed.
|
||||
func streamWAV(s synthesizer, req *pb.TTSRequest, defaultLanguage string, out chan<- []byte) error {
|
||||
if err := validateTTSRequest(req); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
rate, err := outputSampleRate(s)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
out <- streamingWAVHeader(rate)
|
||||
|
||||
return s.synthesize(req, defaultLanguage, func(chunk []byte) bool {
|
||||
out <- chunk
|
||||
return true
|
||||
})
|
||||
}
|
||||
|
||||
// TTS synthesizes req.Text and writes a WAV to req.Dst.
|
||||
//
|
||||
// The whole body runs inside withEngine, so the family check and the C calls
|
||||
// that trust the handle happen under a single acquisition of engineMu. See the
|
||||
// handoff notes at the bottom of nemospeech.go: Free runs without the backend
|
||||
// lock, so anything that checks the family and then releases the lock before
|
||||
// calling C can have the handle destroyed underneath it.
|
||||
func (n *NemoSpeech) TTS(req *pb.TTSRequest) error {
|
||||
return n.withEngine(familyTTS, func() error {
|
||||
return synthesizeWAV(&cSynthesizer{handle: n.synth}, req, n.opts.languageCode)
|
||||
})
|
||||
}
|
||||
|
||||
// TTSStream synthesizes req.Text and emits the audio on results as it is
|
||||
// produced.
|
||||
//
|
||||
// results is closed on EVERY path, including the family rejection and a
|
||||
// validation failure, and the close is deferred outside withEngine so that a
|
||||
// rejected family still closes it. pkg/grpc/server.go drains this channel from a
|
||||
// goroutine and then blocks on that goroutine finishing, so a channel left open
|
||||
// does not fail the request, it hangs the RPC and, because the backend lock is
|
||||
// still held, every request behind it.
|
||||
//
|
||||
// Holding engineMu for the whole stream is deliberate and is the consequence
|
||||
// documented on the locking protocol: an unload waits for the stream to end
|
||||
// rather than destroying the synthesizer underneath it. There is no unbounded
|
||||
// wait here, because unlike the ASR streams this one is driven by the runtime
|
||||
// and ends when the text does, not when a client decides to stop sending.
|
||||
func (n *NemoSpeech) TTSStream(req *pb.TTSRequest, results chan []byte) error {
|
||||
defer close(results)
|
||||
|
||||
return n.withEngine(familyTTS, func() error {
|
||||
return streamWAV(&cSynthesizer{handle: n.synth}, req, n.opts.languageCode, results)
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,727 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"go/ast"
|
||||
"go/parser"
|
||||
"go/token"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"unsafe"
|
||||
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
|
||||
laudio "github.com/mudler/LocalAI/pkg/audio"
|
||||
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
|
||||
)
|
||||
|
||||
// puregoCallbackTableSize is the hard ceiling purego compiles callbacks into:
|
||||
// maxCB in purego/syscall_sysv.go, which panics rather than growing once it is
|
||||
// full and never releases an entry. Read off the module source for v0.10.0
|
||||
// rather than assumed, because the whole point of the specs below is that
|
||||
// exceeding it kills the process.
|
||||
const puregoCallbackTableSize = 2000
|
||||
|
||||
// fakeSynthesizer scripts what the TTS C API emits for one synthesis.
|
||||
//
|
||||
// There is no MagpieTTS GGUF in the tree, so this is the only way the logic on
|
||||
// top of the ABI (validation, WAV framing, chunk ordering, channel closure)
|
||||
// gets tested at all. It fakes the C contract, not the model: chunks is
|
||||
// whatever nemo_speech_tts_synthesize_text would have handed the callback.
|
||||
type fakeSynthesizer struct {
|
||||
rate int32
|
||||
chunks [][]byte
|
||||
err error
|
||||
|
||||
calls int
|
||||
gotReq *pb.TTSRequest
|
||||
gotLang string
|
||||
cancelled bool
|
||||
}
|
||||
|
||||
func (f *fakeSynthesizer) sampleRate() int32 { return f.rate }
|
||||
|
||||
func (f *fakeSynthesizer) synthesize(req *pb.TTSRequest, defaultLanguage string, sink ttsSink) error {
|
||||
f.calls++
|
||||
f.gotReq = req
|
||||
f.gotLang = defaultLanguage
|
||||
for _, c := range f.chunks {
|
||||
if !sink(c) {
|
||||
f.cancelled = true
|
||||
break
|
||||
}
|
||||
}
|
||||
return f.err
|
||||
}
|
||||
|
||||
var _ = Describe("resolveSpeaker", func() {
|
||||
It("passes a numeric voice through as a speaker index", func() {
|
||||
idx, name := resolveSpeaker("3")
|
||||
Expect(idx).To(Equal(int32(3)))
|
||||
Expect(name).To(BeEmpty())
|
||||
})
|
||||
|
||||
It("passes a named voice through as a name with no index", func() {
|
||||
// voice_name is ignored whenever speaker >= 0 (tts.h, and
|
||||
// synthesizer.cpp only calls resolve_speaker for a negative speaker), so
|
||||
// a named voice must leave the index negative or the name is dropped.
|
||||
idx, name := resolveSpeaker("Aria")
|
||||
Expect(idx).To(Equal(int32(-1)))
|
||||
Expect(name).To(Equal("Aria"))
|
||||
})
|
||||
|
||||
It("leaves both unset for an empty voice so the synthesizer default wins", func() {
|
||||
idx, name := resolveSpeaker("")
|
||||
Expect(idx).To(Equal(int32(-1)))
|
||||
Expect(name).To(BeEmpty())
|
||||
})
|
||||
|
||||
// A negative number is the C API's sentinel for "use the default", not a
|
||||
// speaker. Passing it through as an index would turn a request naming an
|
||||
// invalid voice into one that quietly synthesizes in the default voice.
|
||||
// Handed on as a name instead, resolve_speaker rejects it.
|
||||
It("does not let a negative number become a speaker index", func() {
|
||||
idx, name := resolveSpeaker("-1")
|
||||
Expect(idx).To(Equal(int32(-1)))
|
||||
Expect(name).To(Equal("-1"))
|
||||
})
|
||||
|
||||
It("treats a non-numeric voice that merely starts with digits as a name", func() {
|
||||
idx, name := resolveSpeaker("3-alpha")
|
||||
Expect(idx).To(Equal(int32(-1)))
|
||||
Expect(name).To(Equal("3-alpha"))
|
||||
})
|
||||
|
||||
It("keeps speaker 0 addressable", func() {
|
||||
// 0 is a real speaker index, and the only sentinel here is < 0.
|
||||
idx, name := resolveSpeaker("0")
|
||||
Expect(idx).To(BeZero())
|
||||
Expect(name).To(BeEmpty())
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("applySynthesisParams", func() {
|
||||
// The struct the runtime hands out: speaker/seed/steps/top_k all -1, the
|
||||
// overrides off. Written literally rather than taken from
|
||||
// TTSSynthesisOptionsDefault so the specs run without the shared libraries.
|
||||
defaults := func() cTTSSynthesisOptions {
|
||||
return cTTSSynthesisOptions{
|
||||
Size: unsafe.Sizeof(cTTSSynthesisOptions{}),
|
||||
Speaker: -1,
|
||||
Seed: -1,
|
||||
Steps: -1,
|
||||
TopK: -1,
|
||||
}
|
||||
}
|
||||
|
||||
It("leaves every default alone for an absent params map", func() {
|
||||
o := defaults()
|
||||
applySynthesisParams(&o, nil)
|
||||
Expect(o).To(Equal(defaults()))
|
||||
})
|
||||
|
||||
It("leaves every default alone for an empty params map", func() {
|
||||
o := defaults()
|
||||
applySynthesisParams(&o, map[string]string{})
|
||||
Expect(o).To(Equal(defaults()))
|
||||
})
|
||||
|
||||
It("maps the five knobs the C options struct actually has", func() {
|
||||
o := defaults()
|
||||
applySynthesisParams(&o, map[string]string{
|
||||
"seed": "42",
|
||||
"steps": "12",
|
||||
"top_k": "80",
|
||||
"temperature": "0.7",
|
||||
"cfg_scale": "1.5",
|
||||
})
|
||||
Expect(o.Seed).To(Equal(int32(42)))
|
||||
Expect(o.Steps).To(Equal(int32(12)))
|
||||
Expect(o.TopK).To(Equal(int32(80)))
|
||||
Expect(o.Temperature).To(BeNumerically("~", 0.7, 1e-6))
|
||||
Expect(o.CFGScale).To(BeNumerically("~", 1.5, 1e-6))
|
||||
})
|
||||
|
||||
// magpietts/runtime.cpp reads options.temperature only when
|
||||
// override_temperature is true and otherwise falls back to the
|
||||
// synthesizer's config, so a temperature written without its flag is
|
||||
// silently discarded and the request looks like it was honoured.
|
||||
It("sets the override flag with the temperature", func() {
|
||||
o := defaults()
|
||||
applySynthesisParams(&o, map[string]string{"temperature": "0.4"})
|
||||
Expect(o.OverrideTemperature).To(BeTrue())
|
||||
Expect(o.OverrideCFGScale).To(BeFalse(), "cfg_scale was not asked for")
|
||||
})
|
||||
|
||||
It("sets the override flag with the cfg scale", func() {
|
||||
o := defaults()
|
||||
applySynthesisParams(&o, map[string]string{"cfg_scale": "2"})
|
||||
Expect(o.OverrideCFGScale).To(BeTrue())
|
||||
Expect(o.OverrideTemperature).To(BeFalse(), "temperature was not asked for")
|
||||
})
|
||||
|
||||
It("keeps the defaults when a value cannot be parsed", func() {
|
||||
o := defaults()
|
||||
applySynthesisParams(&o, map[string]string{
|
||||
"seed": "many",
|
||||
"steps": "",
|
||||
"top_k": "8.5",
|
||||
"temperature": "warm",
|
||||
"cfg_scale": "-",
|
||||
})
|
||||
Expect(o).To(Equal(defaults()))
|
||||
})
|
||||
|
||||
// The runtime takes a request's seed only when it is >= 0 and its steps and
|
||||
// top_k only when they are > 0. Writing a parsed 0 or a negative would not
|
||||
// be ignored downstream, it would erase the sentinel that means "use the
|
||||
// synthesizer's value".
|
||||
It("refuses values that would erase a sentinel", func() {
|
||||
o := defaults()
|
||||
applySynthesisParams(&o, map[string]string{
|
||||
"seed": "-5",
|
||||
"steps": "0",
|
||||
"top_k": "0",
|
||||
})
|
||||
Expect(o.Seed).To(Equal(int32(-1)))
|
||||
Expect(o.Steps).To(Equal(int32(-1)))
|
||||
Expect(o.TopK).To(Equal(int32(-1)))
|
||||
})
|
||||
|
||||
It("keeps seed 0, which is a real seed", func() {
|
||||
o := defaults()
|
||||
applySynthesisParams(&o, map[string]string{"seed": "0"})
|
||||
Expect(o.Seed).To(BeZero())
|
||||
})
|
||||
|
||||
// TTSRequest carries fields with no equivalent in
|
||||
// nemo_speech_tts_synthesis_options. They must not be smuggled in through a
|
||||
// param name that happens to match.
|
||||
It("ignores params the C options struct has no field for", func() {
|
||||
o := defaults()
|
||||
applySynthesisParams(&o, map[string]string{
|
||||
"top_p": "0.9",
|
||||
"repetition_penalty": "1.1",
|
||||
"speed": "1.2",
|
||||
"instructions": "cheerful",
|
||||
})
|
||||
Expect(o).To(Equal(defaults()))
|
||||
})
|
||||
})
|
||||
|
||||
// The PCM callback is the one resource in this backend with a hard, silent,
|
||||
// process-wide ceiling: purego compiles each into a fixed table of 2000 entries
|
||||
// and never releases one, so a callback built per request takes the whole
|
||||
// backend process down with a panic after 2000 syntheses. Nothing about a
|
||||
// handful of manual calls shows that.
|
||||
var _ = Describe("ttsPCMCallback", func() {
|
||||
It("compiles a usable callback", func() {
|
||||
Expect(ttsPCMCallback()).ToNot(BeZero())
|
||||
})
|
||||
|
||||
It("compiles exactly one callback however many times it is asked", func() {
|
||||
first := ttsPCMCallback()
|
||||
|
||||
// One more than the table holds: a callback compiled per call panics
|
||||
// with "purego: the maximum number of callbacks has been reached"
|
||||
// before this loop ends, which is precisely the production failure.
|
||||
for i := 0; i <= puregoCallbackTableSize; i++ {
|
||||
Expect(ttsPCMCallback()).To(Equal(first),
|
||||
"call %d returned a different callback, so a new one was compiled", i)
|
||||
}
|
||||
})
|
||||
|
||||
// A source-level assertion, deliberately, because the failure it guards
|
||||
// against is invisible from inside the process: the way a per-request
|
||||
// callback gets reintroduced is by someone calling purego.NewCallback at the
|
||||
// synthesis site instead of going through ttsPCMCallback, and no in-process
|
||||
// spec can reach that call without a MagpieTTS GGUF to synthesize with.
|
||||
// Funnelling every compile through one accessor is what the whole design
|
||||
// rests on, so the single call site is the invariant worth pinning.
|
||||
It("compiles callbacks from exactly one place in the TTS path", func() {
|
||||
fset := token.NewFileSet()
|
||||
file, err := parser.ParseFile(fset, "tts.go", nil, 0)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
|
||||
// Counted over the syntax tree rather than by grepping the text: the
|
||||
// doc comment on ttsPCMCallback names purego.NewCallback too, and a
|
||||
// spec that cannot tell an explanation from a call would be pinning the
|
||||
// prose.
|
||||
var sites []string
|
||||
ast.Inspect(file, func(n ast.Node) bool {
|
||||
call, ok := n.(*ast.CallExpr)
|
||||
if !ok {
|
||||
return true
|
||||
}
|
||||
sel, ok := call.Fun.(*ast.SelectorExpr)
|
||||
if !ok || sel.Sel.Name != "NewCallback" {
|
||||
return true
|
||||
}
|
||||
if pkg, ok := sel.X.(*ast.Ident); ok && pkg.Name == "purego" {
|
||||
sites = append(sites, fset.Position(call.Pos()).String())
|
||||
}
|
||||
return true
|
||||
})
|
||||
Expect(sites).To(HaveLen(1),
|
||||
"every callback must be compiled through ttsPCMCallback, which memoises it")
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("the PCM sink table", func() {
|
||||
It("routes a chunk to the sink registered for that id", func() {
|
||||
var got []byte
|
||||
id, release := ttsSinks.register(func(pcm []byte) bool {
|
||||
got = pcm
|
||||
return true
|
||||
})
|
||||
defer release()
|
||||
|
||||
src := []byte{1, 2, 3, 4}
|
||||
Expect(ttsDeliverPCM(unsafe.Pointer(&src[0]), uint64(len(src)), id)).To(BeTrue())
|
||||
Expect(got).To(Equal([]byte{1, 2, 3, 4}))
|
||||
})
|
||||
|
||||
// The pointer addresses a std::string the runtime reuses for the next
|
||||
// chunk, so a slice over it would be rewritten under the consumer.
|
||||
It("copies the chunk out of the runtime's buffer", func() {
|
||||
var got []byte
|
||||
id, release := ttsSinks.register(func(pcm []byte) bool {
|
||||
got = pcm
|
||||
return true
|
||||
})
|
||||
defer release()
|
||||
|
||||
src := []byte{9, 8, 7}
|
||||
Expect(ttsDeliverPCM(unsafe.Pointer(&src[0]), uint64(len(src)), id)).To(BeTrue())
|
||||
src[0], src[1], src[2] = 0, 0, 0
|
||||
Expect(got).To(Equal([]byte{9, 8, 7}))
|
||||
})
|
||||
|
||||
It("gives each registration its own id", func() {
|
||||
idA, releaseA := ttsSinks.register(func([]byte) bool { return true })
|
||||
defer releaseA()
|
||||
idB, releaseB := ttsSinks.register(func([]byte) bool { return true })
|
||||
defer releaseB()
|
||||
|
||||
Expect(idA).ToNot(Equal(idB))
|
||||
Expect(idA).ToNot(BeZero(), "id 0 is what a zeroed user_data would carry")
|
||||
Expect(idB).ToNot(BeZero())
|
||||
})
|
||||
|
||||
// Two models synthesizing at once share one callback, and engineMu is
|
||||
// per-model, so nothing serialises them against each other.
|
||||
It("keeps concurrent sinks apart", func() {
|
||||
var mu sync.Mutex
|
||||
got := map[uintptr][]byte{}
|
||||
|
||||
var wg sync.WaitGroup
|
||||
for i := range 16 {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer GinkgoRecover()
|
||||
defer wg.Done()
|
||||
|
||||
src := []byte{byte(i)}
|
||||
var mine []byte
|
||||
id, release := ttsSinks.register(func(pcm []byte) bool {
|
||||
mine = pcm
|
||||
return true
|
||||
})
|
||||
defer release()
|
||||
|
||||
Expect(ttsDeliverPCM(unsafe.Pointer(&src[0]), 1, id)).To(BeTrue())
|
||||
mu.Lock()
|
||||
defer mu.Unlock()
|
||||
got[id] = mine
|
||||
}()
|
||||
}
|
||||
wg.Wait()
|
||||
|
||||
Expect(got).To(HaveLen(16))
|
||||
for id, pcm := range got {
|
||||
Expect(pcm).To(HaveLen(1), "sink %d received the wrong chunk", id)
|
||||
}
|
||||
})
|
||||
|
||||
// A released id means the request has returned. Answering true would leave
|
||||
// the runtime synthesizing into nothing while the RPC that owns the lock
|
||||
// waits for it.
|
||||
It("cancels the synthesis when the sink is gone", func() {
|
||||
id, release := ttsSinks.register(func([]byte) bool { return true })
|
||||
release()
|
||||
|
||||
src := []byte{1}
|
||||
Expect(ttsDeliverPCM(unsafe.Pointer(&src[0]), 1, id)).To(BeFalse())
|
||||
})
|
||||
|
||||
It("cancels for a user_data that was never registered", func() {
|
||||
src := []byte{1}
|
||||
Expect(ttsDeliverPCM(unsafe.Pointer(&src[0]), 1, 0)).To(BeFalse())
|
||||
})
|
||||
|
||||
It("accepts an empty chunk without touching the pointer", func() {
|
||||
id, release := ttsSinks.register(func([]byte) bool {
|
||||
Fail("an empty chunk must not reach the sink")
|
||||
return true
|
||||
})
|
||||
defer release()
|
||||
|
||||
Expect(ttsDeliverPCM(nil, 0, id)).To(BeTrue())
|
||||
})
|
||||
|
||||
It("passes the sink's cancellation back to the runtime", func() {
|
||||
id, release := ttsSinks.register(func([]byte) bool { return false })
|
||||
defer release()
|
||||
|
||||
src := []byte{1}
|
||||
Expect(ttsDeliverPCM(unsafe.Pointer(&src[0]), 1, id)).To(BeFalse())
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("ttsModelConfig", func() {
|
||||
// Three adjacent same-typed path fields: swapping two changes neither the
|
||||
// struct size nor any offset, so abi_test.go's layout assertions cannot see
|
||||
// it and the runtime would load the codec as the acoustic model.
|
||||
It("assigns each path to its own field", func() {
|
||||
cfg := ttsModelConfig(1, 2, 3, 4)
|
||||
Expect(cfg.MagpieModel).To(Equal(uintptr(1)))
|
||||
Expect(cfg.CodecModel).To(Equal(uintptr(2)))
|
||||
Expect(cfg.TokenizerModelDir).To(Equal(uintptr(3)))
|
||||
Expect(cfg.TextNormalizerModelDir).To(Equal(uintptr(4)))
|
||||
})
|
||||
|
||||
// A config sent with the wrong size has every field past it ignored by
|
||||
// HAS_FIELD, and the model loads with defaults instead of failing.
|
||||
It("declares the size the runtime validates against", func() {
|
||||
Expect(ttsModelConfig(1, 2, 3, 4).Size).To(Equal(unsafe.Sizeof(cTTSModelConfig{})))
|
||||
})
|
||||
|
||||
It("leaves an unset text normalizer null", func() {
|
||||
Expect(ttsModelConfig(1, 2, 3, 0).TextNormalizerModelDir).To(BeZero())
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("ttsRuntimeBackend", func() {
|
||||
// -1 is this backend's documented "CPU" everywhere (asr.h: "-1 = CPU") and
|
||||
// it is also the default, so it has to pin the preference rather than leave
|
||||
// the runtime free to pick CUDA.
|
||||
It("pins CPU for a negative gpu option", func() {
|
||||
Expect(ttsRuntimeBackend(-1)).To(Equal(ttsBackendCPU))
|
||||
})
|
||||
|
||||
// nemo_speech_tts_runtime_config has no device index at all, so a request
|
||||
// for a particular device cannot be honoured and AUTO is the honest answer.
|
||||
It("leaves the choice to the runtime when a device was named", func() {
|
||||
Expect(ttsRuntimeBackend(0)).To(Equal(ttsBackendAuto))
|
||||
Expect(ttsRuntimeBackend(3)).To(Equal(ttsBackendAuto))
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("WAV framing", func() {
|
||||
// 16-bit mono little-endian, the format the runtime's callback delivers.
|
||||
pcm := []byte{0x01, 0x00, 0xff, 0x7f, 0x00, 0x80}
|
||||
|
||||
It("writes a header the audio helpers can read back", func() {
|
||||
out, err := wavFile(pcm, 22050)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
|
||||
body, rate := laudio.ParseWAV(out)
|
||||
Expect(rate).To(Equal(22050))
|
||||
Expect(body).To(Equal(pcm))
|
||||
})
|
||||
|
||||
It("describes the payload it actually carries", func() {
|
||||
out, err := wavFile(pcm, 22050)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(out).To(HaveLen(laudio.WAVHeaderSize + len(pcm)))
|
||||
|
||||
Expect(string(out[0:4])).To(Equal("RIFF"))
|
||||
Expect(string(out[8:12])).To(Equal("WAVE"))
|
||||
Expect(binary.LittleEndian.Uint32(out[4:8])).To(Equal(uint32(36 + len(pcm))))
|
||||
Expect(binary.LittleEndian.Uint32(out[40:44])).To(Equal(uint32(len(pcm))))
|
||||
Expect(binary.LittleEndian.Uint16(out[22:24])).To(Equal(uint16(1)), "mono")
|
||||
Expect(binary.LittleEndian.Uint16(out[34:36])).To(Equal(uint16(16)), "16-bit")
|
||||
Expect(binary.LittleEndian.Uint32(out[24:28])).To(Equal(uint32(22050)))
|
||||
// byte rate = sample rate * block align, and a wrong one plays back at
|
||||
// the wrong speed in players that trust it.
|
||||
Expect(binary.LittleEndian.Uint32(out[28:32])).To(Equal(uint32(22050 * 2)))
|
||||
})
|
||||
|
||||
It("carries whatever rate the synthesizer reported", func() {
|
||||
out, err := wavFile(pcm, 44100)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
_, rate := laudio.ParseWAV(out)
|
||||
Expect(rate).To(Equal(44100))
|
||||
})
|
||||
|
||||
Describe("the streaming header", func() {
|
||||
It("is a complete header on its own", func() {
|
||||
h := streamingWAVHeader(22050)
|
||||
Expect(h).To(HaveLen(laudio.WAVHeaderSize))
|
||||
Expect(string(h[0:4])).To(Equal("RIFF"))
|
||||
Expect(string(h[8:12])).To(Equal("WAVE"))
|
||||
Expect(binary.LittleEndian.Uint32(h[24:28])).To(Equal(uint32(22050)))
|
||||
})
|
||||
|
||||
// NewWAVHeaderWithRate derives ChunkSize from the payload length, so
|
||||
// leaving it alone would write 36 + 0xFFFFFFFF, which wraps to 35: a
|
||||
// RIFF size smaller than the header itself.
|
||||
It("leaves both sizes unknown rather than wrapping", func() {
|
||||
h := streamingWAVHeader(22050)
|
||||
Expect(binary.LittleEndian.Uint32(h[4:8])).To(Equal(uint32(0xFFFFFFFF)))
|
||||
Expect(binary.LittleEndian.Uint32(h[40:44])).To(Equal(uint32(0xFFFFFFFF)))
|
||||
})
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("synthesizeWAV", func() {
|
||||
var dst string
|
||||
|
||||
BeforeEach(func() {
|
||||
dst = filepath.Join(GinkgoT().TempDir(), "out.wav")
|
||||
})
|
||||
|
||||
It("writes one WAV holding every chunk in order", func() {
|
||||
s := &fakeSynthesizer{rate: 22050, chunks: [][]byte{{1, 0}, {2, 0}, {3, 0}}}
|
||||
Expect(synthesizeWAV(s, &pb.TTSRequest{Text: "hello", Dst: dst}, "")).To(Succeed())
|
||||
|
||||
out, err := os.ReadFile(dst)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
body, rate := laudio.ParseWAV(out)
|
||||
Expect(rate).To(Equal(22050))
|
||||
Expect(body).To(Equal([]byte{1, 0, 2, 0, 3, 0}))
|
||||
})
|
||||
|
||||
It("hands the request and the model default language to the runtime", func() {
|
||||
s := &fakeSynthesizer{rate: 22050, chunks: [][]byte{{1, 0}}}
|
||||
req := &pb.TTSRequest{Text: "hello", Dst: dst, Voice: "Aria"}
|
||||
Expect(synthesizeWAV(s, req, "it-IT")).To(Succeed())
|
||||
Expect(s.gotReq).To(Equal(req))
|
||||
Expect(s.gotLang).To(Equal("it-IT"))
|
||||
})
|
||||
|
||||
It("rejects an empty text before it reaches the runtime", func() {
|
||||
s := &fakeSynthesizer{rate: 22050}
|
||||
err := synthesizeWAV(s, &pb.TTSRequest{Dst: dst}, "")
|
||||
Expect(status.Code(err)).To(Equal(codes.InvalidArgument))
|
||||
Expect(s.calls).To(BeZero())
|
||||
})
|
||||
|
||||
// instructions has no equivalent in nemo_speech_tts_synthesis_options, which
|
||||
// conditions on a speaker rather than a prose style. Dropping it must not
|
||||
// fail the request: the caller still wants the audio it can have.
|
||||
It("synthesizes anyway for a request carrying instructions it cannot honour", func() {
|
||||
s := &fakeSynthesizer{rate: 22050, chunks: [][]byte{{1, 0}}}
|
||||
instructions := "speak cheerfully"
|
||||
Expect(synthesizeWAV(s, &pb.TTSRequest{
|
||||
Text: "hello",
|
||||
Dst: dst,
|
||||
Instructions: &instructions,
|
||||
}, "")).To(Succeed())
|
||||
Expect(dst).To(BeAnExistingFile())
|
||||
})
|
||||
|
||||
It("rejects a request with no destination", func() {
|
||||
s := &fakeSynthesizer{rate: 22050}
|
||||
err := synthesizeWAV(s, &pb.TTSRequest{Text: "hello"}, "")
|
||||
Expect(status.Code(err)).To(Equal(codes.InvalidArgument))
|
||||
Expect(s.calls).To(BeZero())
|
||||
})
|
||||
|
||||
// A zero rate is what a null handle reports. The file it would produce is
|
||||
// undecodable, and the synthesis that produced it would be wasted.
|
||||
It("refuses to write a file at an unusable sample rate", func() {
|
||||
s := &fakeSynthesizer{rate: 0, chunks: [][]byte{{1, 0}}}
|
||||
err := synthesizeWAV(s, &pb.TTSRequest{Text: "hello", Dst: dst}, "")
|
||||
Expect(status.Code(err)).To(Equal(codes.Internal))
|
||||
Expect(s.calls).To(BeZero())
|
||||
Expect(dst).ToNot(BeAnExistingFile())
|
||||
})
|
||||
|
||||
It("propagates a synthesis failure and writes nothing", func() {
|
||||
boom := errors.New("boom")
|
||||
s := &fakeSynthesizer{rate: 22050, chunks: [][]byte{{1, 0}}, err: boom}
|
||||
Expect(synthesizeWAV(s, &pb.TTSRequest{Text: "hello", Dst: dst}, "")).To(MatchError(boom))
|
||||
Expect(dst).ToNot(BeAnExistingFile())
|
||||
})
|
||||
|
||||
// An empty WAV is a valid file, so this would otherwise reach the user as
|
||||
// silence with no error anywhere.
|
||||
It("fails rather than write a silent file when nothing was produced", func() {
|
||||
s := &fakeSynthesizer{rate: 22050}
|
||||
err := synthesizeWAV(s, &pb.TTSRequest{Text: "hello", Dst: dst}, "")
|
||||
Expect(status.Code(err)).To(Equal(codes.Internal))
|
||||
Expect(dst).ToNot(BeAnExistingFile())
|
||||
})
|
||||
|
||||
It("reports a destination it cannot write", func() {
|
||||
s := &fakeSynthesizer{rate: 22050, chunks: [][]byte{{1, 0}}}
|
||||
bad := filepath.Join(GinkgoT().TempDir(), "no-such-dir", "out.wav")
|
||||
err := synthesizeWAV(s, &pb.TTSRequest{Text: "hello", Dst: bad}, "")
|
||||
Expect(status.Code(err)).To(Equal(codes.Internal))
|
||||
Expect(err.Error()).To(ContainSubstring("out.wav"))
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("streamWAV", func() {
|
||||
// drain collects everything streamWAV emits. The channel is buffered
|
||||
// because streamWAV sends inline, so an unbuffered one would deadlock the
|
||||
// spec rather than fail it.
|
||||
drain := func(s synthesizer, req *pb.TTSRequest) ([][]byte, error) {
|
||||
out := make(chan []byte, 16)
|
||||
err := streamWAV(s, req, "", out)
|
||||
close(out)
|
||||
|
||||
var got [][]byte
|
||||
for c := range out {
|
||||
got = append(got, c)
|
||||
}
|
||||
return got, err
|
||||
}
|
||||
|
||||
It("emits the header first, then each chunk as it arrives", func() {
|
||||
s := &fakeSynthesizer{rate: 22050, chunks: [][]byte{{1, 0}, {2, 0}}}
|
||||
got, err := drain(s, &pb.TTSRequest{Text: "hello"})
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
|
||||
Expect(got).To(HaveLen(3))
|
||||
Expect(got[0]).To(Equal(streamingWAVHeader(22050)))
|
||||
Expect(got[1]).To(Equal([]byte{1, 0}))
|
||||
Expect(got[2]).To(Equal([]byte{2, 0}))
|
||||
})
|
||||
|
||||
// pkg/grpc/server.go only ever sets Reply.Audio, so core/backend's own
|
||||
// header branch (keyed on Reply.Message) never runs and a backend that
|
||||
// emitted bare PCM would stream something no client could decode.
|
||||
It("owns the header rather than leaving it to the caller", func() {
|
||||
s := &fakeSynthesizer{rate: 44100, chunks: [][]byte{{1, 0}}}
|
||||
got, err := drain(s, &pb.TTSRequest{Text: "hello"})
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(string(got[0][0:4])).To(Equal("RIFF"))
|
||||
Expect(binary.LittleEndian.Uint32(got[0][24:28])).To(Equal(uint32(44100)))
|
||||
})
|
||||
|
||||
It("rejects an empty text before emitting anything", func() {
|
||||
s := &fakeSynthesizer{rate: 22050, chunks: [][]byte{{1, 0}}}
|
||||
got, err := drain(s, &pb.TTSRequest{})
|
||||
Expect(status.Code(err)).To(Equal(codes.InvalidArgument))
|
||||
Expect(got).To(BeEmpty())
|
||||
Expect(s.calls).To(BeZero())
|
||||
})
|
||||
|
||||
It("emits no header at an unusable sample rate", func() {
|
||||
s := &fakeSynthesizer{rate: 0, chunks: [][]byte{{1, 0}}}
|
||||
got, err := drain(s, &pb.TTSRequest{Text: "hello"})
|
||||
Expect(status.Code(err)).To(Equal(codes.Internal))
|
||||
Expect(got).To(BeEmpty())
|
||||
})
|
||||
|
||||
It("propagates a synthesis failure after the chunks it did emit", func() {
|
||||
boom := errors.New("boom")
|
||||
s := &fakeSynthesizer{rate: 22050, chunks: [][]byte{{1, 0}}, err: boom}
|
||||
got, err := drain(s, &pb.TTSRequest{Text: "hello"})
|
||||
Expect(err).To(MatchError(boom))
|
||||
Expect(got).To(HaveLen(2))
|
||||
})
|
||||
|
||||
// streamWAV must not close the channel: TTSStream owns it, and closing in
|
||||
// one of two places depending on how far the request got is how a stream
|
||||
// ends up double-closed.
|
||||
It("leaves the channel open for its caller to close", func() {
|
||||
s := &fakeSynthesizer{rate: 22050, chunks: [][]byte{{1, 0}}}
|
||||
out := make(chan []byte, 4)
|
||||
Expect(streamWAV(s, &pb.TTSRequest{Text: "hello"}, "", out)).To(Succeed())
|
||||
Expect(func() { close(out) }).ToNot(Panic())
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("the TTS RPCs", func() {
|
||||
It("refuses TTS on a model loaded as another family", func() {
|
||||
n := &NemoSpeech{fam: familyASR}
|
||||
err := n.TTS(&pb.TTSRequest{Text: "hello", Dst: "/tmp/out.wav"})
|
||||
Expect(status.Code(err)).To(Equal(codes.Unimplemented))
|
||||
})
|
||||
|
||||
It("refuses TTS on an unloaded model", func() {
|
||||
n := &NemoSpeech{}
|
||||
Expect(status.Code(n.TTS(&pb.TTSRequest{Text: "hello", Dst: "/tmp/out.wav"}))).
|
||||
To(Equal(codes.Unimplemented))
|
||||
})
|
||||
|
||||
It("releases the engine lock after a refusal", func() {
|
||||
n := &NemoSpeech{fam: familyASR}
|
||||
Expect(n.TTS(&pb.TTSRequest{Text: "hello", Dst: "/tmp/out.wav"})).ToNot(Succeed())
|
||||
Expect(n.engineMu.TryLock()).To(BeTrue())
|
||||
n.engineMu.Unlock()
|
||||
})
|
||||
|
||||
// pkg/grpc/server.go drains this channel from a goroutine and then blocks
|
||||
// on that goroutine finishing, so a channel left open does not fail the
|
||||
// request, it hangs the RPC with the backend lock still held. Every exit
|
||||
// path has to close it.
|
||||
Describe("TTSStream channel closure", func() {
|
||||
// streamed runs TTSStream the way the server does and returns once the
|
||||
// channel has been closed, so a spec that hangs is a real hang.
|
||||
streamed := func(n *NemoSpeech, req *pb.TTSRequest) ([][]byte, error) {
|
||||
ch := make(chan []byte, 16)
|
||||
done := make(chan [][]byte, 1)
|
||||
go func() {
|
||||
defer GinkgoRecover()
|
||||
var got [][]byte
|
||||
for c := range ch {
|
||||
got = append(got, c)
|
||||
}
|
||||
done <- got
|
||||
}()
|
||||
|
||||
err := n.TTSStream(req, ch)
|
||||
var got [][]byte
|
||||
Eventually(done).Should(Receive(&got))
|
||||
return got, err
|
||||
}
|
||||
|
||||
It("closes the channel when the family does not match", func() {
|
||||
n := &NemoSpeech{fam: familyASR}
|
||||
got, err := streamed(n, &pb.TTSRequest{Text: "hello"})
|
||||
Expect(status.Code(err)).To(Equal(codes.Unimplemented))
|
||||
Expect(got).To(BeEmpty())
|
||||
})
|
||||
|
||||
It("closes the channel when the model was never loaded", func() {
|
||||
n := &NemoSpeech{}
|
||||
_, err := streamed(n, &pb.TTSRequest{Text: "hello"})
|
||||
Expect(status.Code(err)).To(Equal(codes.Unimplemented))
|
||||
})
|
||||
|
||||
// familyTTS with a zero handle: validation has to reject this before
|
||||
// anything reaches the C entry points, which are nil function values
|
||||
// until openLibraries has bound them.
|
||||
It("closes the channel when the request is rejected", func() {
|
||||
n := &NemoSpeech{fam: familyTTS}
|
||||
got, err := streamed(n, &pb.TTSRequest{})
|
||||
Expect(status.Code(err)).To(Equal(codes.InvalidArgument))
|
||||
Expect(got).To(BeEmpty())
|
||||
})
|
||||
|
||||
It("releases the engine lock afterwards", func() {
|
||||
n := &NemoSpeech{fam: familyTTS}
|
||||
_, err := streamed(n, &pb.TTSRequest{})
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(n.engineMu.TryLock()).To(BeTrue())
|
||||
n.engineMu.Unlock()
|
||||
})
|
||||
})
|
||||
|
||||
// The same guard on the offline path: a rejected request must not reach a
|
||||
// nil C function through a zero handle.
|
||||
It("rejects an invalid TTS request without touching the runtime", func() {
|
||||
n := &NemoSpeech{fam: familyTTS}
|
||||
Expect(status.Code(n.TTS(&pb.TTSRequest{Dst: "/tmp/out.wav"}))).To(Equal(codes.InvalidArgument))
|
||||
Expect(status.Code(n.TTS(&pb.TTSRequest{Text: "hello"}))).To(Equal(codes.InvalidArgument))
|
||||
})
|
||||
})
|
||||
@@ -1,6 +1,6 @@
|
||||
# parakeet-cpp backend Makefile.
|
||||
#
|
||||
# Upstream pin lives below as PARAKEET_VERSION?=1bfbebfaaf493866f49597cd3b7901959d395c60
|
||||
# Upstream pin lives below as PARAKEET_VERSION?=e75de9b6b9b688fd293aa22f7e27aa724ea286f8
|
||||
# (.github/bump_deps.sh) can find and update it - matches the
|
||||
# whisper.cpp / ds4 / vibevoice-cpp convention.
|
||||
#
|
||||
@@ -15,7 +15,7 @@
|
||||
# That's what the L0 smoke test uses. The default target below does the
|
||||
# proper clone-at-pin + cmake build so CI doesn't need a side-checkout.
|
||||
|
||||
PARAKEET_VERSION?=1bfbebfaaf493866f49597cd3b7901959d395c60
|
||||
PARAKEET_VERSION?=e75de9b6b9b688fd293aa22f7e27aa724ea286f8
|
||||
PARAKEET_REPO?=https://github.com/mudler/parakeet.cpp
|
||||
|
||||
GOCMD?=go
|
||||
@@ -49,6 +49,8 @@ else ifeq ($(BUILD_TYPE),hipblas)
|
||||
CMAKE_ARGS+=-DPARAKEET_GGML_HIP=ON
|
||||
else ifeq ($(BUILD_TYPE),vulkan)
|
||||
CMAKE_ARGS+=-DPARAKEET_GGML_VULKAN=ON
|
||||
else ifeq ($(BUILD_TYPE),metal)
|
||||
CMAKE_ARGS+=-DPARAKEET_GGML_METAL=ON
|
||||
endif
|
||||
|
||||
.PHONY: parakeet-cpp-grpc package build clean purge test all
|
||||
|
||||
@@ -8,7 +8,7 @@ JOBS?=$(shell nproc --ignore=1)
|
||||
|
||||
# stablediffusion.cpp (ggml)
|
||||
STABLEDIFFUSION_GGML_REPO?=https://github.com/leejet/stable-diffusion.cpp
|
||||
STABLEDIFFUSION_GGML_VERSION?=db99efdd6d2a43c7937fd55b3359206c680a75b0
|
||||
STABLEDIFFUSION_GGML_VERSION?=97d2990807fe6d558e395f8764198d7c7e7b411c
|
||||
|
||||
CMAKE_ARGS+=-DGGML_MAX_NAME=128
|
||||
|
||||
@@ -42,13 +42,9 @@ else ifeq ($(BUILD_TYPE),hipblas)
|
||||
CMAKE_ARGS+=-DSD_HIPBLAS=ON -DGGML_HIPBLAS=ON -DAMDGPU_TARGETS=$(AMDGPU_TARGETS)
|
||||
else ifeq ($(BUILD_TYPE),vulkan)
|
||||
CMAKE_ARGS+=-DSD_VULKAN=ON -DGGML_VULKAN=ON
|
||||
else ifeq ($(OS),Darwin)
|
||||
ifneq ($(BUILD_TYPE),metal)
|
||||
CMAKE_ARGS+=-DSD_METAL=OFF -DGGML_METAL=OFF
|
||||
else
|
||||
CMAKE_ARGS+=-DSD_METAL=ON -DGGML_METAL=ON
|
||||
CMAKE_ARGS+=-DGGML_METAL_EMBED_LIBRARY=ON
|
||||
endif
|
||||
else ifeq ($(BUILD_TYPE),metal)
|
||||
CMAKE_ARGS+=-DSD_METAL=ON -DGGML_METAL=ON
|
||||
CMAKE_ARGS+=-DGGML_METAL_EMBED_LIBRARY=ON
|
||||
endif
|
||||
|
||||
ifeq ($(BUILD_TYPE),sycl_f16)
|
||||
@@ -72,7 +68,6 @@ sources/stablediffusion-ggml.cpp:
|
||||
git checkout $(STABLEDIFFUSION_GGML_VERSION) && \
|
||||
git submodule update --init --recursive --depth 1 --single-branch
|
||||
|
||||
# Detect OS
|
||||
UNAME_S := $(shell uname -s)
|
||||
|
||||
# Only build CPU variants on Linux
|
||||
@@ -134,4 +129,4 @@ libgosd-custom: CMakeLists.txt cpp/gosd.cpp cpp/gosd.h
|
||||
(mv build-$(SO_TARGET)/libgosd.so ./$(SO_TARGET) 2>/dev/null || \
|
||||
mv build-$(SO_TARGET)/libgosd.dylib ./$(SO_TARGET) 2>/dev/null)
|
||||
|
||||
all: stablediffusion-ggml package
|
||||
all: stablediffusion-ggml package
|
||||
+138
-26
@@ -11,7 +11,30 @@ JOBS?=$(shell nproc --ignore=1 2>/dev/null || sysctl -n hw.ncpu 2>/dev/null || e
|
||||
|
||||
# vllm.cpp version
|
||||
VLLM_CPP_REPO?=https://github.com/mudler/vllm.cpp
|
||||
VLLM_CPP_VERSION?=9e1c9025ae61167a3335454d7cc0de6093c21845
|
||||
VLLM_CPP_VERSION?=438305e1577768ec0f75729456a4c8b9f425e2ee
|
||||
|
||||
# MLX GEMM provider (darwin/metal only; see the metal branch below for why).
|
||||
# Consumed as the prebuilt pip wheel: building MLX from source needs `xcrun
|
||||
# metal`, i.e. a full Xcode the macOS runners do not have, while the wheel ships
|
||||
# include/, lib/libmlx.dylib and the compiled mlx.metallib ready to link.
|
||||
#
|
||||
# DEFAULT ON, but ONLY because VLLM_CPP_VERSION above is pinned at or past
|
||||
# vllm.cpp 89c46aeb, which SHAPE-GATES the provider to prefill. The ordering is
|
||||
# load-bearing, not incidental:
|
||||
#
|
||||
# pin >= 89c46aeb, MLX on -> 99.1% of MLX-LM (gated: prefill only)
|
||||
# pin < 89c46aeb, MLX on -> ~51% (ungated: it also takes decode)
|
||||
#
|
||||
# MLX's steel GEMM wins prefill (537 ms TTFT against 602) and loses decode badly,
|
||||
# because the provider pays an mx::eval sync plus an output memcpy per call and
|
||||
# decode makes ~112 calls per TOKEN. Ungated it does both; gated it does only the
|
||||
# good half. So if this pin is ever moved BACKWARDS, this default must go with it.
|
||||
VLLM_CPP_MLX?=on
|
||||
MLX_VERSION?=0.29.4
|
||||
MLX_VENV?=$(abspath ./mlx-venv)
|
||||
# Resolved lazily (recursive `=`, not `:=`): the glob only matches once the venv
|
||||
# target has run, and the interpreter version in the path varies per runner.
|
||||
MLX_ROOT=$(shell echo $(MLX_VENV)/lib/python*/site-packages/mlx)
|
||||
|
||||
# The backend consumes only the stable C ABI (libvllm + include/vllm.h), so the
|
||||
# server, examples and tests of the engine are never built here.
|
||||
@@ -24,31 +47,57 @@ CMAKE_ARGS+=-DCMAKE_BUILD_TYPE=Release
|
||||
UNAME_M := $(shell uname -m)
|
||||
|
||||
ifeq ($(BUILD_TYPE),cublas)
|
||||
# Blackwell-family targets only: other CUDA arches are build-supported
|
||||
# upstream but have no runtime-proven fast path. amd64 gets the consumer
|
||||
# (120a) + GB10 (121a) fat binary; arm64 CUDA (l4t-style images, DGX
|
||||
# Spark) is GB10 only. Triton-AOT GDN cubins are vendored per-arch, no
|
||||
# Python needed to consume them.
|
||||
# Every CUDA architecture upstream builds that the platform can actually
|
||||
# host, split by where the silicon exists: Jetson (87 Orin, 110 Thor) is
|
||||
# arm64-only, desktop 120a is amd64-only, and 90a/100a appear on both
|
||||
# because of the SBSA parts (GH200, GB200).
|
||||
#
|
||||
# This deliberately matches vllm.cpp's own release archive rather than
|
||||
# narrowing to the boxes we benchmark on. A narrower list does not degrade
|
||||
# on an unlisted card, it dies at the first request with "no kernel image
|
||||
# is available for execution on the device", long after `backends install`
|
||||
# reported success -- so an arch we merely lack numbers for still belongs
|
||||
# in the binary.
|
||||
#
|
||||
# Triton-AOT stays ON for both. A fat build is supported on the BUILDER
|
||||
# path: it embeds every vendored cubin tree (sm_80/86/89/90a/100a/121a) and
|
||||
# selects by exact SM at runtime, so the arches with no tree (87, 103a,
|
||||
# 110, 120a) take the portable CUDA kernels and can never load a
|
||||
# neighbouring cubin. Only maintainer REGEN needs a single pinned arch.
|
||||
# See vllm.cpp cmake/TritonAOT.cmake `_triton_aot_arch_names`.
|
||||
#
|
||||
# CUDA builds REQUIRE the CUDA 13 toolchain: 12.x nvcc lacks compute_121a
|
||||
# (GB10) and its ptxas rejects the sm_120a NVFP4 MMA kernels ("Vector type
|
||||
# too large"), so no cuda-12 variant is shipped.
|
||||
ifeq ($(CUDA_MAJOR_VERSION),12)
|
||||
$(error vllm.cpp needs the CUDA 13 toolchain: CUDA 12.x cannot compile the Blackwell fp4 kernels)
|
||||
endif
|
||||
ifeq ($(UNAME_M),x86_64)
|
||||
# NO -DVLLM_CPP_TRITON on fat builds: the vendored Triton-AOT cubin
|
||||
# trees are per-arch and the engine refuses a multi-arch build unless
|
||||
# pinned to one tree (unsound for the other arch). The non-AOT GDN
|
||||
# path serves the fat binary; single-arch builds keep the cubins.
|
||||
#
|
||||
# CUDA builds REQUIRE the CUDA 13 toolchain: 12.x nvcc lacks
|
||||
# compute_121a (GB10) and its ptxas rejects the sm_120a NVFP4 MMA
|
||||
# kernels ("Vector type too large"), so no cuda-12 variant is shipped.
|
||||
ifeq ($(CUDA_MAJOR_VERSION),12)
|
||||
$(error vllm.cpp needs the CUDA 13 toolchain: CUDA 12.x cannot compile the Blackwell fp4 kernels)
|
||||
endif
|
||||
CMAKE_ARGS+=-DVLLM_CPP_CUDA=ON "-DVLLM_CPP_CUDA_ARCHITECTURES=120a;121a"
|
||||
CMAKE_ARGS+=-DVLLM_CPP_CUDA=ON "-DVLLM_CPP_CUDA_ARCHITECTURES=80;86;89;90a;100a;103a;120a;121a" -DVLLM_CPP_TRITON=ON
|
||||
else
|
||||
CMAKE_ARGS+=-DVLLM_CPP_CUDA=ON -DVLLM_CPP_CUDA_ARCHITECTURES=121a -DVLLM_CPP_TRITON=ON
|
||||
CMAKE_ARGS+=-DVLLM_CPP_CUDA=ON "-DVLLM_CPP_CUDA_ARCHITECTURES=87;90a;100a;110;121a" -DVLLM_CPP_TRITON=ON
|
||||
endif
|
||||
else ifeq ($(BUILD_TYPE),vulkan)
|
||||
CMAKE_ARGS+=-DVLLM_CPP_VULKAN=ON -DVLLM_CPP_CUDA=OFF
|
||||
else ifeq ($(BUILD_TYPE),metal)
|
||||
CMAKE_ARGS+=-DVLLM_CPP_METAL=ON
|
||||
# The optional MLX GEMM provider. vllm.cpp keeps it OFF by default because it
|
||||
# is a ~19 MB libmlx.dylib plus a ~105 MB mlx.metallib, and upstream's
|
||||
# position is that it must earn that cost by measurement. It does, on the
|
||||
# only hardware this build targets: measured on an Apple M4 against the
|
||||
# native MSL GEMM in the SAME binary (arms toggled by
|
||||
# VT_OP_PROVIDER_DISABLE=mlx), Qwen3-1.7B-bf16 p=512 g=128, it is 1.5x to
|
||||
# 2.2x aggregate throughput and 2x to 3x faster TTFT, at equal peak memory
|
||||
# and bit-identical output on every parity shape. See vllm.cpp
|
||||
# docs/BENCHMARKS.md "MLX GEMM provider A/B on Apple M4".
|
||||
#
|
||||
# MLX delegates the dense GEMM ONLY: kPagedAttention stays vllm.cpp's own
|
||||
# kernel, because MLX has no paged-KV primitive at all.
|
||||
#
|
||||
# Set VLLM_CPP_MLX=off for a Metal build without it (smaller image, slower).
|
||||
ifeq ($(VLLM_CPP_MLX),on)
|
||||
MLX_ENABLED=1
|
||||
endif
|
||||
else
|
||||
CMAKE_ARGS+=-DVLLM_CPP_CUDA=OFF
|
||||
endif
|
||||
@@ -56,39 +105,102 @@ endif
|
||||
UNAME_S := $(shell uname -s)
|
||||
ifeq ($(UNAME_S),Darwin)
|
||||
LIB=libvllm.dylib
|
||||
# Apple Clang diagnoses a pair of constant-folded array bounds in the Metal
|
||||
# build as a GNU extension. Disable that diagnostic for both Objective-C and
|
||||
# C++ because vllm.cpp appends target-local -Werror after these global flags.
|
||||
CMAKE_ARGS+=-DCMAKE_CXX_FLAGS=-Wno-gnu-folding-constant
|
||||
CMAKE_ARGS+=-DCMAKE_OBJC_FLAGS=-Wno-gnu-folding-constant
|
||||
CMAKE_ARGS+=-DCMAKE_OBJCXX_FLAGS=-Wno-gnu-folding-constant
|
||||
else
|
||||
LIB=libvllm.so
|
||||
endif
|
||||
|
||||
sources/vllm.cpp:
|
||||
# patches/ carries fixes the pinned engine SHA does not have yet. `git apply`
|
||||
# is deliberately unguarded: a patch that no longer applies must FAIL the clone
|
||||
# loudly, because the alternative is a pin that silently ships without a fix it
|
||||
# is documented to carry. Each patch header says which pin retires it.
|
||||
VLLM_CPP_PATCHES=$(wildcard patches/*.patch)
|
||||
|
||||
sources/vllm.cpp: $(VLLM_CPP_PATCHES)
|
||||
rm -rf sources/vllm.cpp
|
||||
mkdir -p sources/vllm.cpp
|
||||
cd sources/vllm.cpp && \
|
||||
git init && \
|
||||
git remote add origin $(VLLM_CPP_REPO) && \
|
||||
git fetch --depth 1 origin $(VLLM_CPP_VERSION) && \
|
||||
git checkout FETCH_HEAD
|
||||
git checkout FETCH_HEAD && \
|
||||
for p in $(VLLM_CPP_PATCHES); do \
|
||||
echo "==> applying $$p"; \
|
||||
git apply ../../$$p || exit 1; \
|
||||
done
|
||||
|
||||
$(LIB): sources/vllm.cpp
|
||||
ifeq ($(MLX_ENABLED),1)
|
||||
# A stamp FILE, not a phony target: a phony prerequisite is always "newer" than
|
||||
# $(LIB) and would re-link libvllm on every invocation. Keyed on the version so
|
||||
# a MLX_VERSION bump reinstalls instead of silently reusing the old wheel.
|
||||
MLX_STAMP=$(MLX_VENV)/.mlx-$(MLX_VERSION).stamp
|
||||
MLX_CMAKE_ARGS=-DVLLM_CPP_MLX=ON -DMLX_ROOT=$(MLX_ROOT)
|
||||
|
||||
$(MLX_STAMP):
|
||||
@if [ ! -x "$(MLX_VENV)/bin/pip" ]; then \
|
||||
python3 -m venv "$(MLX_VENV)" || { echo "vllm-cpp: python3 with venv is required to build the MLX provider; pass VLLM_CPP_MLX=off to build Metal without it" >&2; exit 1; }; \
|
||||
fi
|
||||
"$(MLX_VENV)"/bin/pip install --quiet --disable-pip-version-check "mlx==$(MLX_VERSION)"
|
||||
@# Resolved in the SHELL, not by $(MLX_ROOT): make expands a whole recipe
|
||||
@# before running its first line, so the glob would still be unmatched here.
|
||||
@# Every later use (the cmake args, package.sh) expands after this target has
|
||||
@# completed, where $(MLX_ROOT) does resolve.
|
||||
@root=$$(echo "$(MLX_VENV)"/lib/python*/site-packages/mlx); \
|
||||
test -f "$$root/lib/libmlx.dylib" -a -f "$$root/include/mlx/array.h" || \
|
||||
{ echo "vllm-cpp: mlx==$(MLX_VERSION) did not provide lib/libmlx.dylib + include/mlx/array.h under $$root" >&2; exit 1; }
|
||||
touch $@
|
||||
else
|
||||
MLX_STAMP=
|
||||
MLX_CMAKE_ARGS=
|
||||
endif
|
||||
|
||||
# govllmcpp.go mirrors vllm.h by hand, and the only guard against the two
|
||||
# drifting apart is the vllm_abi_version check inside registerLib - which fires
|
||||
# at runtime, on the user's machine, taking down every model load (issue
|
||||
# #11379). Compare the two here instead, so moving VLLM_CPP_VERSION past the
|
||||
# mirrors turns the build red while the header is still around to diff.
|
||||
abi-check: sources/vllm.cpp
|
||||
@engine=$$(sed -n 's/^#define VLLM_ABI_VERSION \([0-9][0-9]*\).*/\1/p' sources/vllm.cpp/include/vllm.h); \
|
||||
backend=$$(sed -n 's/^const abiVersion = \([0-9][0-9]*\).*/\1/p' govllmcpp.go); \
|
||||
if [ -z "$$engine" ] || [ -z "$$backend" ]; then \
|
||||
echo "vllm-cpp: cannot read the ABI version (engine='$$engine' backend='$$backend')" >&2; exit 1; \
|
||||
fi; \
|
||||
if [ "$$engine" != "$$backend" ]; then \
|
||||
echo "vllm-cpp: ABI mismatch: vllm.cpp $(VLLM_CPP_VERSION) is v$$engine, govllmcpp.go mirrors v$$backend." >&2; \
|
||||
echo " Update the struct mirrors and abiVersion in govllmcpp.go (and the offsets in vllmcpp_test.go) to v$$engine." >&2; \
|
||||
exit 1; \
|
||||
fi; \
|
||||
echo "vllm-cpp: ABI v$$engine matches the pinned engine"
|
||||
|
||||
$(LIB): sources/vllm.cpp $(MLX_STAMP)
|
||||
$(MAKE) abi-check
|
||||
mkdir -p build && \
|
||||
cd build && \
|
||||
cmake ../sources/vllm.cpp $(CMAKE_ARGS) && \
|
||||
cmake ../sources/vllm.cpp $(CMAKE_ARGS) $(MLX_CMAKE_ARGS) && \
|
||||
cmake --build . --config Release -j$(JOBS) --target vllm_shared
|
||||
cp -fL build/$(LIB) ./$(LIB)
|
||||
|
||||
vllm-cpp: main.go govllmcpp.go backend.go options.go $(LIB)
|
||||
vllm-cpp: main.go govllmcpp.go backend.go chat.go options.go video.go $(LIB)
|
||||
CGO_ENABLED=0 $(GOCMD) build -tags "$(GO_TAGS)" -o vllm-cpp ./
|
||||
|
||||
package: vllm-cpp
|
||||
bash package.sh
|
||||
MLX_ROOT="$(MLX_ROOT)" bash package.sh
|
||||
|
||||
build: package
|
||||
|
||||
clean: purge
|
||||
rm -rf libvllm.so libvllm.dylib package sources/vllm.cpp vllm-cpp
|
||||
rm -rf libvllm.so libvllm.dylib package sources/vllm.cpp vllm-cpp "$(MLX_VENV)"
|
||||
|
||||
purge:
|
||||
rm -rf build
|
||||
|
||||
.PHONY: abi-check
|
||||
|
||||
.NOTPARALLEL:
|
||||
|
||||
# The unit specs are pure Go (struct mirrors, option mapping, load
|
||||
|
||||
@@ -1,12 +1,15 @@
|
||||
# vllm-cpp backend
|
||||
|
||||
LocalAI text-generation backend for [vllm.cpp](https://github.com/mudler/vllm.cpp),
|
||||
LocalAI backend for [vllm.cpp](https://github.com/mudler/vllm.cpp),
|
||||
the LocalAI-team C++20 port of vLLM (paged KV cache, continuous batching,
|
||||
safetensors + GGUF loading, CUDA / CPU / Metal / Vulkan) with no Python at
|
||||
inference time.
|
||||
|
||||
It serves two things: text generation, and MiniMax-H3 joint video+audio
|
||||
generation.
|
||||
|
||||
The backend dlopens the engine's stable C ABI (`libvllm`, `include/vllm.h`,
|
||||
ABI v2) through purego:
|
||||
ABI v20) through purego:
|
||||
|
||||
- `Load` -> `vllm_engine_load`: accepts a `.gguf` file or a HF-style model
|
||||
directory (`config.json` + safetensors). `context_size` maps to
|
||||
@@ -29,6 +32,18 @@ ABI v2) through purego:
|
||||
LocalAI's Go-side grammar-constrained tool calling; JSON-schema / regex /
|
||||
choice constraints are also exposed by the ABI.
|
||||
|
||||
`patches/` carries fixes the pinned engine SHA does not have yet, applied to
|
||||
the clone the same way `longcat-video` patches its upstream. `git apply` is
|
||||
unguarded on purpose: a patch that stops applying must fail the clone loudly
|
||||
rather than leave a pin silently missing a fix it is documented to carry. Each
|
||||
patch header says what retires it.
|
||||
|
||||
The struct mirrors in `govllmcpp.go` are hand-written against one ABI version,
|
||||
and the engine refuses to load against any other. Moving `VLLM_CPP_VERSION` in
|
||||
the Makefile therefore means updating `abiVersion` plus the mirrors (and their
|
||||
offsets in `vllmcpp_test.go`) in the same change; `make abi-check` compares the
|
||||
pinned header against the bindings and the library build runs it first.
|
||||
|
||||
Model config example:
|
||||
|
||||
```yaml
|
||||
@@ -41,5 +56,112 @@ options:
|
||||
- max_num_seqs:16
|
||||
```
|
||||
|
||||
## MiniMax-H3 video+audio generation
|
||||
|
||||
`GenerateVideo` -> `vllm_video_generate` (ABI v12). H3 renders picture and sound
|
||||
together, so the output MP4 carries a real AAC track.
|
||||
|
||||
The video engine is a SECOND handle (`vllm_video_engine`), not a mode of the
|
||||
text one, because H3 is a checkpoint SET rather than a model directory: the DiT,
|
||||
the text encoder and two VAEs are separate artifacts, and vllm.cpp has the two
|
||||
loaders refuse each other's checkpoints. `Load` takes the video branch when the
|
||||
model config carries any of the video options below; `parameters.model` is the
|
||||
DiT and everything else is named in `options:`.
|
||||
|
||||
```yaml
|
||||
name: minimax-h3-fl2va-q4
|
||||
backend: vllm-cpp
|
||||
cuda: true
|
||||
known_usecases: [video]
|
||||
parameters:
|
||||
model: minimax-h3/MiniMax-H3-FL2VA-Q4_K_M.gguf
|
||||
options:
|
||||
- video_encoder:minimax-h3/qwen3vl-32B-MiniMax-H3-Q4_K_M.gguf
|
||||
- video_tokenizer:minimax-h3/tokenizer.json
|
||||
- video_vae:minimax-h3/video_vae.safetensors
|
||||
- video_vae_config:minimax-h3/video_vae_config.json
|
||||
- audio_vae:minimax-h3/audio_vae.safetensors
|
||||
- audio_vae_config:minimax-h3/audio_vae_config.json
|
||||
- video_partition:fl2va
|
||||
- video_device:cuda
|
||||
- video_dequant_bf16:true
|
||||
- video_width:1344
|
||||
- video_height:768
|
||||
- video_num_frames:124
|
||||
```
|
||||
|
||||
Three things are worth knowing before touching this path.
|
||||
|
||||
**The partition is declared, not detected, and a mismatch does not fail
|
||||
cleanly.** The FL2VA DiT serves `t2va` and `fl2va`; `ref2va` is a different
|
||||
checkpoint. The community GGUF/NVFP4 quantisations strip the release metadata
|
||||
and the two DiTs are byte-structurally identical, so the engine refuses every
|
||||
generate until `video_partition` says which one it has. Handing reference
|
||||
conditioning to an FL2VA DiT renders for hours and returns a coloured lattice
|
||||
over the frame, so `checkPartitionConditioning` refuses that combination here,
|
||||
before the engine is called.
|
||||
|
||||
**ffmpeg comes from the host.** libvllm writes the frames and the WAV and
|
||||
COMPOSES the mux argv, then spawns nothing — that process boundary is upstream's
|
||||
decision. `muxVideo` takes the composed argv, substitutes `argv[0]` with the
|
||||
resolved binary and execs it; the backend image is `FROM scratch` and carries no
|
||||
ffmpeg, the same arrangement `vibevoice-cpp` uses for transcoding. ffmpeg also
|
||||
converts a `start_image`/`end_image` upload into the binary PPM at the exact
|
||||
output canvas the engine requires, since libvllm vendors neither an image codec
|
||||
nor a resampler.
|
||||
|
||||
**It is slow.** Roughly 176 s per denoise step at 1344x768 on a 20-SM device, so
|
||||
the 50-step default is hours. Nothing here imposes a deadline.
|
||||
|
||||
Geometry mirrors the engine so the two agree: the canvas is truncated onto a
|
||||
32-pixel grid, the frame count sits on the 17n+5 grid, and an unspecified canvas
|
||||
with a keyframe is derived from that image's aspect on a 768-pixel short edge
|
||||
(`MiniMaxH3ResolveShape`, `minimax_h3_planner.cpp`).
|
||||
|
||||
## Apple Silicon: the MLX GEMM provider (ON by default, gated to prefill)
|
||||
|
||||
`BUILD_TYPE=metal` builds vllm.cpp's MLX provider for the dense GEMM
|
||||
(`VLLM_CPP_MLX=on`, the default here). It is on because upstream now SHAPE-GATES
|
||||
it to prefill; it was briefly off in this branch's history, and that was correct
|
||||
at the time for an ungated provider.
|
||||
|
||||
The gate matters more than the flag. MLX's steel GEMM wins prefill but loses
|
||||
decode, because the provider pays an `mx::eval` synchronisation plus an output
|
||||
memcpy on every call and decode makes ~112 calls *per token*. Measured on an
|
||||
Apple M4, Qwen3-1.7B-bf16 warm at p=512 g=128:
|
||||
|
||||
| configuration | prefill TTFT | warm throughput |
|
||||
|---|--:|--:|
|
||||
| MLX **gated to prefill** (pin >= 89c46aeb) | **524.5 ms** | **24.37 tok/s, 97.6% of MLX-LM** |
|
||||
| MLX ungated (older pins) | 537 ms | 12.7 tok/s |
|
||||
| MLX off | 602 ms | 23.9 tok/s, 95.9% |
|
||||
|
||||
Ratios are against an MLX-LM baseline measured INTERLEAVED with ours over four
|
||||
ABBA blocks (its spread 0.34%, ours 0.12%). An earlier revision of this file
|
||||
claimed 99.1%; that used a two-run MLX-LM baseline containing an outlier and
|
||||
overstated us by about 1.5 points.
|
||||
|
||||
**`VLLM_CPP_VERSION` and this flag are coupled.** Moving the pin back before
|
||||
`89c46aeb` while leaving `VLLM_CPP_MLX=on` would take the middle row — roughly
|
||||
half throughput. If you roll the pin back, roll the default back with it.
|
||||
|
||||
One caveat: MLX's GEMM is not bit-identical to the native kernel, so an MLX build
|
||||
produces a different greedy sequence than a non-MLX one. That is a property of the
|
||||
provider, not of the gate, and it predates this packaging. Full disposition in
|
||||
vllm.cpp `docs/BENCHMARKS.md`.
|
||||
|
||||
Build knobs:
|
||||
|
||||
- `VLLM_CPP_MLX=off` builds Metal without the provider: ~124 MB smaller, and
|
||||
96.4% of MLX-LM instead of 99.1%.
|
||||
- `MLX_VERSION` pins the wheel (default `0.29.4`). MLX is consumed as the
|
||||
prebuilt pip wheel because building it from source needs `xcrun metal`, i.e. a
|
||||
full Xcode the macOS runners do not have.
|
||||
|
||||
Packaging vendors `libmlx.dylib`, `mlx.metallib` and MLX's MIT license into
|
||||
`package/lib/`, and rewrites `libvllm.dylib`'s rpath to `@loader_path/lib`
|
||||
(re-signing it, since `install_name_tool` invalidates the signature). The
|
||||
metallib must stay beside `libmlx.dylib`: MLX looks for it there.
|
||||
|
||||
Testing: `make test` runs the unit specs; export `VLLM_CPP_MODEL=<model>` (and
|
||||
optionally `VLLM_CPP_LIBRARY=<libvllm path>`) to enable the e2e specs.
|
||||
@@ -28,7 +28,12 @@ type VllmCpp struct {
|
||||
base.Base
|
||||
|
||||
engine uintptr
|
||||
opts loadOptions
|
||||
// videoEngine is the MiniMax-H3 handle (ABI v12). It is deliberately a
|
||||
// SECOND handle, not a mode of the first: H3 is a checkpoint set rather
|
||||
// than a model directory, and vllm.cpp has the two loaders refuse each
|
||||
// other's checkpoints. Exactly one of the two is ever non-zero.
|
||||
videoEngine uintptr
|
||||
opts loadOptions
|
||||
}
|
||||
|
||||
// Stream registry: the per-request bridge between the C token callback and
|
||||
@@ -109,6 +114,24 @@ func (v *VllmCpp) Load(opts *pb.ModelOptions) error {
|
||||
|
||||
v.opts = parseOptions(opts)
|
||||
|
||||
// MiniMax-H3 is a checkpoint SET behind its own engine handle, so the
|
||||
// branch is taken before any text-engine knob is resolved. The two loaders
|
||||
// refuse each other's checkpoints, which is why this is decided from the
|
||||
// config rather than probed.
|
||||
if v.opts.video.engaged() {
|
||||
return v.loadVideo(opts, model)
|
||||
}
|
||||
|
||||
// A DFlash draft is a second checkpoint the engine opens by path, and the
|
||||
// engine never downloads one. Resolve it against LocalAI's models directory
|
||||
// now so a repo-id spelling works, and so a missing draft fails here with an
|
||||
// actionable message rather than as an HF-cache miss inside the load.
|
||||
resolvedSpec, err := resolveDraftModelPath(v.opts.speculativeConfig, opts.ModelPath)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
v.opts.speculativeConfig = resolvedSpec
|
||||
|
||||
mp := defaultModelParams()
|
||||
if v.opts.blockSize > 0 {
|
||||
mp.BlockSize = v.opts.blockSize
|
||||
@@ -116,34 +139,62 @@ func (v *VllmCpp) Load(opts *pb.ModelOptions) error {
|
||||
if v.opts.numBlocks > 0 {
|
||||
mp.NumBlocks = v.opts.numBlocks
|
||||
}
|
||||
// Sequence-length precedence, narrowest source last: context_size is the
|
||||
// generic LocalAI knob every backend honours, max_model_len is the
|
||||
// vLLM-specific one, and engine_args.max_model_len is the explicit
|
||||
// vllm-cpp override.
|
||||
if opts.ContextSize > 0 {
|
||||
mp.MaxModelLen = opts.ContextSize
|
||||
}
|
||||
if opts.MaxModelLen > 0 {
|
||||
mp.MaxModelLen = opts.MaxModelLen
|
||||
}
|
||||
if v.opts.maxModelLen > 0 {
|
||||
mp.MaxModelLen = v.opts.maxModelLen
|
||||
}
|
||||
if v.opts.maxNumSeqs > 0 {
|
||||
mp.MaxNumSeqs = v.opts.maxNumSeqs
|
||||
}
|
||||
if v.opts.maxNumBatchedTokens > 0 {
|
||||
mp.MaxNumBatchedTokens = v.opts.maxNumBatchedTokens
|
||||
}
|
||||
mp.EnablePrefixCaching = v.opts.enablePrefixCaching
|
||||
mp.EnableJumpForward = v.opts.enableJumpForward
|
||||
|
||||
// Every string below is borrowed by C for the duration of the load call
|
||||
// only (the library copies what it keeps), so the backing slices just have
|
||||
// to outlive vllmEngineLoad - hence the single KeepAlive after it.
|
||||
modelC := cString(model)
|
||||
mp.ModelPath = uintptr(unsafe.Pointer(&modelC[0])) // #nosec G103 -- borrowed by C for the load call only
|
||||
var toolParserC, reasoningParserC []byte
|
||||
if v.opts.toolParser != "" {
|
||||
toolParserC = cString(v.opts.toolParser)
|
||||
mp.ToolParser = uintptr(unsafe.Pointer(&toolParserC[0])) // #nosec G103 -- borrowed by C for the load call only
|
||||
}
|
||||
if v.opts.reasoningParser != "" {
|
||||
reasoningParserC = cString(v.opts.reasoningParser)
|
||||
mp.ReasoningParser = uintptr(unsafe.Pointer(&reasoningParserC[0])) // #nosec G103 -- borrowed by C for the load call only
|
||||
keep := [][]byte{modelC}
|
||||
setStr := func(dst *uintptr, s string) {
|
||||
if s == "" {
|
||||
return
|
||||
}
|
||||
b := cString(s)
|
||||
keep = append(keep, b)
|
||||
*dst = uintptr(unsafe.Pointer(&b[0])) // #nosec G103 -- borrowed by C for the load call only
|
||||
}
|
||||
setStr(&mp.ToolParser, v.opts.toolParser)
|
||||
setStr(&mp.ReasoningParser, v.opts.reasoningParser)
|
||||
setStr(&mp.SpeculativeConfig, v.opts.speculativeConfig)
|
||||
setStr(&mp.KVTransferConfig, v.opts.kvTransferConfig)
|
||||
setStr(&mp.SchedulingPolicy, v.opts.schedulingPolicy)
|
||||
setStr(&mp.TokenizerConfigPath, v.opts.tokenizerConfigPath)
|
||||
|
||||
xlog.Info("[vllm-cpp] Load", "model", model, "engine", vllmVersion(),
|
||||
"blockSize", mp.BlockSize, "numBlocks", mp.NumBlocks,
|
||||
"maxModelLen", mp.MaxModelLen, "maxNumSeqs", mp.MaxNumSeqs)
|
||||
"maxModelLen", mp.MaxModelLen, "maxNumSeqs", mp.MaxNumSeqs,
|
||||
"maxNumBatchedTokens", mp.MaxNumBatchedTokens,
|
||||
"prefixCaching", triStateName(mp.EnablePrefixCaching),
|
||||
"jumpForward", triStateName(mp.EnableJumpForward),
|
||||
"schedulingPolicy", v.opts.schedulingPolicy,
|
||||
"speculativeConfig", v.opts.speculativeConfig,
|
||||
"kvTransferConfig", v.opts.kvTransferConfig)
|
||||
|
||||
var engine uintptr
|
||||
rc := vllmEngineLoad(unsafe.Pointer(&mp), unsafe.Pointer(&engine)) // #nosec G103 -- POD out-params
|
||||
runtime.KeepAlive(modelC)
|
||||
runtime.KeepAlive(toolParserC)
|
||||
runtime.KeepAlive(reasoningParserC)
|
||||
runtime.KeepAlive(keep)
|
||||
if rc != vllmOK {
|
||||
return fmt.Errorf("vllm-cpp: engine load failed: %s", vllmLastError())
|
||||
}
|
||||
@@ -156,6 +207,10 @@ func (v *VllmCpp) Free() error {
|
||||
vllmEngineFree(v.engine)
|
||||
v.engine = 0
|
||||
}
|
||||
if v.videoEngine != 0 {
|
||||
vllmVideoEngineFree(v.videoEngine)
|
||||
v.videoEngine = 0
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
package main
|
||||
|
||||
// purego bindings for the vllm.cpp stable C ABI (include/vllm.h, ABI v2).
|
||||
// purego bindings for the vllm.cpp stable C ABI (include/vllm.h, ABI v21).
|
||||
//
|
||||
// The structs below are hand-mirrored PODs of the C declarations, with
|
||||
// explicit padding so the Go layout matches the C layout on linux/darwin
|
||||
@@ -17,29 +17,77 @@ import (
|
||||
"github.com/ebitengine/purego"
|
||||
)
|
||||
|
||||
// abiVersion is the VLLM_ABI_VERSION this file mirrors (vllm.h).
|
||||
const abiVersion = 5
|
||||
// abiVersion is the VLLM_ABI_VERSION this file mirrors (vllm.h). It must track
|
||||
// the header of the VLLM_CPP_VERSION pinned in the Makefile: the build checks
|
||||
// the two against each other, because a mismatch is only caught at runtime by
|
||||
// registerLib, where it takes the backend down on every load (issue #11379).
|
||||
const abiVersion = 21
|
||||
|
||||
// The ABI's tri-state toggles (enable_prefix_caching ABI v7,
|
||||
// enable_jump_forward ABI v10) share one encoding: 0 is NOT "off", it is
|
||||
// "defer" - to the model capability for prefix caching, to the environment for
|
||||
// jump forward. Only 2 is an explicit off.
|
||||
const (
|
||||
triStateDefer int32 = 0
|
||||
triStateOn int32 = 1
|
||||
triStateOff int32 = 2
|
||||
)
|
||||
|
||||
// triStateName renders a tri-state for the load log line, where "0" would
|
||||
// otherwise read as "off" rather than "whatever the default resolves to".
|
||||
func triStateName(state int32) string {
|
||||
switch state {
|
||||
case triStateOn:
|
||||
return "on"
|
||||
case triStateOff:
|
||||
return "off"
|
||||
default:
|
||||
return "model-default"
|
||||
}
|
||||
}
|
||||
|
||||
// vllm_status (vllm.h).
|
||||
const (
|
||||
vllmOK = 0
|
||||
)
|
||||
|
||||
// cModelParams mirrors vllm_model_params.
|
||||
// cModelParams mirrors vllm_model_params. The int32 fields sit in pairs so the
|
||||
// interior needs no padding on LP64, but the struct is 8-aligned (it holds
|
||||
// pointers) and ends on a lone int32, so the trailing pad is explicit. Offsets
|
||||
// and total size are asserted in vllmcpp_test.go.
|
||||
type cModelParams struct {
|
||||
ModelPath uintptr // const char*
|
||||
TokenizerConfigPath uintptr // const char*
|
||||
TokenizerConfigPath uintptr // const char*; NULL = <model_dir>/... (ABI v9)
|
||||
BlockSize int32
|
||||
NumBlocks int32
|
||||
MaxModelLen int32
|
||||
MaxNumSeqs int32
|
||||
ToolParser uintptr // const char*; NULL = auto-detect (ABI v4)
|
||||
ReasoningParser uintptr // const char*; NULL = auto-detect (ABI v5)
|
||||
SpeculativeConfig uintptr // const char* JSON; NULL = no speculation (ABI v6)
|
||||
EnablePrefixCaching int32 // tri-state 0/1/2 (ABI v7)
|
||||
MaxNumBatchedTokens int32 // <= 0 = per-arch default (ABI v9)
|
||||
SchedulingPolicy uintptr // const char*; NULL = "fcfs" (ABI v9)
|
||||
KVTransferConfig uintptr // const char* JSON; NULL = no connector (ABI v9)
|
||||
OffloadConfig uintptr // const char* JSON; NULL = no weight offload
|
||||
EnableJumpForward int32 // tri-state 0/1/2 (ABI v10)
|
||||
// v14/v16 tail. LocalAI sets none of these (0 is "auto" for the device and
|
||||
// "unset" for both sizing knobs, i.e. the pre-v14 engine byte for byte), but
|
||||
// the fields MUST be mirrored: the C side reads sizeof(vllm_model_params)
|
||||
// bytes off the pointer we hand it, so a Go struct that stopped at
|
||||
// EnableJumpForward would have vllm_engine_load read 24 bytes past our
|
||||
// allocation and size the KV pool from whatever sat there.
|
||||
Device int32 // 0 auto, 1 cpu, 2 cuda (ABI v14)
|
||||
GPUMemoryUtil float64 // 0 => 0.92 (ABI v16)
|
||||
KVCacheMemoryBytes int64 // 0 => unset (ABI v16)
|
||||
LanguageModelOnly int32 // 0 = multimodal inputs enabled (ABI v19)
|
||||
_ [4]byte
|
||||
LimitMMPerPrompt uintptr // const char* JSON; NULL = default limits (ABI v19)
|
||||
}
|
||||
|
||||
// cSamplingParams mirrors vllm_sampling_params (ABI v2, structured fields
|
||||
// included). Padding matches the C compiler's: the uint64 seed is 8-aligned,
|
||||
// and each pointer following an int32 is 8-aligned.
|
||||
// cSamplingParams mirrors vllm_sampling_params (structured fields included).
|
||||
// Padding matches the C compiler's: the uint64 seed is 8-aligned, and each
|
||||
// pointer following an int32 is 8-aligned.
|
||||
type cSamplingParams struct {
|
||||
Temperature float32
|
||||
TopP float32
|
||||
@@ -65,6 +113,12 @@ type cSamplingParams struct {
|
||||
StructuredGrammar uintptr // const char*
|
||||
StructuredJSONObject int32
|
||||
_ [4]byte
|
||||
// ABI v8 tail. LocalAI installs no custom logits processor, but the fields
|
||||
// MUST be mirrored: the C side reads them off the pointer we hand it, so a
|
||||
// Go struct that stopped at StructuredJSONObject would have the engine read
|
||||
// 16 bytes past our allocation and call whatever garbage sat there.
|
||||
LogitsProcessor uintptr // vllm_logits_processor; NULL = none
|
||||
LogitsProcessorUserData uintptr // void*
|
||||
}
|
||||
|
||||
// cCompletion mirrors vllm_completion.
|
||||
@@ -75,6 +129,88 @@ type cCompletion struct {
|
||||
CompletionTokens int32
|
||||
}
|
||||
|
||||
// ── Video+audio generation (ABI v12, MiniMax-H3) ────────────────────────────
|
||||
//
|
||||
// A video engine is a SEPARATE handle from vllm_engine: H3 is a checkpoint SET
|
||||
// (DiT + text encoder + two VAEs), not one model directory, and the two loaders
|
||||
// refuse each other's checkpoints on purpose. Offsets are asserted in
|
||||
// video_test.go the same way the text PODs are in vllmcpp_test.go.
|
||||
|
||||
// cVideoModelParams mirrors vllm_video_model_params. Nine pointers then three
|
||||
// int32s, so only the trailing pad is implicit.
|
||||
type cVideoModelParams struct {
|
||||
DitPath uintptr // const char*
|
||||
EncoderPath uintptr // const char*
|
||||
TokenizerPath uintptr // const char*
|
||||
VideoVaePath uintptr // const char*
|
||||
VideoVaeConfigPath uintptr // const char*
|
||||
AudioVaePath uintptr // const char*
|
||||
AudioVaeConfigPath uintptr // const char*
|
||||
PromptEmbedsPath uintptr // const char*
|
||||
Partition uintptr // const char*; "fl2va" | "ref2va", REQUIRED
|
||||
Device int32 // 0 cpu, 1 cuda
|
||||
DequantBf16 int32 // 0 keep-quant, 1 dequant/stream bf16
|
||||
Fp4Resident int32 // NVFP4+cuda: keep FP4 packed, Marlin W4A16
|
||||
_ [4]byte
|
||||
Family uintptr // const char*; NULL = detect (ABI v18)
|
||||
ExtraKeys uintptr // const char* const* (ABI v18)
|
||||
ExtraValues uintptr // const char* const* (ABI v18)
|
||||
NExtras int32 // 0 = none (ABI v18)
|
||||
_ [4]byte // trailing pad to the struct's 8-byte alignment
|
||||
}
|
||||
|
||||
// cVideoParams mirrors vllm_video_params. `width`/`height` and `num_frames`/
|
||||
// `steps` pair up into 8-byte slots; the uint64 seed forces the alignment after
|
||||
// them, and the float noise_aug leaves a pad before output_dir.
|
||||
type cVideoParams struct {
|
||||
Prompt uintptr // const char*
|
||||
Width int32
|
||||
Height int32
|
||||
NumFrames int32 // <= 1 => per-task default (124 for t2va/fl2va)
|
||||
Steps int32 // <= 0 => the H3 default (50)
|
||||
Seed uint64
|
||||
HasSeed int32
|
||||
_ [4]byte
|
||||
FirstFrame uintptr // const char*; fl2va keyframe, binary PPM (P6)
|
||||
LastFrame uintptr // const char*
|
||||
RefImage uintptr // const char*; ref2va only
|
||||
RefVideo uintptr // const char*; ref2va only, a frame_%06d.ppm DIRECTORY
|
||||
RefAudio uintptr // const char*; ref2va only, 16-bit PCM WAV
|
||||
NoiseAug float32 // <= 0 => 1.0
|
||||
_ [4]byte
|
||||
OutputDir uintptr // const char*; REQUIRED
|
||||
ExtraKeys uintptr // const char* const* (ABI v18)
|
||||
ExtraValues uintptr // const char* const* (ABI v18)
|
||||
NExtras int32 // 0 = none (ABI v18)
|
||||
_ [4]byte
|
||||
}
|
||||
|
||||
// cVideoResult mirrors vllm_video_result. Every member is library-allocated and
|
||||
// released together by vllm_video_result_free.
|
||||
type cVideoResult struct {
|
||||
FrameDir uintptr // char*, holds frame_%06d.ppm
|
||||
AudioPath uintptr // char*, 16-bit PCM WAV
|
||||
FrameCount int32
|
||||
Width int32
|
||||
Height int32
|
||||
Fps int32
|
||||
SampleRate int32
|
||||
_ [4]byte
|
||||
MuxArgv uintptr // char**, NULL-terminated at MuxArgc
|
||||
MuxArgc int32
|
||||
_ [4]byte
|
||||
}
|
||||
|
||||
// cVideoMuxParams mirrors vllm_video_mux_params. The library composes the argv;
|
||||
// spawning it is the CALLER's job, which is why no ffmpeg lives in libvllm.
|
||||
type cVideoMuxParams struct {
|
||||
Frames uintptr // const char*; printf pattern, dir/frame_%06d.ppm
|
||||
AudioPath uintptr // const char*; NULL/empty => a silent clip
|
||||
OutputPath uintptr // const char*; the .mp4 to write
|
||||
Fps int32 // <= 0 => the H3 default (24)
|
||||
Crf int32 // <= 0 => the library default (18)
|
||||
}
|
||||
|
||||
// defaultSamplingParams mirrors vllm_sampling_params_default().
|
||||
func defaultSamplingParams() cSamplingParams {
|
||||
return cSamplingParams{
|
||||
@@ -106,6 +242,14 @@ var (
|
||||
vllmLastError func() string
|
||||
vllmVersion func() string
|
||||
vllmABIVersion func() int32
|
||||
|
||||
// Video+audio generation (ABI v12).
|
||||
vllmVideoEngineLoad func(params, out unsafe.Pointer) int32
|
||||
vllmVideoEngineFree func(engine uintptr)
|
||||
vllmVideoGenerate func(engine uintptr, params, out unsafe.Pointer) int32
|
||||
vllmVideoResultFree func(out unsafe.Pointer)
|
||||
vllmVideoMuxArgv func(params, outArgv, outArgc unsafe.Pointer) int32
|
||||
vllmVideoMuxArgvFre func(argv uintptr, argc int32)
|
||||
)
|
||||
|
||||
type libFunc struct {
|
||||
@@ -133,6 +277,12 @@ func registerLib(libName string) error {
|
||||
{&vllmLastError, "vllm_last_error"},
|
||||
{&vllmVersion, "vllm_version"},
|
||||
{&vllmABIVersion, "vllm_abi_version"},
|
||||
{&vllmVideoEngineLoad, "vllm_video_engine_load"},
|
||||
{&vllmVideoEngineFree, "vllm_video_engine_free"},
|
||||
{&vllmVideoGenerate, "vllm_video_generate"},
|
||||
{&vllmVideoResultFree, "vllm_video_result_free"},
|
||||
{&vllmVideoMuxArgv, "vllm_video_mux_argv"},
|
||||
{&vllmVideoMuxArgvFre, "vllm_video_mux_argv_free"},
|
||||
} {
|
||||
purego.RegisterLibFunc(lf.ptr, lib, lf.name)
|
||||
}
|
||||
@@ -180,3 +330,19 @@ func goString(p uintptr) string {
|
||||
}
|
||||
return string(unsafe.Slice((*byte)(base), n))
|
||||
}
|
||||
|
||||
// goStringSlice copies a C `char*` array of n entries. Used for the ffmpeg argv
|
||||
// the library composes: it is copied out immediately so the caller can free the
|
||||
// C allocation before ever spawning the process.
|
||||
func goStringSlice(p uintptr, n int32) []string {
|
||||
if p == 0 || n <= 0 {
|
||||
return nil
|
||||
}
|
||||
//nolint:govet // C-owned pointer handed over by purego, valid for this call
|
||||
entries := unsafe.Slice((**byte)(unsafe.Pointer(p)), int(n)) // #nosec G103 -- C-owned, copied out immediately
|
||||
out := make([]string, 0, n)
|
||||
for _, e := range entries {
|
||||
out = append(out, goString(uintptr(unsafe.Pointer(e)))) // #nosec G103 -- ditto
|
||||
}
|
||||
return out
|
||||
}
|
||||
@@ -1,30 +1,140 @@
|
||||
package main
|
||||
|
||||
// Engine-sizing knobs carried through the model config's free-form
|
||||
// `options:` list ("key:value" entries), mirroring how the other in-house
|
||||
// backends pass engine-specific settings that have no proto field.
|
||||
// Load-time engine configuration, from two config surfaces:
|
||||
//
|
||||
// - `engine_args:` (ModelOptions.EngineArgs, a JSON object) is the canonical
|
||||
// one. Keys are spelled exactly as vLLM's own CLI flags, so a config written
|
||||
// against vLLM works verbatim here - `speculative_config` and
|
||||
// `kv_transfer_config` in particular take the same JSON documents vLLM's
|
||||
// --speculative-config / --kv-transfer-config accept, and are handed to the
|
||||
// engine unparsed.
|
||||
// - `options:` (the free-form "key:value" list) is the older surface this
|
||||
// backend shipped with. It is still honoured so existing configs keep
|
||||
// working; engine_args wins on any key set in both.
|
||||
//
|
||||
// Anything unrecognised is ignored rather than fatal: the engine validates the
|
||||
// documents it is given and reports a precise error at load, and a config that
|
||||
// also carries knobs for a different backend must not fail the load here.
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"path"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
|
||||
"github.com/mudler/xlog"
|
||||
)
|
||||
|
||||
type loadOptions struct {
|
||||
blockSize int32 // KV block size (tokens/block); engine default 32.
|
||||
numBlocks int32 // KV blocks to allocate; engine default 256.
|
||||
maxNumSeqs int32 // max concurrent sequences; engine default 8.
|
||||
// Max sequence length. Also settable through the model config's
|
||||
// context_size / max_model_len; see Load for the precedence.
|
||||
maxModelLen int32
|
||||
// Per-step chunked-prefill token budget (ABI v9). 0 = the engine's
|
||||
// bounded per-arch default.
|
||||
maxNumBatchedTokens int32
|
||||
// Automatic prefix caching tri-state (ABI v7): 0 = the model-capability
|
||||
// default, 1 = force on, 2 = force off.
|
||||
enablePrefixCaching int32
|
||||
// Jump-forward decoding tri-state (ABI v10), SGLang's grammar-speed subset:
|
||||
// 0 = defer to the environment (VT_ENABLE_JUMP_FORWARD, default off),
|
||||
// 1 = force on, 2 = force off.
|
||||
enableJumpForward int32
|
||||
// Scheduler admission policy (ABI v9): "" = fcfs, else fcfs|priority|lpm.
|
||||
schedulingPolicy string
|
||||
// Engine-side parser selection (ABI v4/v5). Empty = the engine
|
||||
// auto-detects from the chat template; "none" disables the reasoning
|
||||
// split; unknown names fail the first chat call.
|
||||
toolParser string
|
||||
reasoningParser string
|
||||
// Speculative decoding (ABI v6), as vLLM's --speculative-config JSON:
|
||||
// {"method":"mtp"|"dflash"|"ngram", ...}. Empty = no speculation.
|
||||
speculativeConfig string
|
||||
// External KV connector / LMCache (ABI v9), as vLLM's --kv-transfer-config
|
||||
// JSON. Empty = no connector.
|
||||
kvTransferConfig string
|
||||
// Override for the tokenizer_config.json the chat template is read from
|
||||
// (ABI v9). Empty = <model_dir>/tokenizer_config.json.
|
||||
tokenizerConfigPath string
|
||||
// MiniMax-H3 video+audio generation (ABI v12). Present only when the config
|
||||
// carries at least one of its keys; see videoOptions.engaged.
|
||||
video videoOptions
|
||||
}
|
||||
|
||||
// videoOptions is the MiniMax-H3 checkpoint SET plus its generation defaults.
|
||||
//
|
||||
// H3 is not one model directory: the DiT, the text encoder and the two VAEs are
|
||||
// separate artifacts, which is why vllm.cpp gives video its own engine handle
|
||||
// (vllm_video_engine, ABI v12) rather than another vllm_engine. The DiT is the
|
||||
// model config's `parameters.model`; everything else arrives through these
|
||||
// options, so one gallery entry can name five files.
|
||||
//
|
||||
// The geometry/frame defaults exist because H3's trained canvas is nothing like
|
||||
// the generic /video defaults: 1344x768 at 124 frames is a ~5.2 s clip, and the
|
||||
// frame count must sit on the 17n+5 grid. A request that leaves a field unset
|
||||
// gets the model's own default from here instead of a canvas the checkpoint was
|
||||
// never trained at.
|
||||
type videoOptions struct {
|
||||
encoderPath string // H3-Encoder GGUF or bf16 shard dir
|
||||
tokenizerPath string // tokenizer.json, needed with an encoder
|
||||
videoVaePath string
|
||||
videoVaeConfig string
|
||||
audioVaePath string
|
||||
audioVaeConfig string
|
||||
promptEmbedsPath string // fallback conditioning when there is no encoder
|
||||
// The served checkpoint PARTITION. Community GGUF/NVFP4 files strip the
|
||||
// release metadata and the FL2VA/Ref2VA DiTs are byte-structurally
|
||||
// identical, so the engine refuses every generate until it is DECLARED.
|
||||
// "fl2va" serves t2va + fl2va; "ref2va" serves reference conditioning.
|
||||
partition string
|
||||
device int32 // 0 cpu, 1 cuda (the ABI's own encoding, no auto slot)
|
||||
deviceSet bool
|
||||
dequantBf16 int32
|
||||
fp4Resident int32
|
||||
// Per-model generation defaults, applied when the request leaves the field
|
||||
// at 0.
|
||||
width int32
|
||||
height int32
|
||||
numFrames int32
|
||||
steps int32
|
||||
// Where frames + WAV are written. Empty = a temporary directory beside the
|
||||
// requested output, removed once the mux succeeds. Set it to keep the
|
||||
// frame_%06d.ppm runs around (they are what ref2va's ref_video consumes).
|
||||
workdir string
|
||||
// The ffmpeg binary the composed mux argv is exec'd with. Empty = "ffmpeg"
|
||||
// from PATH. libvllm composes the argv and spawns nothing, by design.
|
||||
ffmpeg string
|
||||
crf int32
|
||||
}
|
||||
|
||||
// engaged reports whether this config describes an H3 video engine. Load uses
|
||||
// it to choose which of the two mutually exclusive engine handles to open: the
|
||||
// checkpoints refuse each other, so guessing is not an option, and every key
|
||||
// below is meaningless to the text engine.
|
||||
func (v videoOptions) engaged() bool {
|
||||
return v.encoderPath != "" || v.tokenizerPath != "" ||
|
||||
v.videoVaePath != "" || v.videoVaeConfig != "" ||
|
||||
v.audioVaePath != "" || v.audioVaeConfig != "" ||
|
||||
v.promptEmbedsPath != "" || v.partition != ""
|
||||
}
|
||||
|
||||
func parseOptions(opts *pb.ModelOptions) loadOptions {
|
||||
lo := loadOptions{}
|
||||
for _, o := range opts.GetOptions() {
|
||||
applyOptionsList(&lo, opts.GetOptions())
|
||||
applyEngineArgs(&lo, opts.GetEngineArgs())
|
||||
return lo
|
||||
}
|
||||
|
||||
// applyOptionsList reads the legacy free-form "key:value" list. strings.Cut
|
||||
// splits on the FIRST colon only, so a JSON object value survives intact.
|
||||
func applyOptionsList(lo *loadOptions, options []string) {
|
||||
for _, o := range options {
|
||||
k, v, found := strings.Cut(o, ":")
|
||||
if !found {
|
||||
continue
|
||||
@@ -36,13 +146,298 @@ func parseOptions(opts *pb.ModelOptions) loadOptions {
|
||||
lo.numBlocks = parseInt32(v, lo.numBlocks)
|
||||
case "max_num_seqs":
|
||||
lo.maxNumSeqs = parseInt32(v, lo.maxNumSeqs)
|
||||
case "tool_parser":
|
||||
case "max_num_batched_tokens":
|
||||
lo.maxNumBatchedTokens = parseInt32(v, lo.maxNumBatchedTokens)
|
||||
case "max_model_len":
|
||||
lo.maxModelLen = parseInt32(v, lo.maxModelLen)
|
||||
case "scheduling_policy", "schedule_policy":
|
||||
lo.schedulingPolicy = strings.TrimSpace(v)
|
||||
case "tool_parser", "tool_call_parser":
|
||||
lo.toolParser = strings.TrimSpace(v)
|
||||
case "reasoning_parser":
|
||||
lo.reasoningParser = strings.TrimSpace(v)
|
||||
case "speculative_config":
|
||||
lo.speculativeConfig = strings.TrimSpace(v)
|
||||
case "kv_transfer_config":
|
||||
lo.kvTransferConfig = strings.TrimSpace(v)
|
||||
case "tokenizer_config", "tokenizer_config_path":
|
||||
lo.tokenizerConfigPath = strings.TrimSpace(v)
|
||||
case "enable_prefix_caching", "enable_radix_attention":
|
||||
if b, err := strconv.ParseBool(strings.TrimSpace(v)); err == nil {
|
||||
lo.enablePrefixCaching = boolTriState(b)
|
||||
}
|
||||
case "enable_jump_forward":
|
||||
if b, err := strconv.ParseBool(strings.TrimSpace(v)); err == nil {
|
||||
lo.enableJumpForward = boolTriState(b)
|
||||
}
|
||||
default:
|
||||
applyVideoOption(&lo.video, strings.TrimSpace(k), v)
|
||||
}
|
||||
}
|
||||
return lo
|
||||
}
|
||||
|
||||
// applyVideoOption reads one MiniMax-H3 key. Split out of applyOptionsList so
|
||||
// the video surface stays legible next to the videoOptions it fills, and so
|
||||
// video_test.go can exercise it directly.
|
||||
func applyVideoOption(vo *videoOptions, key, value string) bool {
|
||||
v := strings.TrimSpace(value)
|
||||
switch key {
|
||||
case "video_encoder":
|
||||
vo.encoderPath = v
|
||||
case "video_tokenizer":
|
||||
vo.tokenizerPath = v
|
||||
case "video_vae":
|
||||
vo.videoVaePath = v
|
||||
case "video_vae_config":
|
||||
vo.videoVaeConfig = v
|
||||
case "audio_vae":
|
||||
vo.audioVaePath = v
|
||||
case "audio_vae_config":
|
||||
vo.audioVaeConfig = v
|
||||
case "video_prompt_embeds":
|
||||
vo.promptEmbedsPath = v
|
||||
case "video_partition":
|
||||
vo.partition = strings.ToLower(v)
|
||||
case "video_device":
|
||||
switch strings.ToLower(v) {
|
||||
case "cpu":
|
||||
vo.device, vo.deviceSet = videoDeviceCPU, true
|
||||
case "cuda", "gpu":
|
||||
vo.device, vo.deviceSet = videoDeviceCUDA, true
|
||||
default:
|
||||
xlog.Warn("[vllm-cpp] ignoring unknown video_device", "value", v)
|
||||
}
|
||||
case "video_dequant_bf16":
|
||||
if b, err := strconv.ParseBool(v); err == nil {
|
||||
vo.dequantBf16 = boolInt32(b)
|
||||
}
|
||||
case "video_fp4_resident":
|
||||
if b, err := strconv.ParseBool(v); err == nil {
|
||||
vo.fp4Resident = boolInt32(b)
|
||||
}
|
||||
case "video_width":
|
||||
vo.width = parseInt32(v, vo.width)
|
||||
case "video_height":
|
||||
vo.height = parseInt32(v, vo.height)
|
||||
case "video_num_frames":
|
||||
vo.numFrames = parseInt32(v, vo.numFrames)
|
||||
case "video_steps":
|
||||
vo.steps = parseInt32(v, vo.steps)
|
||||
case "video_workdir":
|
||||
vo.workdir = v
|
||||
case "video_crf":
|
||||
vo.crf = parseInt32(v, vo.crf)
|
||||
case "ffmpeg", "ffmpeg_path":
|
||||
vo.ffmpeg = v
|
||||
default:
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// videoScalarString renders an engine_args scalar so the video keys can share
|
||||
// one parser with the "key:value" list. Objects and arrays have no video
|
||||
// meaning and are left to the caller's unknown-key path.
|
||||
func videoScalarString(v any) (string, bool) {
|
||||
switch t := v.(type) {
|
||||
case string:
|
||||
return t, true
|
||||
case bool:
|
||||
return strconv.FormatBool(t), true
|
||||
case float64:
|
||||
return strconv.FormatFloat(t, 'f', -1, 64), true
|
||||
default:
|
||||
return "", false
|
||||
}
|
||||
}
|
||||
|
||||
func boolInt32(b bool) int32 {
|
||||
if b {
|
||||
return 1
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
// applyEngineArgs overlays the `engine_args:` JSON object. A document that does
|
||||
// not parse is logged and skipped: engine_args is shared with the other engines
|
||||
// (the vLLM and SGLang backends read the same field), so a stray key must not
|
||||
// take the model down.
|
||||
func applyEngineArgs(lo *loadOptions, engineArgs string) {
|
||||
if strings.TrimSpace(engineArgs) == "" {
|
||||
return
|
||||
}
|
||||
var args map[string]any
|
||||
if err := json.Unmarshal([]byte(engineArgs), &args); err != nil {
|
||||
xlog.Warn("[vllm-cpp] ignoring unparseable engine_args", "error", err)
|
||||
return
|
||||
}
|
||||
for k, v := range args {
|
||||
switch k {
|
||||
case "block_size":
|
||||
lo.blockSize = jsonInt32(v, lo.blockSize)
|
||||
case "num_blocks":
|
||||
lo.numBlocks = jsonInt32(v, lo.numBlocks)
|
||||
case "max_num_seqs":
|
||||
lo.maxNumSeqs = jsonInt32(v, lo.maxNumSeqs)
|
||||
case "max_num_batched_tokens":
|
||||
lo.maxNumBatchedTokens = jsonInt32(v, lo.maxNumBatchedTokens)
|
||||
case "max_model_len":
|
||||
lo.maxModelLen = jsonInt32(v, lo.maxModelLen)
|
||||
case "scheduling_policy", "schedule_policy":
|
||||
lo.schedulingPolicy = jsonString(v, lo.schedulingPolicy)
|
||||
case "tool_parser", "tool_call_parser":
|
||||
lo.toolParser = jsonString(v, lo.toolParser)
|
||||
case "reasoning_parser":
|
||||
lo.reasoningParser = jsonString(v, lo.reasoningParser)
|
||||
case "tokenizer_config", "tokenizer_config_path":
|
||||
lo.tokenizerConfigPath = jsonString(v, lo.tokenizerConfigPath)
|
||||
case "speculative_config":
|
||||
lo.speculativeConfig = jsonDocument(v, lo.speculativeConfig, k)
|
||||
case "kv_transfer_config":
|
||||
lo.kvTransferConfig = jsonDocument(v, lo.kvTransferConfig, k)
|
||||
case "enable_prefix_caching", "enable_radix_attention":
|
||||
if b, ok := v.(bool); ok {
|
||||
lo.enablePrefixCaching = boolTriState(b)
|
||||
}
|
||||
case "enable_jump_forward":
|
||||
if b, ok := v.(bool); ok {
|
||||
lo.enableJumpForward = boolTriState(b)
|
||||
}
|
||||
default:
|
||||
if s, ok := videoScalarString(v); ok && applyVideoOption(&lo.video, k, s) {
|
||||
continue
|
||||
}
|
||||
xlog.Debug("[vllm-cpp] ignoring unknown engine_args key", "key", k)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// boolTriState maps a YAML/JSON boolean onto the ABI's tri-state encoding. An
|
||||
// explicit `false` must reach the engine as force-OFF (2), NOT as the 0 that
|
||||
// means "defer". The difference is real in both directions: prefix caching
|
||||
// defaults ON for dense archs and OFF for hybrid ones, and jump forward defers
|
||||
// to VT_ENABLE_JUMP_FORWARD.
|
||||
func boolTriState(on bool) int32 {
|
||||
if on {
|
||||
return triStateOn
|
||||
}
|
||||
return triStateOff
|
||||
}
|
||||
|
||||
// jsonDocument normalises an object-valued engine_args entry to a JSON string
|
||||
// for the C ABI. YAML nesting arrives as a map (the natural spelling); a
|
||||
// pre-encoded JSON string is accepted too, since a config round-tripped through
|
||||
// a flat store may carry it that way.
|
||||
func jsonDocument(v any, fallback string, key string) string {
|
||||
switch t := v.(type) {
|
||||
case string:
|
||||
if strings.TrimSpace(t) == "" {
|
||||
return fallback
|
||||
}
|
||||
return t
|
||||
default:
|
||||
buf, err := json.Marshal(t)
|
||||
if err != nil {
|
||||
xlog.Warn("[vllm-cpp] ignoring unencodable engine_args value", "key", key, "error", err)
|
||||
return fallback
|
||||
}
|
||||
return string(buf)
|
||||
}
|
||||
}
|
||||
|
||||
func jsonString(v any, fallback string) string {
|
||||
s, ok := v.(string)
|
||||
if !ok {
|
||||
return fallback
|
||||
}
|
||||
return strings.TrimSpace(s)
|
||||
}
|
||||
|
||||
// jsonInt32 accepts the float64 a JSON number decodes to, plus the string
|
||||
// spelling a YAML config may produce. Non-positive values keep the fallback:
|
||||
// every knob this covers uses "<= 0 means the engine default".
|
||||
func jsonInt32(v any, fallback int32) int32 {
|
||||
switch t := v.(type) {
|
||||
case float64:
|
||||
if t <= 0 || t > 1<<31-1 {
|
||||
return fallback
|
||||
}
|
||||
return int32(t)
|
||||
case string:
|
||||
return parseInt32(t, fallback)
|
||||
default:
|
||||
return fallback
|
||||
}
|
||||
}
|
||||
|
||||
// resolveDraftModelPath rewrites a DFlash draft reference into an absolute path
|
||||
// the engine can actually open.
|
||||
//
|
||||
// The engine resolves `speculative_config.model` against a directory containing
|
||||
// config.json, or against ~/.cache/huggingface/hub/models--<org>--<repo>/
|
||||
// snapshots/* - and it NEVER downloads. LocalAI keeps models in its own
|
||||
// directory, so a bare HF repo id (the spelling the vLLM docs teach) misses the
|
||||
// HF cache and dies deep in the load with "draft checkpoint not found", which
|
||||
// reads like a broken checkpoint rather than a missing download.
|
||||
//
|
||||
// So: try the reference as given, then the last path segment under the models
|
||||
// dir (`z-lab/Qwen3.6-27B-DFlash` -> `<models>/Qwen3.6-27B-DFlash`, which is
|
||||
// what LocalAI's own downloader produces), then the whole reference under the
|
||||
// models dir. If none exist, fail HERE with a message naming both what was
|
||||
// asked for and where we looked.
|
||||
//
|
||||
// mtp and ngram carry no separate draft checkpoint, so they pass through. A
|
||||
// document that does not parse also passes through: the engine owns config
|
||||
// validation and produces the better error.
|
||||
func resolveDraftModelPath(speculativeConfig, modelsDir string) (string, error) {
|
||||
if strings.TrimSpace(speculativeConfig) == "" {
|
||||
return speculativeConfig, nil
|
||||
}
|
||||
var spec map[string]any
|
||||
if err := json.Unmarshal([]byte(speculativeConfig), &spec); err != nil {
|
||||
return speculativeConfig, nil
|
||||
}
|
||||
if method, _ := spec["method"].(string); !strings.EqualFold(method, "dflash") {
|
||||
return speculativeConfig, nil
|
||||
}
|
||||
|
||||
ref, _ := spec["model"].(string)
|
||||
ref = strings.TrimSpace(ref)
|
||||
if ref == "" {
|
||||
return "", fmt.Errorf(
|
||||
"vllm-cpp: speculative_config method %q requires a \"model\" key naming the draft checkpoint", "dflash")
|
||||
}
|
||||
|
||||
candidates := []string{ref}
|
||||
if modelsDir != "" {
|
||||
if base := path.Base(filepath.ToSlash(ref)); base != "" && base != "." && base != "/" {
|
||||
candidates = append(candidates, filepath.Join(modelsDir, base))
|
||||
}
|
||||
candidates = append(candidates, filepath.Join(modelsDir, filepath.FromSlash(ref)))
|
||||
}
|
||||
|
||||
for _, c := range candidates {
|
||||
if _, err := os.Stat(filepath.Join(c, "config.json")); err != nil {
|
||||
continue
|
||||
}
|
||||
abs, err := filepath.Abs(c)
|
||||
if err != nil {
|
||||
abs = c
|
||||
}
|
||||
spec["model"] = abs
|
||||
out, err := json.Marshal(spec)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("vllm-cpp: re-encoding speculative_config: %w", err)
|
||||
}
|
||||
xlog.Info("[vllm-cpp] resolved DFlash draft checkpoint", "reference", ref, "path", abs)
|
||||
return string(out), nil
|
||||
}
|
||||
|
||||
return "", fmt.Errorf(
|
||||
"vllm-cpp: DFlash draft checkpoint %q not found (looked in: %s). "+
|
||||
"The engine does not download drafts - install the draft model into LocalAI first, "+
|
||||
"or set speculative_config.model to an absolute path to a directory containing config.json",
|
||||
ref, strings.Join(candidates, ", "))
|
||||
}
|
||||
|
||||
func parseInt32(s string, fallback int32) int32 {
|
||||
|
||||
@@ -43,6 +43,50 @@ elif [ -f "/lib/ld-linux-aarch64.so.1" ]; then
|
||||
cp -arfLv /lib/aarch64-linux-gnu/libpthread.so.0 $CURDIR/package/lib/libpthread.so.0
|
||||
elif [ $(uname -s) = "Darwin" ]; then
|
||||
echo "Detected Darwin"
|
||||
# Vendor the optional MLX GEMM provider, when libvllm was built against it.
|
||||
# Three facts drive every line below, each verified on an Apple M4 before it
|
||||
# was written:
|
||||
# 1. libvllm.dylib carries an LC_LOAD_DYLIB on @rpath/libmlx.dylib, and its
|
||||
# build-time LC_RPATH points inside the build venv. That path does not
|
||||
# exist on a user's machine, so it must become @loader_path/lib.
|
||||
# 2. MLX finds its ~100 MB mlx.metallib beside its OWN dylib, so the two
|
||||
# files have to land in the same directory or every Metal op dies with
|
||||
# "Failed to load the default metallib".
|
||||
# 3. install_name_tool invalidates the code signature, and macOS refuses to
|
||||
# load an arm64 image whose signature does not match, so the patched
|
||||
# library must be re-signed ad-hoc afterwards.
|
||||
if otool -L "$CURDIR/package/libvllm.dylib" 2>/dev/null | grep -q "libmlx.dylib"; then
|
||||
MLX_LIB_DIR="${MLX_ROOT}/lib"
|
||||
if [ ! -f "$MLX_LIB_DIR/libmlx.dylib" ] || [ ! -f "$MLX_LIB_DIR/mlx.metallib" ]; then
|
||||
echo "Error: libvllm.dylib links libmlx.dylib but $MLX_LIB_DIR is missing libmlx.dylib/mlx.metallib" >&2
|
||||
exit 1
|
||||
fi
|
||||
echo "Vendoring the MLX GEMM provider from $MLX_LIB_DIR"
|
||||
cp -fLv "$MLX_LIB_DIR/libmlx.dylib" "$CURDIR/package/lib/"
|
||||
cp -fLv "$MLX_LIB_DIR/mlx.metallib" "$CURDIR/package/lib/"
|
||||
# MLX is MIT and we redistribute its binaries, so its license ships with
|
||||
# them. mlx-metal is the wheel carrying the dylib and the metallib.
|
||||
MLX_LICENSE=$(ls "${MLX_ROOT}"/../mlx_metal-*.dist-info/licenses/LICENSE 2>/dev/null | head -1)
|
||||
if [ -z "$MLX_LICENSE" ]; then
|
||||
MLX_LICENSE=$(ls "${MLX_ROOT}"/../mlx-*.dist-info/licenses/LICENSE 2>/dev/null | head -1)
|
||||
fi
|
||||
if [ -z "$MLX_LICENSE" ]; then
|
||||
echo "Error: could not find the MLX LICENSE to redistribute alongside libmlx.dylib" >&2
|
||||
exit 1
|
||||
fi
|
||||
cp -fLv "$MLX_LICENSE" "$CURDIR/package/lib/LICENSE.mlx"
|
||||
# Drop every build-tree rpath, then point at the packaged copy.
|
||||
otool -l "$CURDIR/package/libvllm.dylib" | awk '/LC_RPATH/{f=1;next} f&&/ path /{print $2;f=0}' | while read -r rp; do
|
||||
install_name_tool -delete_rpath "$rp" "$CURDIR/package/libvllm.dylib" 2>/dev/null || true
|
||||
done
|
||||
install_name_tool -add_rpath "@loader_path/lib" "$CURDIR/package/libvllm.dylib"
|
||||
codesign -f -s - "$CURDIR/package/libvllm.dylib"
|
||||
# A broken rpath must fail the BUILD, not the user's first inference.
|
||||
if ! otool -l "$CURDIR/package/libvllm.dylib" | grep -q "@loader_path/lib"; then
|
||||
echo "Error: libvllm.dylib did not get the @loader_path/lib rpath" >&2
|
||||
exit 1
|
||||
fi
|
||||
fi
|
||||
else
|
||||
echo "Error: Could not detect architecture"
|
||||
exit 1
|
||||
|
||||
@@ -0,0 +1,634 @@
|
||||
package main
|
||||
|
||||
// MiniMax-H3 video+audio generation over the vllm.cpp C ABI (v12).
|
||||
//
|
||||
// Two things make this different from the text path, and both come from the
|
||||
// engine's own shape rather than from LocalAI:
|
||||
//
|
||||
// 1. A video engine is loaded from a checkpoint SET - the DiT, the text
|
||||
// encoder and two VAEs are separate artifacts - so it is its own handle
|
||||
// (vllm_video_engine) and its own Load branch. The two loaders refuse each
|
||||
// other's checkpoints on purpose.
|
||||
// 2. libvllm writes frames + a WAV and COMPOSES the ffmpeg argv, but spawns
|
||||
// nothing. That process boundary is deliberate upstream, so the mux lives
|
||||
// here: we take the composed argv, substitute argv[0], and exec it. ffmpeg
|
||||
// comes from PATH the same way the vibevoice-cpp backend takes it.
|
||||
//
|
||||
// Generation is SLOW - roughly 176 s per denoise step at 1344x768 on a 20-SM
|
||||
// device, so a default 50-step render is hours, not seconds. Nothing here
|
||||
// imposes a deadline: GenerateVideo blocks for as long as the engine needs and
|
||||
// the gRPC call carries LocalAI's application context.
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"image"
|
||||
"math"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strconv"
|
||||
"strings"
|
||||
"unsafe"
|
||||
|
||||
// Registered for image.DecodeConfig only: a staged keyframe arrives as
|
||||
// whatever the caller uploaded, and we need its geometry to size the canvas.
|
||||
_ "image/gif"
|
||||
_ "image/jpeg"
|
||||
_ "image/png"
|
||||
|
||||
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
|
||||
"github.com/mudler/xlog"
|
||||
)
|
||||
|
||||
// vllm_video_model_params.device (vllm.h): no auto slot, unlike the text
|
||||
// engine's v14 device field.
|
||||
const (
|
||||
videoDeviceCPU int32 = 0
|
||||
videoDeviceCUDA int32 = 1
|
||||
)
|
||||
|
||||
// H3's shipped geometry. The canvas is truncated onto a 32-pixel grid and the
|
||||
// frame count onto the 17n+5 grid by the engine itself
|
||||
// (MiniMaxH3ResolveShape / MiniMaxH3AlignFrameCount in
|
||||
// src/vllm/model_executor/models/minimax_h3_planner.cpp); mirrored here only so
|
||||
// a keyframe can be resampled to the exact canvas the engine will render at.
|
||||
const (
|
||||
h3CanvasMultiple int32 = 32
|
||||
h3FrameGrid int32 = 17
|
||||
h3FrameOffset int32 = 5
|
||||
h3ShortEdge int32 = 768
|
||||
)
|
||||
|
||||
// videoPartitions are the two DECLARED partitions of the H3 release. The FL2VA
|
||||
// checkpoint serves t2va and fl2va; ref2va is a different checkpoint. Passing
|
||||
// reference conditioning against an fl2va DiT is a partition mismatch that
|
||||
// renders a coloured lattice over the frame rather than failing cleanly, which
|
||||
// is why it is refused here before the engine is ever called.
|
||||
const (
|
||||
partitionFL2VA = "fl2va"
|
||||
partitionRef2VA = "ref2va"
|
||||
)
|
||||
|
||||
// videoRequestParams are the per-request `params` keys this backend accepts.
|
||||
// Unknown keys are an error rather than a silent drop: a misspelled reference
|
||||
// path would otherwise produce a perfectly successful render of the wrong
|
||||
// thing, hours later.
|
||||
var videoRequestParams = []string{"noise_aug", "ref_image", "ref_video", "crf"}
|
||||
|
||||
// loadVideo opens the H3 checkpoint set. `dit` is the model config's
|
||||
// parameters.model; every other artifact comes from the options.
|
||||
func (v *VllmCpp) loadVideo(opts *pb.ModelOptions, dit string) error {
|
||||
vo := &v.opts.video
|
||||
|
||||
// Relative option paths resolve against LocalAI's models directory, which
|
||||
// is where the gallery lands the five H3 files.
|
||||
resolve := func(p string) string {
|
||||
if p == "" || filepath.IsAbs(p) || opts.ModelPath == "" {
|
||||
return p
|
||||
}
|
||||
return filepath.Join(opts.ModelPath, p)
|
||||
}
|
||||
vo.encoderPath = resolve(vo.encoderPath)
|
||||
vo.tokenizerPath = resolve(vo.tokenizerPath)
|
||||
vo.videoVaePath = resolve(vo.videoVaePath)
|
||||
vo.videoVaeConfig = resolve(vo.videoVaeConfig)
|
||||
vo.audioVaePath = resolve(vo.audioVaePath)
|
||||
vo.audioVaeConfig = resolve(vo.audioVaeConfig)
|
||||
vo.promptEmbedsPath = resolve(vo.promptEmbedsPath)
|
||||
vo.workdir = resolve(vo.workdir)
|
||||
|
||||
// A VAE config carries the per-channel latents_mean/latents_std and the
|
||||
// temporal clip_length/token_drop; decode is wrong without it. The release
|
||||
// ships it beside the weights, so default to that rather than making every
|
||||
// config repeat it.
|
||||
if vo.videoVaeConfig == "" && vo.videoVaePath != "" {
|
||||
vo.videoVaeConfig = siblingConfigJSON(vo.videoVaePath)
|
||||
}
|
||||
if vo.audioVaeConfig == "" && vo.audioVaePath != "" {
|
||||
vo.audioVaeConfig = siblingConfigJSON(vo.audioVaePath)
|
||||
}
|
||||
|
||||
if vo.partition == "" {
|
||||
// The community GGUF/NVFP4 quantisations strip the release metadata and
|
||||
// the two DiTs are byte-structurally identical, so the engine cannot
|
||||
// infer this and refuses every generate until it is declared. The
|
||||
// shipped FL2VA checkpoint is the one the gallery entry installs.
|
||||
vo.partition = partitionFL2VA
|
||||
xlog.Warn("[vllm-cpp] video partition not declared, assuming the FL2VA checkpoint",
|
||||
"hint", "set options: [video_partition:fl2va] or [video_partition:ref2va] to match the DiT you installed")
|
||||
}
|
||||
if vo.partition != partitionFL2VA && vo.partition != partitionRef2VA {
|
||||
return fmt.Errorf("vllm-cpp: video_partition must be %q or %q, got %q",
|
||||
partitionFL2VA, partitionRef2VA, vo.partition)
|
||||
}
|
||||
if vo.videoVaePath == "" || vo.audioVaePath == "" {
|
||||
return fmt.Errorf("vllm-cpp: MiniMax-H3 needs both VAEs: set options: " +
|
||||
"[video_vae:<video vae .safetensors>, audio_vae:<audio vae .safetensors>]")
|
||||
}
|
||||
if vo.encoderPath == "" && vo.promptEmbedsPath == "" {
|
||||
return fmt.Errorf("vllm-cpp: MiniMax-H3 needs text conditioning: set options: " +
|
||||
"[video_encoder:<encoder .gguf>, video_tokenizer:<tokenizer.json>] " +
|
||||
"or [video_prompt_embeds:<f32 embeddings>]")
|
||||
}
|
||||
if !vo.deviceSet && opts.GetCUDA() {
|
||||
vo.device = videoDeviceCUDA
|
||||
}
|
||||
|
||||
mp := cVideoModelParams{
|
||||
Device: vo.device,
|
||||
DequantBf16: vo.dequantBf16,
|
||||
Fp4Resident: vo.fp4Resident,
|
||||
}
|
||||
var keep [][]byte
|
||||
setStr := func(dst *uintptr, s string) {
|
||||
if s == "" {
|
||||
return
|
||||
}
|
||||
b := cString(s)
|
||||
keep = append(keep, b)
|
||||
*dst = uintptr(unsafe.Pointer(&b[0])) // #nosec G103 -- borrowed by C for the load call only
|
||||
}
|
||||
setStr(&mp.DitPath, dit)
|
||||
setStr(&mp.EncoderPath, vo.encoderPath)
|
||||
setStr(&mp.TokenizerPath, vo.tokenizerPath)
|
||||
setStr(&mp.VideoVaePath, vo.videoVaePath)
|
||||
setStr(&mp.VideoVaeConfigPath, vo.videoVaeConfig)
|
||||
setStr(&mp.AudioVaePath, vo.audioVaePath)
|
||||
setStr(&mp.AudioVaeConfigPath, vo.audioVaeConfig)
|
||||
setStr(&mp.PromptEmbedsPath, vo.promptEmbedsPath)
|
||||
setStr(&mp.Partition, vo.partition)
|
||||
|
||||
xlog.Info("[vllm-cpp] Load (MiniMax-H3 video)", "dit", dit, "engine", vllmVersion(),
|
||||
"encoder", vo.encoderPath, "tokenizer", vo.tokenizerPath,
|
||||
"videoVae", vo.videoVaePath, "audioVae", vo.audioVaePath,
|
||||
"partition", vo.partition, "device", videoDeviceName(vo.device),
|
||||
"dequantBf16", vo.dequantBf16 == 1, "fp4Resident", vo.fp4Resident == 1)
|
||||
|
||||
var engine uintptr
|
||||
rc := vllmVideoEngineLoad(unsafe.Pointer(&mp), unsafe.Pointer(&engine)) // #nosec G103 -- POD out-params
|
||||
runtime.KeepAlive(keep)
|
||||
if rc != vllmOK {
|
||||
return fmt.Errorf("vllm-cpp: video engine load failed: %s", vllmLastError())
|
||||
}
|
||||
v.videoEngine = engine
|
||||
return nil
|
||||
}
|
||||
|
||||
// GenerateVideo renders one clip and muxes it to opts.Dst as an MP4 carrying
|
||||
// H3's jointly generated AAC audio track. It blocks for the whole render.
|
||||
func (v *VllmCpp) GenerateVideo(opts *pb.GenerateVideoRequest) error {
|
||||
if v.videoEngine == 0 {
|
||||
return fmt.Errorf("vllm-cpp: this model is not a MiniMax-H3 video engine " +
|
||||
"(load it with the video_vae / audio_vae / video_encoder options)")
|
||||
}
|
||||
if strings.TrimSpace(opts.GetPrompt()) == "" {
|
||||
return fmt.Errorf("vllm-cpp: video generation needs a prompt")
|
||||
}
|
||||
dst := opts.GetDst()
|
||||
if dst == "" {
|
||||
return fmt.Errorf("vllm-cpp: video generation needs an output path")
|
||||
}
|
||||
vo := v.opts.video
|
||||
|
||||
extra, err := parseVideoRequestParams(opts.GetParams())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := checkPartitionConditioning(vo.partition, opts, extra); err != nil {
|
||||
return err
|
||||
}
|
||||
if opts.GetNegativePrompt() != "" {
|
||||
xlog.Warn("[vllm-cpp] MiniMax-H3 has no negative prompt; ignoring it")
|
||||
}
|
||||
if opts.GetCfgScale() != 0 {
|
||||
xlog.Warn("[vllm-cpp] MiniMax-H3 has no classifier-free guidance scale; ignoring cfg_scale")
|
||||
}
|
||||
|
||||
workdir, cleanup, err := v.videoWorkdir(dst)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer cleanup()
|
||||
|
||||
width, height := firstPositive(opts.GetWidth(), vo.width), firstPositive(opts.GetHeight(), vo.height)
|
||||
frames := firstPositive(opts.GetNumFrames(), vo.numFrames)
|
||||
steps := firstPositive(opts.GetStep(), vo.steps)
|
||||
|
||||
vp := cVideoParams{
|
||||
NumFrames: frames,
|
||||
Steps: steps,
|
||||
NoiseAug: extra.noiseAug,
|
||||
}
|
||||
if opts.GetSeed() > 0 {
|
||||
vp.Seed = uint64(opts.GetSeed())
|
||||
vp.HasSeed = 1
|
||||
}
|
||||
if aligned := alignFrameCount(frames); aligned != frames {
|
||||
xlog.Warn("[vllm-cpp] frame count is not on H3's 17n+5 grid; the engine rounds up",
|
||||
"requested", frames, "rendered", aligned)
|
||||
}
|
||||
|
||||
// Keyframes must be binary PPM (P6) at the exact output canvas: no image
|
||||
// codec and no resampler is vendored in libvllm. Resolve the canvas first,
|
||||
// then stage the frames through ffmpeg into it.
|
||||
//
|
||||
// The REQUEST's geometry is what is honoured here, not the model-level
|
||||
// default: that default is a t2va canvas, and applying it to a keyframe
|
||||
// would stretch a portrait photo into a 1344x768 letterbox. With no
|
||||
// requested geometry the canvas comes from the keyframe's own aspect, which
|
||||
// is the rule the engine itself applies (MiniMaxH3ResolveShape).
|
||||
first, last := opts.GetStartImage(), opts.GetEndImage()
|
||||
if first != "" || last != "" {
|
||||
width, height, err = resolveCanvas(opts.GetWidth(), opts.GetHeight(), first, last)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if first, err = stageKeyframe(vo.ffmpeg, first, width, height, workdir, "first"); err != nil {
|
||||
return err
|
||||
}
|
||||
if last, err = stageKeyframe(vo.ffmpeg, last, width, height, workdir, "last"); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
vp.Width, vp.Height = truncateToGrid(width), truncateToGrid(height)
|
||||
|
||||
var keep [][]byte
|
||||
setStr := func(dst *uintptr, s string) {
|
||||
if s == "" {
|
||||
return
|
||||
}
|
||||
b := cString(s)
|
||||
keep = append(keep, b)
|
||||
*dst = uintptr(unsafe.Pointer(&b[0])) // #nosec G103 -- borrowed by C for the call only
|
||||
}
|
||||
setStr(&vp.Prompt, opts.GetPrompt())
|
||||
setStr(&vp.OutputDir, workdir)
|
||||
setStr(&vp.FirstFrame, first)
|
||||
setStr(&vp.LastFrame, last)
|
||||
setStr(&vp.RefImage, extra.refImage)
|
||||
setStr(&vp.RefVideo, extra.refVideo)
|
||||
setStr(&vp.RefAudio, opts.GetAudio())
|
||||
|
||||
xlog.Info("[vllm-cpp] GenerateVideo", "dst", dst, "workdir", workdir,
|
||||
"width", vp.Width, "height", vp.Height, "frames", vp.NumFrames,
|
||||
"steps", vp.Steps, "seeded", vp.HasSeed == 1, "partition", vo.partition)
|
||||
|
||||
var out cVideoResult
|
||||
rc := vllmVideoGenerate(v.videoEngine, unsafe.Pointer(&vp), unsafe.Pointer(&out)) // #nosec G103 -- POD in/out params
|
||||
runtime.KeepAlive(keep)
|
||||
if rc != vllmOK {
|
||||
return fmt.Errorf("vllm-cpp: video generation failed: %s", vllmLastError())
|
||||
}
|
||||
defer vllmVideoResultFree(unsafe.Pointer(&out)) // #nosec G103 -- frees the library-owned members
|
||||
|
||||
frameDir, audioPath := goString(out.FrameDir), goString(out.AudioPath)
|
||||
xlog.Info("[vllm-cpp] rendered", "frames", out.FrameCount,
|
||||
"width", out.Width, "height", out.Height, "fps", out.Fps,
|
||||
"audio", audioPath, "sampleRate", out.SampleRate)
|
||||
if opts.GetFps() > 0 && opts.GetFps() != out.Fps {
|
||||
// Muxing at any other rate desynchronises the jointly generated audio.
|
||||
xlog.Warn("[vllm-cpp] MiniMax-H3 renders at a fixed frame rate; ignoring the requested fps",
|
||||
"requested", opts.GetFps(), "rendered", out.Fps)
|
||||
}
|
||||
|
||||
return v.muxVideo(frameDir, audioPath, dst, out.Fps, extra.crf)
|
||||
}
|
||||
|
||||
// muxVideo execs the argv libvllm composed. The encoding contract (h264 /
|
||||
// yuv420p + AAC, -shortest, +faststart) belongs to the library; only the spawn
|
||||
// is ours.
|
||||
func (v *VllmCpp) muxVideo(frameDir, audioPath, dst string, fps, crf int32) error {
|
||||
mx := cVideoMuxParams{Fps: fps, Crf: crf}
|
||||
var keep [][]byte
|
||||
setStr := func(dst *uintptr, s string) {
|
||||
if s == "" {
|
||||
return
|
||||
}
|
||||
b := cString(s)
|
||||
keep = append(keep, b)
|
||||
*dst = uintptr(unsafe.Pointer(&b[0])) // #nosec G103 -- borrowed by C for the call only
|
||||
}
|
||||
setStr(&mx.Frames, filepath.Join(frameDir, "frame_%06d.ppm"))
|
||||
setStr(&mx.AudioPath, audioPath)
|
||||
setStr(&mx.OutputPath, dst)
|
||||
|
||||
var argvPtr uintptr
|
||||
var argc int32
|
||||
rc := vllmVideoMuxArgv(unsafe.Pointer(&mx), unsafe.Pointer(&argvPtr), unsafe.Pointer(&argc)) // #nosec G103 -- POD out-params
|
||||
runtime.KeepAlive(keep)
|
||||
if rc != vllmOK {
|
||||
return fmt.Errorf("vllm-cpp: composing the mux command failed: %s", vllmLastError())
|
||||
}
|
||||
argv := goStringSlice(argvPtr, argc)
|
||||
vllmVideoMuxArgvFre(argvPtr, argc)
|
||||
if len(argv) == 0 {
|
||||
return fmt.Errorf("vllm-cpp: the library composed an empty mux command")
|
||||
}
|
||||
|
||||
ffmpegBin, err := resolveFfmpeg(v.opts.video.ffmpeg)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
argv[0] = ffmpegBin
|
||||
|
||||
xlog.Debug("[vllm-cpp] muxing", "argv", argv)
|
||||
output, err := exec.Command(argv[0], argv[1:]...).CombinedOutput() // #nosec G204 -- argv is composed by libvllm, argv[0] is a resolved binary
|
||||
if err != nil {
|
||||
return fmt.Errorf("vllm-cpp: ffmpeg mux failed: %w (output: %s)", err, strings.TrimSpace(string(output)))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// resolveFfmpeg locates the mux binary. The backend image is FROM scratch and
|
||||
// carries no ffmpeg, exactly like vibevoice-cpp's transcode path: the host must
|
||||
// provide one, and saying so plainly beats a bare "exec: not found" after an
|
||||
// hours-long render.
|
||||
func resolveFfmpeg(configured string) (string, error) {
|
||||
name := configured
|
||||
if name == "" {
|
||||
name = "ffmpeg"
|
||||
}
|
||||
path, err := exec.LookPath(name)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("vllm-cpp: %q not found: MiniMax-H3 output is muxed with ffmpeg, "+
|
||||
"install it on the host or point options: [ffmpeg:<path>] at a binary: %w", name, err)
|
||||
}
|
||||
return path, nil
|
||||
}
|
||||
|
||||
// videoWorkdir returns the directory the engine writes frame_%06d.ppm and
|
||||
// audio.wav into, plus its cleanup.
|
||||
//
|
||||
// It is ALWAYS a fresh directory. Reusing one would leave a longer previous
|
||||
// run's trailing frames in place for the mux to pick up, silently splicing two
|
||||
// renders together. With video_workdir set the run is kept (its frames are what
|
||||
// ref2va's ref_video consumes); otherwise it is removed once the mux succeeds.
|
||||
func (v *VllmCpp) videoWorkdir(dst string) (string, func(), error) {
|
||||
parent := v.opts.video.workdir
|
||||
keep := parent != ""
|
||||
if parent == "" {
|
||||
parent = filepath.Dir(dst)
|
||||
}
|
||||
if err := os.MkdirAll(parent, 0o750); err != nil {
|
||||
return "", nil, fmt.Errorf("vllm-cpp: creating the video work directory: %w", err)
|
||||
}
|
||||
dir, err := os.MkdirTemp(parent, "vllm-cpp-h3-")
|
||||
if err != nil {
|
||||
return "", nil, fmt.Errorf("vllm-cpp: creating the video work directory: %w", err)
|
||||
}
|
||||
if keep {
|
||||
return dir, func() {}, nil
|
||||
}
|
||||
return dir, func() {
|
||||
if err := os.RemoveAll(dir); err != nil {
|
||||
xlog.Warn("[vllm-cpp] could not remove the video work directory", "dir", dir, "error", err)
|
||||
}
|
||||
}, nil
|
||||
}
|
||||
|
||||
// videoExtraParams holds the per-request knobs that have no proto field.
|
||||
type videoExtraParams struct {
|
||||
noiseAug float32
|
||||
refImage string
|
||||
refVideo string
|
||||
crf int32
|
||||
}
|
||||
|
||||
func parseVideoRequestParams(params map[string]string) (videoExtraParams, error) {
|
||||
var extra videoExtraParams
|
||||
for k, raw := range params {
|
||||
v := strings.TrimSpace(raw)
|
||||
switch k {
|
||||
case "noise_aug":
|
||||
f, err := strconv.ParseFloat(v, 32)
|
||||
if err != nil {
|
||||
return extra, fmt.Errorf("vllm-cpp: params.noise_aug must be a number, got %q", raw)
|
||||
}
|
||||
extra.noiseAug = float32(f)
|
||||
case "ref_image":
|
||||
extra.refImage = v
|
||||
case "ref_video":
|
||||
extra.refVideo = v
|
||||
case "crf":
|
||||
n, err := strconv.ParseInt(v, 10, 32)
|
||||
if err != nil {
|
||||
return extra, fmt.Errorf("vllm-cpp: params.crf must be an integer, got %q", raw)
|
||||
}
|
||||
extra.crf = int32(n)
|
||||
default:
|
||||
return extra, fmt.Errorf("vllm-cpp: unknown params key %q (accepted: %s)",
|
||||
k, strings.Join(videoRequestParams, ", "))
|
||||
}
|
||||
}
|
||||
return extra, nil
|
||||
}
|
||||
|
||||
// checkPartitionConditioning refuses conditioning the loaded checkpoint cannot
|
||||
// serve.
|
||||
//
|
||||
// This is the failure this backend most needs to catch early. The FL2VA
|
||||
// partition serves t2va and fl2va; handing it a reference image or audio is a
|
||||
// partition mismatch, and H3 does not fail cleanly on one - it renders, for
|
||||
// hours, and returns a coloured lattice over the frame. The engine's own #77
|
||||
// guard covers a missing declaration; this covers a declaration that does not
|
||||
// match the request.
|
||||
func checkPartitionConditioning(partition string, opts *pb.GenerateVideoRequest, extra videoExtraParams) error {
|
||||
hasKeyframe := opts.GetStartImage() != "" || opts.GetEndImage() != ""
|
||||
hasReference := extra.refImage != "" || extra.refVideo != "" || opts.GetAudio() != ""
|
||||
|
||||
if hasKeyframe && hasReference {
|
||||
return fmt.Errorf("vllm-cpp: fl2va keyframes (start_image/end_image) and ref2va reference " +
|
||||
"conditioning (params.ref_image/params.ref_video/audio) are exclusive in the H3 pipeline")
|
||||
}
|
||||
switch partition {
|
||||
case partitionFL2VA:
|
||||
if hasReference {
|
||||
return fmt.Errorf("vllm-cpp: the FL2VA checkpoint serves t2va and fl2va only - " +
|
||||
"reference conditioning (params.ref_image/params.ref_video/audio) needs a ref2va DiT. " +
|
||||
"Use start_image for first-frame conditioning instead")
|
||||
}
|
||||
case partitionRef2VA:
|
||||
if hasKeyframe {
|
||||
return fmt.Errorf("vllm-cpp: the Ref2VA checkpoint does not serve fl2va keyframes - " +
|
||||
"pass the image as params.ref_image, or install the FL2VA checkpoint")
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// resolveCanvas settles the output geometry BEFORE a keyframe is resampled,
|
||||
// because the two have to agree exactly: the engine refuses a keyframe that is
|
||||
// not already at the output resolution, and when no geometry is requested it
|
||||
// derives one from the keyframe's own aspect. Mirrors _resolve_shape
|
||||
// (src/vllm/model_executor/models/minimax_h3_planner.cpp:264-308).
|
||||
func resolveCanvas(width, height int32, keyframes ...string) (int32, int32, error) {
|
||||
if width > 0 && height > 0 {
|
||||
return width, height, nil
|
||||
}
|
||||
for _, k := range keyframes {
|
||||
if k == "" {
|
||||
continue
|
||||
}
|
||||
w, h, err := imageDimensions(k)
|
||||
if err != nil {
|
||||
return 0, 0, err
|
||||
}
|
||||
if w <= 0 || h <= 0 {
|
||||
continue
|
||||
}
|
||||
// A 768 short edge, the long edge snapped onto the 32 grid.
|
||||
if w >= h {
|
||||
return alignMultiple(float64(h3ShortEdge)*float64(w)/float64(h), h3CanvasMultiple), h3ShortEdge, nil
|
||||
}
|
||||
return h3ShortEdge, alignMultiple(float64(h3ShortEdge)*float64(h)/float64(w), h3CanvasMultiple), nil
|
||||
}
|
||||
// The shipped canvas.
|
||||
return 1344, h3ShortEdge, nil
|
||||
}
|
||||
|
||||
// stageKeyframe converts a staged upload into the binary PPM (P6) at exactly
|
||||
// width x height that the engine requires. libvllm vendors no image codec and
|
||||
// no resampler, so ffmpeg does both; a P6 already at the canvas passes through
|
||||
// untouched.
|
||||
func stageKeyframe(ffmpegPath, src string, width, height int32, workdir, name string) (string, error) {
|
||||
if src == "" {
|
||||
return "", nil
|
||||
}
|
||||
if w, h, err := ppmDimensions(src); err == nil && w == width && h == height {
|
||||
return src, nil
|
||||
}
|
||||
ffmpegBin, err := resolveFfmpeg(ffmpegPath)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("converting the %s keyframe to PPM: %w", name, err)
|
||||
}
|
||||
out := filepath.Join(workdir, name+"_frame.ppm")
|
||||
// -frames:v 1 because an animated upload (GIF) would otherwise write a
|
||||
// sequence; -pix_fmt rgb24 is what the image2/ppm muxer needs for P6.
|
||||
cmd := exec.Command(ffmpegBin, "-y", "-loglevel", "error", "-i", src, // #nosec G204 -- the binary is resolved, the rest are literals and staged paths
|
||||
"-frames:v", "1",
|
||||
"-vf", fmt.Sprintf("scale=%d:%d", width, height),
|
||||
"-pix_fmt", "rgb24", "-f", "image2", out)
|
||||
if output, err := cmd.CombinedOutput(); err != nil {
|
||||
return "", fmt.Errorf("vllm-cpp: converting the %s keyframe to PPM failed: %w (output: %s)",
|
||||
name, err, strings.TrimSpace(string(output)))
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// imageDimensions reads geometry from a staged upload, PPM included (the Go
|
||||
// standard library has no netpbm decoder).
|
||||
func imageDimensions(path string) (int32, int32, error) {
|
||||
if w, h, err := ppmDimensions(path); err == nil {
|
||||
return w, h, nil
|
||||
}
|
||||
f, err := os.Open(path) // #nosec G304 -- a path staged by LocalAI for this request
|
||||
if err != nil {
|
||||
return 0, 0, fmt.Errorf("vllm-cpp: reading the keyframe %q: %w", path, err)
|
||||
}
|
||||
defer func() { _ = f.Close() }()
|
||||
cfg, _, err := image.DecodeConfig(f)
|
||||
if err != nil {
|
||||
return 0, 0, fmt.Errorf("vllm-cpp: the keyframe %q is not a PNG, JPEG, GIF or binary PPM: %w", path, err)
|
||||
}
|
||||
return int32(cfg.Width), int32(cfg.Height), nil
|
||||
}
|
||||
|
||||
// ppmDimensions parses a binary PPM (P6) header: magic, then width, height and
|
||||
// maxval as ASCII decimals separated by whitespace, with # comments allowed.
|
||||
func ppmDimensions(path string) (int32, int32, error) {
|
||||
f, err := os.Open(path) // #nosec G304 -- a path staged by LocalAI for this request
|
||||
if err != nil {
|
||||
return 0, 0, err
|
||||
}
|
||||
defer func() { _ = f.Close() }()
|
||||
|
||||
// A P6 header is a handful of bytes; 512 covers any sane comment run.
|
||||
buf := make([]byte, 512)
|
||||
n, err := f.Read(buf)
|
||||
if n < 2 || (err != nil && n == 0) {
|
||||
return 0, 0, fmt.Errorf("not a PPM")
|
||||
}
|
||||
if buf[0] != 'P' || buf[1] != '6' {
|
||||
return 0, 0, fmt.Errorf("not a binary PPM (P6)")
|
||||
}
|
||||
fields := make([]int32, 0, 2)
|
||||
for i := 2; i < n && len(fields) < 2; {
|
||||
switch {
|
||||
case buf[i] == '#':
|
||||
for i < n && buf[i] != '\n' {
|
||||
i++
|
||||
}
|
||||
case buf[i] >= '0' && buf[i] <= '9':
|
||||
value := int32(0)
|
||||
for i < n && buf[i] >= '0' && buf[i] <= '9' {
|
||||
value = value*10 + int32(buf[i]-'0')
|
||||
i++
|
||||
}
|
||||
fields = append(fields, value)
|
||||
default:
|
||||
i++
|
||||
}
|
||||
}
|
||||
if len(fields) < 2 {
|
||||
return 0, 0, fmt.Errorf("truncated PPM header")
|
||||
}
|
||||
return fields[0], fields[1], nil
|
||||
}
|
||||
|
||||
// alignMultiple mirrors MiniMaxH3AlignMultiple: round-half-to-even onto the
|
||||
// multiple, floored at one multiple. Half-to-even, not half-away-from-zero,
|
||||
// because the reference pipeline uses Python's round().
|
||||
func alignMultiple(value float64, multiple int32) int32 {
|
||||
snapped := int32(math.RoundToEven(value/float64(multiple))) * multiple
|
||||
if snapped < multiple {
|
||||
return multiple
|
||||
}
|
||||
return snapped
|
||||
}
|
||||
|
||||
// truncateToGrid mirrors the engine's canvas snap: truncation, not rounding.
|
||||
func truncateToGrid(v int32) int32 {
|
||||
if v <= 0 {
|
||||
return 0
|
||||
}
|
||||
return v / h3CanvasMultiple * h3CanvasMultiple
|
||||
}
|
||||
|
||||
// alignFrameCount mirrors MiniMaxH3AlignFrameCount: the next value on the
|
||||
// 17n+5 grid. Used only to warn - the engine does the real alignment.
|
||||
func alignFrameCount(frames int32) int32 {
|
||||
if frames <= 0 {
|
||||
return frames
|
||||
}
|
||||
for frames%h3FrameGrid != h3FrameOffset {
|
||||
frames++
|
||||
}
|
||||
return frames
|
||||
}
|
||||
|
||||
func firstPositive(values ...int32) int32 {
|
||||
for _, v := range values {
|
||||
if v > 0 {
|
||||
return v
|
||||
}
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func videoDeviceName(device int32) string {
|
||||
if device == videoDeviceCUDA {
|
||||
return "cuda"
|
||||
}
|
||||
return "cpu"
|
||||
}
|
||||
|
||||
// siblingConfigJSON is the release layout: each VAE ships its config.json in
|
||||
// the directory holding its weights.
|
||||
func siblingConfigJSON(weights string) string {
|
||||
candidate := filepath.Join(filepath.Dir(weights), "config.json")
|
||||
if _, err := os.Stat(candidate); err != nil {
|
||||
return ""
|
||||
}
|
||||
return candidate
|
||||
}
|
||||
@@ -0,0 +1,305 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"unsafe"
|
||||
|
||||
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
// The video PODs carry the same contract as the text ones in vllmcpp_test.go:
|
||||
// these are the C offsets of vllm.h on LP64, and a drift here is silent memory
|
||||
// corruption rather than a compile error.
|
||||
var _ = Describe("C ABI video struct mirrors", func() {
|
||||
It("cVideoModelParams matches vllm_video_model_params", func() {
|
||||
var p cVideoModelParams
|
||||
Expect(unsafe.Offsetof(p.DitPath)).To(Equal(uintptr(0)))
|
||||
Expect(unsafe.Offsetof(p.EncoderPath)).To(Equal(uintptr(8)))
|
||||
Expect(unsafe.Offsetof(p.TokenizerPath)).To(Equal(uintptr(16)))
|
||||
Expect(unsafe.Offsetof(p.VideoVaePath)).To(Equal(uintptr(24)))
|
||||
Expect(unsafe.Offsetof(p.VideoVaeConfigPath)).To(Equal(uintptr(32)))
|
||||
Expect(unsafe.Offsetof(p.AudioVaePath)).To(Equal(uintptr(40)))
|
||||
Expect(unsafe.Offsetof(p.AudioVaeConfigPath)).To(Equal(uintptr(48)))
|
||||
Expect(unsafe.Offsetof(p.PromptEmbedsPath)).To(Equal(uintptr(56)))
|
||||
Expect(unsafe.Offsetof(p.Partition)).To(Equal(uintptr(64)))
|
||||
Expect(unsafe.Offsetof(p.Device)).To(Equal(uintptr(72)))
|
||||
Expect(unsafe.Offsetof(p.DequantBf16)).To(Equal(uintptr(76)))
|
||||
Expect(unsafe.Offsetof(p.Fp4Resident)).To(Equal(uintptr(80)))
|
||||
Expect(unsafe.Offsetof(p.Family)).To(Equal(uintptr(88)))
|
||||
Expect(unsafe.Offsetof(p.ExtraKeys)).To(Equal(uintptr(96)))
|
||||
Expect(unsafe.Offsetof(p.ExtraValues)).To(Equal(uintptr(104)))
|
||||
Expect(unsafe.Offsetof(p.NExtras)).To(Equal(uintptr(112)))
|
||||
Expect(unsafe.Sizeof(p)).To(Equal(uintptr(120)))
|
||||
})
|
||||
|
||||
It("cVideoParams matches vllm_video_params", func() {
|
||||
var p cVideoParams
|
||||
Expect(unsafe.Offsetof(p.Prompt)).To(Equal(uintptr(0)))
|
||||
Expect(unsafe.Offsetof(p.Width)).To(Equal(uintptr(8)))
|
||||
Expect(unsafe.Offsetof(p.Height)).To(Equal(uintptr(12)))
|
||||
Expect(unsafe.Offsetof(p.NumFrames)).To(Equal(uintptr(16)))
|
||||
Expect(unsafe.Offsetof(p.Steps)).To(Equal(uintptr(20)))
|
||||
Expect(unsafe.Offsetof(p.Seed)).To(Equal(uintptr(24)))
|
||||
Expect(unsafe.Offsetof(p.HasSeed)).To(Equal(uintptr(32)))
|
||||
Expect(unsafe.Offsetof(p.FirstFrame)).To(Equal(uintptr(40)))
|
||||
Expect(unsafe.Offsetof(p.LastFrame)).To(Equal(uintptr(48)))
|
||||
Expect(unsafe.Offsetof(p.RefImage)).To(Equal(uintptr(56)))
|
||||
Expect(unsafe.Offsetof(p.RefVideo)).To(Equal(uintptr(64)))
|
||||
Expect(unsafe.Offsetof(p.RefAudio)).To(Equal(uintptr(72)))
|
||||
Expect(unsafe.Offsetof(p.NoiseAug)).To(Equal(uintptr(80)))
|
||||
Expect(unsafe.Offsetof(p.OutputDir)).To(Equal(uintptr(88)))
|
||||
Expect(unsafe.Offsetof(p.ExtraKeys)).To(Equal(uintptr(96)))
|
||||
Expect(unsafe.Offsetof(p.ExtraValues)).To(Equal(uintptr(104)))
|
||||
Expect(unsafe.Offsetof(p.NExtras)).To(Equal(uintptr(112)))
|
||||
Expect(unsafe.Sizeof(p)).To(Equal(uintptr(120)))
|
||||
})
|
||||
|
||||
It("cVideoResult matches vllm_video_result", func() {
|
||||
var r cVideoResult
|
||||
Expect(unsafe.Offsetof(r.FrameDir)).To(Equal(uintptr(0)))
|
||||
Expect(unsafe.Offsetof(r.AudioPath)).To(Equal(uintptr(8)))
|
||||
Expect(unsafe.Offsetof(r.FrameCount)).To(Equal(uintptr(16)))
|
||||
Expect(unsafe.Offsetof(r.Width)).To(Equal(uintptr(20)))
|
||||
Expect(unsafe.Offsetof(r.Height)).To(Equal(uintptr(24)))
|
||||
Expect(unsafe.Offsetof(r.Fps)).To(Equal(uintptr(28)))
|
||||
Expect(unsafe.Offsetof(r.SampleRate)).To(Equal(uintptr(32)))
|
||||
Expect(unsafe.Offsetof(r.MuxArgv)).To(Equal(uintptr(40)))
|
||||
Expect(unsafe.Offsetof(r.MuxArgc)).To(Equal(uintptr(48)))
|
||||
Expect(unsafe.Sizeof(r)).To(Equal(uintptr(56)))
|
||||
})
|
||||
|
||||
It("cVideoMuxParams matches vllm_video_mux_params", func() {
|
||||
var p cVideoMuxParams
|
||||
Expect(unsafe.Offsetof(p.Frames)).To(Equal(uintptr(0)))
|
||||
Expect(unsafe.Offsetof(p.AudioPath)).To(Equal(uintptr(8)))
|
||||
Expect(unsafe.Offsetof(p.OutputPath)).To(Equal(uintptr(16)))
|
||||
Expect(unsafe.Offsetof(p.Fps)).To(Equal(uintptr(24)))
|
||||
Expect(unsafe.Offsetof(p.Crf)).To(Equal(uintptr(28)))
|
||||
Expect(unsafe.Sizeof(p)).To(Equal(uintptr(32)))
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("video load options", func() {
|
||||
It("stays disengaged for a plain text config", func() {
|
||||
lo := parseOptions(&pb.ModelOptions{Options: []string{"max_num_seqs:16"}})
|
||||
Expect(lo.video.engaged()).To(BeFalse())
|
||||
})
|
||||
|
||||
It("reads the H3 checkpoint set from the options list", func() {
|
||||
lo := parseOptions(&pb.ModelOptions{Options: []string{
|
||||
"video_encoder:qwen3vl-32B-MiniMax-H3-Q4_K_M.gguf",
|
||||
"video_tokenizer:tokenizer.json",
|
||||
"video_vae:vae/diffusion_pytorch_model.safetensors",
|
||||
"audio_vae:audio_vae/model.safetensors",
|
||||
"video_partition:fl2va",
|
||||
"video_device:cuda",
|
||||
"video_dequant_bf16:true",
|
||||
"video_width:1344",
|
||||
"video_height:768",
|
||||
"video_num_frames:124",
|
||||
"video_steps:50",
|
||||
}})
|
||||
Expect(lo.video.engaged()).To(BeTrue())
|
||||
Expect(lo.video.encoderPath).To(Equal("qwen3vl-32B-MiniMax-H3-Q4_K_M.gguf"))
|
||||
Expect(lo.video.tokenizerPath).To(Equal("tokenizer.json"))
|
||||
Expect(lo.video.videoVaePath).To(Equal("vae/diffusion_pytorch_model.safetensors"))
|
||||
Expect(lo.video.audioVaePath).To(Equal("audio_vae/model.safetensors"))
|
||||
Expect(lo.video.partition).To(Equal(partitionFL2VA))
|
||||
Expect(lo.video.device).To(Equal(videoDeviceCUDA))
|
||||
Expect(lo.video.deviceSet).To(BeTrue())
|
||||
Expect(lo.video.dequantBf16).To(Equal(int32(1)))
|
||||
Expect(lo.video.width).To(Equal(int32(1344)))
|
||||
Expect(lo.video.height).To(Equal(int32(768)))
|
||||
Expect(lo.video.numFrames).To(Equal(int32(124)))
|
||||
Expect(lo.video.steps).To(Equal(int32(50)))
|
||||
})
|
||||
|
||||
It("reads the same keys from engine_args", func() {
|
||||
lo := parseOptions(&pb.ModelOptions{
|
||||
EngineArgs: `{"video_vae":"vae/v.safetensors","audio_vae":"a.safetensors","video_num_frames":124,"video_dequant_bf16":true}`,
|
||||
})
|
||||
Expect(lo.video.engaged()).To(BeTrue())
|
||||
Expect(lo.video.videoVaePath).To(Equal("vae/v.safetensors"))
|
||||
Expect(lo.video.audioVaePath).To(Equal("a.safetensors"))
|
||||
Expect(lo.video.numFrames).To(Equal(int32(124)))
|
||||
Expect(lo.video.dequantBf16).To(Equal(int32(1)))
|
||||
})
|
||||
|
||||
It("ignores an unknown video_device rather than guessing", func() {
|
||||
lo := parseOptions(&pb.ModelOptions{Options: []string{"video_vae:v", "video_device:tpu"}})
|
||||
Expect(lo.video.deviceSet).To(BeFalse())
|
||||
Expect(lo.video.device).To(Equal(videoDeviceCPU))
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("per-request params", func() {
|
||||
It("maps the accepted keys", func() {
|
||||
extra, err := parseVideoRequestParams(map[string]string{
|
||||
"noise_aug": "0.5", "ref_image": "/tmp/ref.ppm", "crf": "20",
|
||||
})
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(extra.noiseAug).To(BeNumerically("~", 0.5, 1e-6))
|
||||
Expect(extra.refImage).To(Equal("/tmp/ref.ppm"))
|
||||
Expect(extra.crf).To(Equal(int32(20)))
|
||||
})
|
||||
|
||||
It("refuses an unknown key instead of dropping it", func() {
|
||||
_, err := parseVideoRequestParams(map[string]string{"resolution": "480p"})
|
||||
Expect(err).To(MatchError(ContainSubstring("unknown params key")))
|
||||
})
|
||||
|
||||
It("refuses a non-numeric noise_aug", func() {
|
||||
_, err := parseVideoRequestParams(map[string]string{"noise_aug": "high"})
|
||||
Expect(err).To(HaveOccurred())
|
||||
})
|
||||
})
|
||||
|
||||
// The partition guard is the correctness rule this backend exists to enforce:
|
||||
// the FL2VA DiT serves t2va and fl2va, and handing it reference conditioning
|
||||
// renders a broken lattice over the frame after a multi-hour generation rather
|
||||
// than failing.
|
||||
var _ = Describe("partition conditioning guard", func() {
|
||||
It("accepts a plain t2va request on fl2va", func() {
|
||||
Expect(checkPartitionConditioning(partitionFL2VA,
|
||||
&pb.GenerateVideoRequest{Prompt: "a llama"}, videoExtraParams{})).To(Succeed())
|
||||
})
|
||||
|
||||
It("accepts fl2va keyframes on fl2va", func() {
|
||||
Expect(checkPartitionConditioning(partitionFL2VA,
|
||||
&pb.GenerateVideoRequest{StartImage: "/tmp/a.png"}, videoExtraParams{})).To(Succeed())
|
||||
})
|
||||
|
||||
It("refuses a reference image on fl2va", func() {
|
||||
err := checkPartitionConditioning(partitionFL2VA,
|
||||
&pb.GenerateVideoRequest{}, videoExtraParams{refImage: "/tmp/ref.ppm"})
|
||||
Expect(err).To(MatchError(ContainSubstring("ref2va")))
|
||||
})
|
||||
|
||||
It("refuses reference audio on fl2va", func() {
|
||||
err := checkPartitionConditioning(partitionFL2VA,
|
||||
&pb.GenerateVideoRequest{Audio: "/tmp/voice.wav"}, videoExtraParams{})
|
||||
Expect(err).To(HaveOccurred())
|
||||
})
|
||||
|
||||
It("refuses fl2va keyframes on ref2va", func() {
|
||||
err := checkPartitionConditioning(partitionRef2VA,
|
||||
&pb.GenerateVideoRequest{StartImage: "/tmp/a.png"}, videoExtraParams{})
|
||||
Expect(err).To(HaveOccurred())
|
||||
})
|
||||
|
||||
It("refuses keyframes and references together on either partition", func() {
|
||||
err := checkPartitionConditioning(partitionRef2VA,
|
||||
&pb.GenerateVideoRequest{StartImage: "/tmp/a.png"}, videoExtraParams{refVideo: "/tmp/clip"})
|
||||
Expect(err).To(MatchError(ContainSubstring("exclusive")))
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("H3 geometry", func() {
|
||||
It("keeps an explicitly requested canvas", func() {
|
||||
w, h, err := resolveCanvas(1280, 720)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(w).To(Equal(int32(1280)))
|
||||
Expect(h).To(Equal(int32(720)))
|
||||
})
|
||||
|
||||
It("falls back to the shipped 1344x768 canvas", func() {
|
||||
w, h, err := resolveCanvas(0, 0)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(w).To(Equal(int32(1344)))
|
||||
Expect(h).To(Equal(int32(768)))
|
||||
})
|
||||
|
||||
It("derives a landscape canvas from a keyframe's aspect", func() {
|
||||
path := writePPM(1920, 1080)
|
||||
w, h, err := resolveCanvas(0, 0, path)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(h).To(Equal(int32(768)))
|
||||
// 768 * 16/9 = 1365.33; /32 = 42.67, round-half-to-even to 43, x32.
|
||||
Expect(w).To(Equal(int32(1376)))
|
||||
})
|
||||
|
||||
It("derives a portrait canvas from a keyframe's aspect", func() {
|
||||
path := writePPM(1080, 1920)
|
||||
w, h, err := resolveCanvas(0, 0, path)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(w).To(Equal(int32(768)))
|
||||
Expect(h).To(Equal(int32(1376)))
|
||||
})
|
||||
|
||||
It("truncates onto the 32 grid the way the engine does", func() {
|
||||
Expect(truncateToGrid(1000)).To(Equal(int32(992)))
|
||||
Expect(truncateToGrid(768)).To(Equal(int32(768)))
|
||||
})
|
||||
|
||||
It("reports the 17n+5 frame grid", func() {
|
||||
Expect(alignFrameCount(124)).To(Equal(int32(124)))
|
||||
Expect(alignFrameCount(120)).To(Equal(int32(124)))
|
||||
Expect(alignFrameCount(100)).To(Equal(int32(107)))
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("keyframe staging", func() {
|
||||
It("parses a binary PPM header, comments included", func() {
|
||||
dir := GinkgoT().TempDir()
|
||||
path := filepath.Join(dir, "commented.ppm")
|
||||
Expect(os.WriteFile(path, []byte("P6\n# made by a test\n64 32\n255\n"), 0o600)).To(Succeed())
|
||||
w, h, err := ppmDimensions(path)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(w).To(Equal(int32(64)))
|
||||
Expect(h).To(Equal(int32(32)))
|
||||
})
|
||||
|
||||
It("refuses an ASCII PPM (P3): the engine reads P6 only", func() {
|
||||
dir := GinkgoT().TempDir()
|
||||
path := filepath.Join(dir, "ascii.ppm")
|
||||
Expect(os.WriteFile(path, []byte("P3\n64 32\n255\n"), 0o600)).To(Succeed())
|
||||
_, _, err := ppmDimensions(path)
|
||||
Expect(err).To(HaveOccurred())
|
||||
})
|
||||
|
||||
It("passes a P6 already at the canvas straight through, without ffmpeg", func() {
|
||||
path := writePPM(64, 32)
|
||||
out, err := stageKeyframe("", path, 64, 32, GinkgoT().TempDir(), "first")
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(out).To(Equal(path))
|
||||
})
|
||||
|
||||
It("is a no-op for an absent keyframe", func() {
|
||||
out, err := stageKeyframe("", "", 64, 32, GinkgoT().TempDir(), "first")
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(out).To(BeEmpty())
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("GenerateVideo preconditions", func() {
|
||||
It("refuses when the model is not a video engine", func() {
|
||||
v := &VllmCpp{}
|
||||
Expect(v.GenerateVideo(&pb.GenerateVideoRequest{Prompt: "x", Dst: "/tmp/o.mp4"})).
|
||||
To(MatchError(ContainSubstring("not a MiniMax-H3 video engine")))
|
||||
})
|
||||
})
|
||||
|
||||
// writePPM writes a valid P6 header of the given geometry. Only the header is
|
||||
// read by anything under test, so the pixel payload is left off.
|
||||
func writePPM(width, height int) string {
|
||||
dir := GinkgoT().TempDir()
|
||||
path := filepath.Join(dir, "frame.ppm")
|
||||
header := []byte("P6\n" + itoa(width) + " " + itoa(height) + "\n255\n")
|
||||
Expect(os.WriteFile(path, header, 0o600)).To(Succeed())
|
||||
return path
|
||||
}
|
||||
|
||||
func itoa(v int) string {
|
||||
if v == 0 {
|
||||
return "0"
|
||||
}
|
||||
digits := ""
|
||||
for v > 0 {
|
||||
digits = string(rune('0'+v%10)) + digits
|
||||
v /= 10
|
||||
}
|
||||
return digits
|
||||
}
|
||||
@@ -16,10 +16,17 @@ func TestVllmCpp(t *testing.T) {
|
||||
RunSpecs(t, "vllm-cpp suite")
|
||||
}
|
||||
|
||||
// The Go POD mirrors must match the C struct layout of vllm.h (ABI v2)
|
||||
// The Go POD mirrors must match the C struct layout of vllm.h (ABI v21)
|
||||
// byte-for-byte: these offsets are the C offsets on LP64 (linux/darwin
|
||||
// amd64+arm64). A failure here means govllmcpp.go drifted from vllm.h.
|
||||
var _ = Describe("C ABI struct mirrors", func() {
|
||||
It("declares the ABI version the pinned engine reports", func() {
|
||||
// VLLM_ABI_VERSION in the vllm.h of VLLM_CPP_VERSION (Makefile).
|
||||
// Moving the pin past this without growing the mirrors below ships a
|
||||
// backend that refuses every load at startup (issue #11379).
|
||||
Expect(abiVersion).To(Equal(21))
|
||||
})
|
||||
|
||||
It("cModelParams matches vllm_model_params", func() {
|
||||
var p cModelParams
|
||||
Expect(unsafe.Offsetof(p.ModelPath)).To(Equal(uintptr(0)))
|
||||
@@ -30,10 +37,24 @@ var _ = Describe("C ABI struct mirrors", func() {
|
||||
Expect(unsafe.Offsetof(p.MaxNumSeqs)).To(Equal(uintptr(28)))
|
||||
Expect(unsafe.Offsetof(p.ToolParser)).To(Equal(uintptr(32)))
|
||||
Expect(unsafe.Offsetof(p.ReasoningParser)).To(Equal(uintptr(40)))
|
||||
Expect(unsafe.Sizeof(p)).To(Equal(uintptr(48)))
|
||||
Expect(unsafe.Offsetof(p.SpeculativeConfig)).To(Equal(uintptr(48)))
|
||||
Expect(unsafe.Offsetof(p.EnablePrefixCaching)).To(Equal(uintptr(56)))
|
||||
Expect(unsafe.Offsetof(p.MaxNumBatchedTokens)).To(Equal(uintptr(60)))
|
||||
Expect(unsafe.Offsetof(p.SchedulingPolicy)).To(Equal(uintptr(64)))
|
||||
Expect(unsafe.Offsetof(p.KVTransferConfig)).To(Equal(uintptr(72)))
|
||||
Expect(unsafe.Offsetof(p.OffloadConfig)).To(Equal(uintptr(80)))
|
||||
Expect(unsafe.Offsetof(p.EnableJumpForward)).To(Equal(uintptr(88)))
|
||||
Expect(unsafe.Offsetof(p.Device)).To(Equal(uintptr(92)))
|
||||
// 96: gpu_memory_utilization is a double, so it takes the next
|
||||
// 8-aligned slot after the int32 pair. Go pads identically.
|
||||
Expect(unsafe.Offsetof(p.GPUMemoryUtil)).To(Equal(uintptr(96)))
|
||||
Expect(unsafe.Offsetof(p.KVCacheMemoryBytes)).To(Equal(uintptr(104)))
|
||||
Expect(unsafe.Offsetof(p.LanguageModelOnly)).To(Equal(uintptr(112)))
|
||||
Expect(unsafe.Offsetof(p.LimitMMPerPrompt)).To(Equal(uintptr(120)))
|
||||
Expect(unsafe.Sizeof(p)).To(Equal(uintptr(128)))
|
||||
})
|
||||
|
||||
It("cSamplingParams matches vllm_sampling_params (ABI v2)", func() {
|
||||
It("cSamplingParams matches vllm_sampling_params (ABI v8)", func() {
|
||||
var p cSamplingParams
|
||||
Expect(unsafe.Offsetof(p.Temperature)).To(Equal(uintptr(0)))
|
||||
Expect(unsafe.Offsetof(p.TopP)).To(Equal(uintptr(4)))
|
||||
@@ -55,7 +76,9 @@ var _ = Describe("C ABI struct mirrors", func() {
|
||||
Expect(unsafe.Offsetof(p.NStructuredChoice)).To(Equal(uintptr(96)))
|
||||
Expect(unsafe.Offsetof(p.StructuredGrammar)).To(Equal(uintptr(104)))
|
||||
Expect(unsafe.Offsetof(p.StructuredJSONObject)).To(Equal(uintptr(112)))
|
||||
Expect(unsafe.Sizeof(p)).To(Equal(uintptr(120)))
|
||||
Expect(unsafe.Offsetof(p.LogitsProcessor)).To(Equal(uintptr(120)))
|
||||
Expect(unsafe.Offsetof(p.LogitsProcessorUserData)).To(Equal(uintptr(128)))
|
||||
Expect(unsafe.Sizeof(p)).To(Equal(uintptr(136)))
|
||||
})
|
||||
|
||||
It("cCompletion matches vllm_completion", func() {
|
||||
@@ -68,6 +91,23 @@ var _ = Describe("C ABI struct mirrors", func() {
|
||||
})
|
||||
})
|
||||
|
||||
// Pin/mirror skew is the failure mode this backend is most exposed to: the Go
|
||||
// PODs above are hand-written against one VLLM_ABI_VERSION, and the Makefile
|
||||
// pins the vllm.cpp commit that produces it. This spec catches drift without
|
||||
// needing model weights - set VLLM_CPP_LIBRARY to a built libvllm and it binds
|
||||
// every symbol and compares the library's reported ABI against the mirrors'.
|
||||
var _ = Describe("real library ABI handshake", func() {
|
||||
It("binds every symbol and reports the ABI the mirrors were written against", func() {
|
||||
lib := os.Getenv("VLLM_CPP_LIBRARY")
|
||||
if lib == "" {
|
||||
Skip("VLLM_CPP_LIBRARY not set; skipping the real-library handshake")
|
||||
}
|
||||
Expect(registerLib(lib)).To(Succeed())
|
||||
Expect(vllmABIVersion()).To(Equal(int32(abiVersion)))
|
||||
Expect(vllmVersion()).NotTo(BeEmpty())
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("parseOptions", func() {
|
||||
It("extracts the engine sizing knobs", func() {
|
||||
lo := parseOptions(&pb.ModelOptions{Options: []string{
|
||||
@@ -83,6 +123,129 @@ var _ = Describe("parseOptions", func() {
|
||||
}})
|
||||
Expect(lo).To(Equal(loadOptions{}))
|
||||
})
|
||||
|
||||
It("carries a speculative_config JSON value through the legacy options list", func() {
|
||||
// strings.Cut splits on the FIRST colon only, so a JSON object value
|
||||
// survives the "key:value" spelling intact.
|
||||
lo := parseOptions(&pb.ModelOptions{Options: []string{
|
||||
`speculative_config:{"method":"mtp","num_speculative_tokens":1}`,
|
||||
}})
|
||||
Expect(lo.speculativeConfig).To(Equal(`{"method":"mtp","num_speculative_tokens":1}`))
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("engine_args", func() {
|
||||
It("maps every load knob onto the C model params", func() {
|
||||
lo := parseOptions(&pb.ModelOptions{EngineArgs: `{
|
||||
"block_size": 64,
|
||||
"num_blocks": 1024,
|
||||
"max_model_len": 16384,
|
||||
"max_num_seqs": 32,
|
||||
"max_num_batched_tokens": 8192,
|
||||
"enable_prefix_caching": true,
|
||||
"scheduling_policy": "lpm",
|
||||
"tool_parser": "qwen3",
|
||||
"reasoning_parser": "deepseek_r1",
|
||||
"tokenizer_config": "/models/tok/tokenizer_config.json"
|
||||
}`})
|
||||
Expect(lo.blockSize).To(Equal(int32(64)))
|
||||
Expect(lo.numBlocks).To(Equal(int32(1024)))
|
||||
Expect(lo.maxModelLen).To(Equal(int32(16384)))
|
||||
Expect(lo.maxNumSeqs).To(Equal(int32(32)))
|
||||
Expect(lo.maxNumBatchedTokens).To(Equal(int32(8192)))
|
||||
Expect(lo.enablePrefixCaching).To(Equal(int32(1)))
|
||||
Expect(lo.schedulingPolicy).To(Equal("lpm"))
|
||||
Expect(lo.toolParser).To(Equal("qwen3"))
|
||||
Expect(lo.reasoningParser).To(Equal("deepseek_r1"))
|
||||
Expect(lo.tokenizerConfigPath).To(Equal("/models/tok/tokenizer_config.json"))
|
||||
})
|
||||
|
||||
It("re-marshals a nested speculative_config object to JSON for the engine", func() {
|
||||
lo := parseOptions(&pb.ModelOptions{EngineArgs: `{
|
||||
"speculative_config": {"method": "mtp", "num_speculative_tokens": 1}
|
||||
}`})
|
||||
Expect(lo.speculativeConfig).To(MatchJSON(`{"method":"mtp","num_speculative_tokens":1}`))
|
||||
})
|
||||
|
||||
It("re-marshals a nested kv_transfer_config object (LMCache) to JSON", func() {
|
||||
lo := parseOptions(&pb.ModelOptions{EngineArgs: `{
|
||||
"kv_transfer_config": {
|
||||
"kv_connector": "LMCacheConnector",
|
||||
"kv_role": "kv_both",
|
||||
"kv_connector_extra_config": {"host": "127.0.0.1", "port": 65432}
|
||||
}
|
||||
}`})
|
||||
Expect(lo.kvTransferConfig).To(MatchJSON(`{
|
||||
"kv_connector":"LMCacheConnector",
|
||||
"kv_role":"kv_both",
|
||||
"kv_connector_extra_config":{"host":"127.0.0.1","port":65432}
|
||||
}`))
|
||||
})
|
||||
|
||||
It("accepts a pre-encoded JSON string for the object-valued knobs", func() {
|
||||
// A config written by hand (or round-tripped through a flat store) may
|
||||
// carry the object as a string; both spellings reach the engine the same.
|
||||
lo := parseOptions(&pb.ModelOptions{EngineArgs: `{
|
||||
"speculative_config": "{\"method\":\"ngram\",\"num_speculative_tokens\":4}"
|
||||
}`})
|
||||
Expect(lo.speculativeConfig).To(MatchJSON(`{"method":"ngram","num_speculative_tokens":4}`))
|
||||
})
|
||||
|
||||
It("maps enable_prefix_caching false onto the force-OFF tri-state", func() {
|
||||
// The C ABI tri-state is 0=model default, 1=on, 2=off, so an explicit
|
||||
// `false` must NOT collapse to the 0 that means "let the model decide".
|
||||
lo := parseOptions(&pb.ModelOptions{EngineArgs: `{"enable_prefix_caching": false}`})
|
||||
Expect(lo.enablePrefixCaching).To(Equal(int32(2)))
|
||||
})
|
||||
|
||||
It("leaves the prefix-caching tri-state at the model default when unset", func() {
|
||||
lo := parseOptions(&pb.ModelOptions{EngineArgs: `{"max_num_seqs": 4}`})
|
||||
Expect(lo.enablePrefixCaching).To(Equal(int32(0)))
|
||||
})
|
||||
|
||||
It("accepts the radix-attention alias upstream documents for prefix caching", func() {
|
||||
lo := parseOptions(&pb.ModelOptions{EngineArgs: `{"enable_radix_attention": true}`})
|
||||
Expect(lo.enablePrefixCaching).To(Equal(int32(1)))
|
||||
})
|
||||
|
||||
It("maps enable_jump_forward onto its own tri-state", func() {
|
||||
// ABI v10. Same tri-state shape as prefix caching, and the same trap:
|
||||
// an explicit false must be force-OFF (2), not the 0 that defers to the
|
||||
// environment.
|
||||
on := parseOptions(&pb.ModelOptions{EngineArgs: `{"enable_jump_forward": true}`})
|
||||
Expect(on.enableJumpForward).To(Equal(int32(1)))
|
||||
off := parseOptions(&pb.ModelOptions{EngineArgs: `{"enable_jump_forward": false}`})
|
||||
Expect(off.enableJumpForward).To(Equal(int32(2)))
|
||||
unset := parseOptions(&pb.ModelOptions{EngineArgs: `{"max_num_seqs": 4}`})
|
||||
Expect(unset.enableJumpForward).To(Equal(int32(0)))
|
||||
})
|
||||
|
||||
It("reads enable_jump_forward from the legacy options list too", func() {
|
||||
lo := parseOptions(&pb.ModelOptions{Options: []string{"enable_jump_forward:true"}})
|
||||
Expect(lo.enableJumpForward).To(Equal(int32(1)))
|
||||
})
|
||||
|
||||
It("lets engine_args override the legacy options list", func() {
|
||||
lo := parseOptions(&pb.ModelOptions{
|
||||
Options: []string{"max_num_seqs:8", "block_size:16"},
|
||||
EngineArgs: `{"max_num_seqs": 64}`,
|
||||
})
|
||||
Expect(lo.maxNumSeqs).To(Equal(int32(64))) // engine_args wins
|
||||
Expect(lo.blockSize).To(Equal(int32(16))) // untouched keys survive
|
||||
})
|
||||
|
||||
It("ignores malformed engine_args rather than failing the load", func() {
|
||||
lo := parseOptions(&pb.ModelOptions{
|
||||
Options: []string{"max_num_seqs:8"},
|
||||
EngineArgs: `{not json`,
|
||||
})
|
||||
Expect(lo.maxNumSeqs).To(Equal(int32(8)))
|
||||
})
|
||||
|
||||
It("ignores unknown keys", func() {
|
||||
lo := parseOptions(&pb.ModelOptions{EngineArgs: `{"gpu_memory_utilization": 0.9}`})
|
||||
Expect(lo).To(Equal(loadOptions{}))
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("samplingFromPredict", func() {
|
||||
@@ -135,6 +298,91 @@ var _ = Describe("samplingFromPredict", func() {
|
||||
})
|
||||
})
|
||||
|
||||
// The engine resolves speculative_config.model against a local directory or
|
||||
// ~/.cache/huggingface/hub ONLY - it never downloads. LocalAI keeps models in
|
||||
// its own directory, so a bare repo id would miss the HF cache and fail deep in
|
||||
// the load with a confusing "draft checkpoint not found". Resolve it here.
|
||||
var _ = Describe("resolveDraftModelPath", func() {
|
||||
var modelsDir string
|
||||
|
||||
BeforeEach(func() {
|
||||
modelsDir = GinkgoT().TempDir()
|
||||
})
|
||||
|
||||
// draftDir creates a plausible draft checkpoint under models/.
|
||||
draftDir := func(name string) string {
|
||||
d := filepath.Join(modelsDir, name)
|
||||
Expect(os.MkdirAll(d, 0o750)).To(Succeed())
|
||||
Expect(os.WriteFile(filepath.Join(d, "config.json"), []byte("{}"), 0o600)).To(Succeed())
|
||||
return d
|
||||
}
|
||||
|
||||
It("rewrites a repo id to the matching directory in the models dir", func() {
|
||||
want := draftDir("Qwen3.6-27B-DFlash")
|
||||
spec := `{"method":"dflash","model":"z-lab/Qwen3.6-27B-DFlash"}`
|
||||
out, err := resolveDraftModelPath(spec, modelsDir)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(out).To(MatchJSON(`{"method":"dflash","model":"` + want + `"}`))
|
||||
})
|
||||
|
||||
It("rewrites a models-dir-relative path", func() {
|
||||
want := draftDir("drafts__dflash")
|
||||
spec := `{"method":"dflash","model":"drafts__dflash"}`
|
||||
out, err := resolveDraftModelPath(spec, modelsDir)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(out).To(ContainSubstring(want))
|
||||
})
|
||||
|
||||
It("leaves an absolute path that already resolves alone", func() {
|
||||
abs := draftDir("elsewhere")
|
||||
spec := `{"method":"dflash","model":"` + abs + `"}`
|
||||
out, err := resolveDraftModelPath(spec, modelsDir)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(out).To(MatchJSON(spec))
|
||||
})
|
||||
|
||||
It("fails with an actionable error when the draft is nowhere on disk", func() {
|
||||
// Silently passing the repo id through would surface as an HF-cache
|
||||
// miss inside the engine, which reads as "your model is broken".
|
||||
spec := `{"method":"dflash","model":"z-lab/Not-Downloaded"}`
|
||||
_, err := resolveDraftModelPath(spec, modelsDir)
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(err.Error()).To(ContainSubstring("z-lab/Not-Downloaded"))
|
||||
Expect(err.Error()).To(ContainSubstring(modelsDir))
|
||||
})
|
||||
|
||||
It("requires a model key for dflash", func() {
|
||||
_, err := resolveDraftModelPath(`{"method":"dflash"}`, modelsDir)
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(err.Error()).To(ContainSubstring("model"))
|
||||
})
|
||||
|
||||
It("leaves mtp and ngram configs untouched", func() {
|
||||
// Neither has a separate draft checkpoint to resolve.
|
||||
for _, spec := range []string{
|
||||
`{"method":"mtp"}`,
|
||||
`{"method":"ngram","num_speculative_tokens":4}`,
|
||||
} {
|
||||
out, err := resolveDraftModelPath(spec, modelsDir)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(out).To(MatchJSON(spec))
|
||||
}
|
||||
})
|
||||
|
||||
It("passes a malformed document through for the engine to reject", func() {
|
||||
// The engine owns config validation and produces the better message.
|
||||
out, err := resolveDraftModelPath(`{not json`, modelsDir)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(out).To(Equal(`{not json`))
|
||||
})
|
||||
|
||||
It("is a no-op on an empty config", func() {
|
||||
out, err := resolveDraftModelPath("", modelsDir)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(out).To(BeEmpty())
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("validModelPath", func() {
|
||||
It("accepts a .gguf file", func() {
|
||||
dir := GinkgoT().TempDir()
|
||||
|
||||
@@ -8,7 +8,7 @@ JOBS?=$(shell nproc --ignore=1)
|
||||
|
||||
# whisper.cpp version
|
||||
WHISPER_REPO?=https://github.com/ggml-org/whisper.cpp
|
||||
WHISPER_CPP_VERSION?=2ca53bb45e38748d07b310eeb36245a7157ac882
|
||||
WHISPER_CPP_VERSION?=4834a2327d008ace3ec5a9ed00f51454bcabbc1c
|
||||
SO_TARGET?=libgowhisper.so
|
||||
|
||||
CMAKE_ARGS+=-DBUILD_SHARED_LIBS=OFF
|
||||
|
||||
+163
-6
@@ -11,6 +11,8 @@
|
||||
- https://github.com/ggerganov/llama.cpp
|
||||
tags:
|
||||
- text-to-text
|
||||
- text-to-speech
|
||||
- TTS
|
||||
- LLM
|
||||
- CPU
|
||||
- GPU
|
||||
@@ -148,6 +150,7 @@
|
||||
- audio-transcription
|
||||
- CPU
|
||||
- CUDA
|
||||
- HIP
|
||||
- Metal
|
||||
# No vulkan key: the vulkan image would carry a Vulkan loader with no Mesa ICD
|
||||
# (see the audio-cpp block in .github/backend-matrix.yml). Pointing a
|
||||
@@ -159,6 +162,7 @@
|
||||
nvidia: "cuda12-audio-cpp"
|
||||
nvidia-cuda-12: "cuda12-audio-cpp"
|
||||
nvidia-cuda-13: "cuda13-audio-cpp"
|
||||
amd: "rocm-audio-cpp"
|
||||
metal: "metal-audio-cpp"
|
||||
metal-darwin-arm64: "metal-audio-cpp"
|
||||
- &whispercpp
|
||||
@@ -193,12 +197,22 @@
|
||||
alias: "vllm-cpp"
|
||||
license: apache-2.0
|
||||
description: |
|
||||
vllm.cpp is a from-scratch C++20 port of vLLM created and maintained by the LocalAI team.
|
||||
It mirrors vLLM's V1 architecture (paged KV cache, continuous batching, prefix caching,
|
||||
scheduler, sampler) on a portable tensor runtime with no Python, PyTorch or ggml at
|
||||
inference time. It loads Hugging Face safetensors and GGUF checkpoints, supports
|
||||
structured output (JSON schema / regex / choice / GBNF grammar) enforced in-engine,
|
||||
and runs on CPU, NVIDIA CUDA (Blackwell-family), Apple Metal and Vulkan.
|
||||
ALPHA development builds. Try it, but llama-cpp stays the recommendation for
|
||||
production use.
|
||||
|
||||
vllm.cpp is an Apache-2.0 C++20 inference engine maintained by the LocalAI team,
|
||||
developed in its own repository and usable without LocalAI. It began as a port of
|
||||
vLLM and keeps vLLM as its reference implementation, checking output against it and
|
||||
benchmarking against it, while growing a featureset of its own. It implements vLLM's
|
||||
V1 architecture (paged KV cache, continuous batching, prefix caching, scheduler,
|
||||
sampler) on a portable tensor runtime with no Python, PyTorch or ggml at inference
|
||||
time. It loads GGUF as well as Hugging Face safetensors, supports structured output
|
||||
(JSON schema / regex / choice / GBNF grammar) enforced in-engine, ships speculative
|
||||
decoding and KV offload, and runs on CPU, NVIDIA CUDA (Blackwell-family), Apple
|
||||
Metal and Vulkan.
|
||||
|
||||
The project is expected to be renamed as it diverges further from vLLM; the new
|
||||
name is still to be decided.
|
||||
urls:
|
||||
- https://github.com/mudler/vllm.cpp
|
||||
tags:
|
||||
@@ -282,6 +296,55 @@
|
||||
nvidia-cuda-12: "cuda12-parakeet-cpp"
|
||||
nvidia-l4t-cuda-12: "nvidia-l4t-arm64-parakeet-cpp"
|
||||
nvidia-l4t-cuda-13: "cuda13-nvidia-l4t-arm64-parakeet-cpp"
|
||||
- &nemospeechcpp
|
||||
name: "nemo-speech-cpp"
|
||||
alias: "nemo-speech-cpp"
|
||||
license: apache-2.0
|
||||
icon: https://avatars.githubusercontent.com/u/1728152?s=200&v=4
|
||||
description: |
|
||||
NVIDIA NeMo-Speech.cpp, a C++/ggml runtime for NVIDIA Nemotron Speech models.
|
||||
One backend serves four model families, selected automatically from the GGUF
|
||||
general.architecture key: automatic speech recognition (offline, cache-aware
|
||||
streaming and live transcription, with optional Silero VAD, punctuation,
|
||||
inverse text normalization and Sortformer speaker diarization attached),
|
||||
standalone Sortformer diarization, MagpieTTS text-to-speech over NanoCodec,
|
||||
and Riva-Translate text translation. Runs on CPU, NVIDIA CUDA, Vulkan,
|
||||
NVIDIA Jetson (L4T) and Apple Metal.
|
||||
urls:
|
||||
- https://github.com/NVIDIA/NeMo-Speech.cpp
|
||||
tags:
|
||||
- audio-transcription
|
||||
- text-to-speech
|
||||
- diarization
|
||||
- text-to-text
|
||||
- CPU
|
||||
- GPU
|
||||
- CUDA
|
||||
- Metal
|
||||
# No amd and no intel key on purpose: upstream NeMo-Speech.cpp has no ROCm/HIP
|
||||
# and no SYCL backend, so there is nothing to point those at. A host reporting
|
||||
# either capability falls through to "default" (SystemState.Capability) and
|
||||
# gets the CPU build, which is the honest answer rather than a broken tag.
|
||||
#
|
||||
# Listing only nvidia-l4t would be a silent downgrade: a Jetson that reports a
|
||||
# CUDA-refined capability would miss the map and fall back to the CPU build.
|
||||
#
|
||||
# The two nvidia-l4t-cuda-* keys point at DIFFERENT images on purpose. The
|
||||
# JetPack r36.4.0 base links ggml against CUDA 12, so serving it to a host that
|
||||
# reports nvidia-l4t-cuda-13 would fail at dlopen on a missing libcudart.so.12.
|
||||
# That is worse than no key at all, since a missing key falls back to a working
|
||||
# CPU build. Hence the separate cuda13 L4T image, as parakeet-cpp and
|
||||
# moss-transcribe-cpp both do.
|
||||
capabilities:
|
||||
default: "cpu-nemo-speech-cpp"
|
||||
nvidia: "cuda12-nemo-speech-cpp"
|
||||
metal: "metal-nemo-speech-cpp"
|
||||
vulkan: "vulkan-nemo-speech-cpp"
|
||||
nvidia-l4t: "nvidia-l4t-arm64-nemo-speech-cpp"
|
||||
nvidia-cuda-13: "cuda13-nemo-speech-cpp"
|
||||
nvidia-cuda-12: "cuda12-nemo-speech-cpp"
|
||||
nvidia-l4t-cuda-12: "nvidia-l4t-arm64-nemo-speech-cpp"
|
||||
nvidia-l4t-cuda-13: "cuda13-nvidia-l4t-arm64-nemo-speech-cpp"
|
||||
- &mosstranscribecpp
|
||||
name: "moss-transcribe-cpp"
|
||||
alias: "moss-transcribe-cpp"
|
||||
@@ -2040,6 +2103,7 @@
|
||||
nvidia: "cuda12-audio-cpp-development"
|
||||
nvidia-cuda-12: "cuda12-audio-cpp-development"
|
||||
nvidia-cuda-13: "cuda13-audio-cpp-development"
|
||||
amd: "rocm-audio-cpp-development"
|
||||
metal: "metal-audio-cpp-development"
|
||||
metal-darwin-arm64: "metal-audio-cpp-development"
|
||||
- !!merge <<: *stablediffusionggml
|
||||
@@ -3264,6 +3328,89 @@
|
||||
uri: "quay.io/go-skynet/local-ai-backends:master-gpu-nvidia-cuda-13-parakeet-cpp"
|
||||
mirrors:
|
||||
- localai/localai-backends:master-gpu-nvidia-cuda-13-parakeet-cpp
|
||||
## nemo-speech-cpp
|
||||
- !!merge <<: *nemospeechcpp
|
||||
name: "nemo-speech-cpp-development"
|
||||
capabilities:
|
||||
default: "cpu-nemo-speech-cpp-development"
|
||||
nvidia: "cuda12-nemo-speech-cpp-development"
|
||||
metal: "metal-nemo-speech-cpp-development"
|
||||
vulkan: "vulkan-nemo-speech-cpp-development"
|
||||
nvidia-l4t: "nvidia-l4t-arm64-nemo-speech-cpp-development"
|
||||
nvidia-cuda-13: "cuda13-nemo-speech-cpp-development"
|
||||
nvidia-cuda-12: "cuda12-nemo-speech-cpp-development"
|
||||
nvidia-l4t-cuda-12: "nvidia-l4t-arm64-nemo-speech-cpp-development"
|
||||
nvidia-l4t-cuda-13: "cuda13-nvidia-l4t-arm64-nemo-speech-cpp-development"
|
||||
- !!merge <<: *nemospeechcpp
|
||||
name: "cpu-nemo-speech-cpp"
|
||||
uri: "quay.io/go-skynet/local-ai-backends:latest-cpu-nemo-speech-cpp"
|
||||
mirrors:
|
||||
- localai/localai-backends:latest-cpu-nemo-speech-cpp
|
||||
- !!merge <<: *nemospeechcpp
|
||||
name: "cpu-nemo-speech-cpp-development"
|
||||
uri: "quay.io/go-skynet/local-ai-backends:master-cpu-nemo-speech-cpp"
|
||||
mirrors:
|
||||
- localai/localai-backends:master-cpu-nemo-speech-cpp
|
||||
- !!merge <<: *nemospeechcpp
|
||||
name: "cuda12-nemo-speech-cpp"
|
||||
uri: "quay.io/go-skynet/local-ai-backends:latest-gpu-nvidia-cuda-12-nemo-speech-cpp"
|
||||
mirrors:
|
||||
- localai/localai-backends:latest-gpu-nvidia-cuda-12-nemo-speech-cpp
|
||||
- !!merge <<: *nemospeechcpp
|
||||
name: "cuda12-nemo-speech-cpp-development"
|
||||
uri: "quay.io/go-skynet/local-ai-backends:master-gpu-nvidia-cuda-12-nemo-speech-cpp"
|
||||
mirrors:
|
||||
- localai/localai-backends:master-gpu-nvidia-cuda-12-nemo-speech-cpp
|
||||
- !!merge <<: *nemospeechcpp
|
||||
name: "cuda13-nemo-speech-cpp"
|
||||
uri: "quay.io/go-skynet/local-ai-backends:latest-gpu-nvidia-cuda-13-nemo-speech-cpp"
|
||||
mirrors:
|
||||
- localai/localai-backends:latest-gpu-nvidia-cuda-13-nemo-speech-cpp
|
||||
- !!merge <<: *nemospeechcpp
|
||||
name: "cuda13-nemo-speech-cpp-development"
|
||||
uri: "quay.io/go-skynet/local-ai-backends:master-gpu-nvidia-cuda-13-nemo-speech-cpp"
|
||||
mirrors:
|
||||
- localai/localai-backends:master-gpu-nvidia-cuda-13-nemo-speech-cpp
|
||||
- !!merge <<: *nemospeechcpp
|
||||
name: "vulkan-nemo-speech-cpp"
|
||||
uri: "quay.io/go-skynet/local-ai-backends:latest-gpu-vulkan-nemo-speech-cpp"
|
||||
mirrors:
|
||||
- localai/localai-backends:latest-gpu-vulkan-nemo-speech-cpp
|
||||
- !!merge <<: *nemospeechcpp
|
||||
name: "vulkan-nemo-speech-cpp-development"
|
||||
uri: "quay.io/go-skynet/local-ai-backends:master-gpu-vulkan-nemo-speech-cpp"
|
||||
mirrors:
|
||||
- localai/localai-backends:master-gpu-vulkan-nemo-speech-cpp
|
||||
- !!merge <<: *nemospeechcpp
|
||||
name: "nvidia-l4t-arm64-nemo-speech-cpp"
|
||||
uri: "quay.io/go-skynet/local-ai-backends:latest-nvidia-l4t-arm64-nemo-speech-cpp"
|
||||
mirrors:
|
||||
- localai/localai-backends:latest-nvidia-l4t-arm64-nemo-speech-cpp
|
||||
- !!merge <<: *nemospeechcpp
|
||||
name: "nvidia-l4t-arm64-nemo-speech-cpp-development"
|
||||
uri: "quay.io/go-skynet/local-ai-backends:master-nvidia-l4t-arm64-nemo-speech-cpp"
|
||||
mirrors:
|
||||
- localai/localai-backends:master-nvidia-l4t-arm64-nemo-speech-cpp
|
||||
- !!merge <<: *nemospeechcpp
|
||||
name: "cuda13-nvidia-l4t-arm64-nemo-speech-cpp"
|
||||
uri: "quay.io/go-skynet/local-ai-backends:latest-nvidia-l4t-cuda-13-arm64-nemo-speech-cpp"
|
||||
mirrors:
|
||||
- localai/localai-backends:latest-nvidia-l4t-cuda-13-arm64-nemo-speech-cpp
|
||||
- !!merge <<: *nemospeechcpp
|
||||
name: "cuda13-nvidia-l4t-arm64-nemo-speech-cpp-development"
|
||||
uri: "quay.io/go-skynet/local-ai-backends:master-nvidia-l4t-cuda-13-arm64-nemo-speech-cpp"
|
||||
mirrors:
|
||||
- localai/localai-backends:master-nvidia-l4t-cuda-13-arm64-nemo-speech-cpp
|
||||
- !!merge <<: *nemospeechcpp
|
||||
name: "metal-nemo-speech-cpp"
|
||||
uri: "quay.io/go-skynet/local-ai-backends:latest-metal-darwin-arm64-nemo-speech-cpp"
|
||||
mirrors:
|
||||
- localai/localai-backends:latest-metal-darwin-arm64-nemo-speech-cpp
|
||||
- !!merge <<: *nemospeechcpp
|
||||
name: "metal-nemo-speech-cpp-development"
|
||||
uri: "quay.io/go-skynet/local-ai-backends:master-metal-darwin-arm64-nemo-speech-cpp"
|
||||
mirrors:
|
||||
- localai/localai-backends:master-metal-darwin-arm64-nemo-speech-cpp
|
||||
## moss-transcribe-cpp
|
||||
- !!merge <<: *mosstranscribecpp
|
||||
name: "moss-transcribe-cpp-development"
|
||||
@@ -6900,6 +7047,16 @@
|
||||
uri: "quay.io/go-skynet/local-ai-backends:master-gpu-nvidia-cuda-13-audio-cpp"
|
||||
mirrors:
|
||||
- localai/localai-backends:master-gpu-nvidia-cuda-13-audio-cpp
|
||||
- !!merge <<: *audiocpp
|
||||
name: "rocm-audio-cpp"
|
||||
uri: "quay.io/go-skynet/local-ai-backends:latest-gpu-rocm-hipblas-audio-cpp"
|
||||
mirrors:
|
||||
- localai/localai-backends:latest-gpu-rocm-hipblas-audio-cpp
|
||||
- !!merge <<: *audiocpp
|
||||
name: "rocm-audio-cpp-development"
|
||||
uri: "quay.io/go-skynet/local-ai-backends:master-gpu-rocm-hipblas-audio-cpp"
|
||||
mirrors:
|
||||
- localai/localai-backends:master-gpu-rocm-hipblas-audio-cpp
|
||||
- !!merge <<: *audiocpp
|
||||
name: "metal-audio-cpp"
|
||||
uri: "quay.io/go-skynet/local-ai-backends:latest-metal-darwin-arm64-audio-cpp"
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
grpcio==1.83.0
|
||||
protobuf
|
||||
certifi
|
||||
packaging==26.2
|
||||
packaging==26.3
|
||||
@@ -1,7 +1,10 @@
|
||||
.PHONY: fish-speech
|
||||
fish-speech:
|
||||
.PHONY: fish-speech test-source-preparation
|
||||
fish-speech: test-source-preparation
|
||||
bash install.sh
|
||||
|
||||
test-source-preparation:
|
||||
bash prepare-source_test.sh
|
||||
|
||||
.PHONY: run
|
||||
run: fish-speech
|
||||
@echo "Running fish-speech..."
|
||||
|
||||
@@ -39,10 +39,10 @@ else
|
||||
cd "${FISH_SPEECH_DIR}" && git pull && cd -
|
||||
fi
|
||||
|
||||
# Remove pyaudio from fish-speech deps — it's only used by the upstream client tool
|
||||
# (tools/api_client.py) for speaker playback, not by our gRPC backend server.
|
||||
# It requires native portaudio libs which aren't available on all build environments.
|
||||
sed -i.bak '/"pyaudio"/d' "${FISH_SPEECH_DIR}/pyproject.toml"
|
||||
# Keep the platform-specific PyTorch installed above. Upstream pins the generic
|
||||
# PyPI torch wheel, which replaces ROCm builds with a CUDA wheel during the
|
||||
# editable install. pyaudio is only used by the upstream playback client.
|
||||
bash "${backend_dir}/prepare-source.sh" "${BUILD_TYPE:-}" "${FISH_SPEECH_DIR}/pyproject.toml"
|
||||
|
||||
# Install fish-speech deps from source (without the package itself since we use PYTHONPATH)
|
||||
ensureVenv
|
||||
|
||||
Executable
+18
@@ -0,0 +1,18 @@
|
||||
#!/bin/bash
|
||||
set -euo pipefail
|
||||
|
||||
build_type=${1:-}
|
||||
pyproject=${2:?usage: prepare-source.sh BUILD_TYPE PYPROJECT}
|
||||
prepared=$(mktemp "${pyproject}.XXXXXX")
|
||||
trap 'rm -f "$prepared"' EXIT
|
||||
|
||||
awk -v build_type="$build_type" '
|
||||
/^dependencies = \[$/ { in_project_dependencies = 1 }
|
||||
build_type == "hipblas" && in_project_dependencies && /^[[:space:]]*"(torch|torchaudio)[^"]*",?[[:space:]]*$/ { next }
|
||||
in_project_dependencies && /^[[:space:]]*"pyaudio",?[[:space:]]*$/ { next }
|
||||
{ print }
|
||||
in_project_dependencies && /^\]$/ { in_project_dependencies = 0 }
|
||||
' "$pyproject" > "$prepared"
|
||||
|
||||
mv "$prepared" "$pyproject"
|
||||
trap - EXIT
|
||||
+66
@@ -0,0 +1,66 @@
|
||||
#!/bin/bash
|
||||
set -euo pipefail
|
||||
|
||||
SCRIPT_DIR=$(dirname "$(realpath "$0")")
|
||||
WORK_DIR=$(mktemp -d)
|
||||
trap 'rm -rf "$WORK_DIR"' EXIT
|
||||
|
||||
write_fixture() {
|
||||
cat > "$1" <<'EOF'
|
||||
[project]
|
||||
dependencies = [
|
||||
"numpy",
|
||||
"torch==2.8.0",
|
||||
"torchaudio==2.8.0",
|
||||
"pyaudio",
|
||||
]
|
||||
|
||||
[project.optional-dependencies]
|
||||
stable = [
|
||||
"torch==2.8.0",
|
||||
"torchaudio",
|
||||
]
|
||||
EOF
|
||||
}
|
||||
|
||||
write_fixture "$WORK_DIR/rocm.toml"
|
||||
write_fixture "$WORK_DIR/cuda.toml"
|
||||
write_fixture "$WORK_DIR/cpu.toml"
|
||||
|
||||
bash "$SCRIPT_DIR/prepare-source.sh" hipblas "$WORK_DIR/rocm.toml"
|
||||
bash "$SCRIPT_DIR/prepare-source.sh" cublas "$WORK_DIR/cuda.toml"
|
||||
bash "$SCRIPT_DIR/prepare-source.sh" "" "$WORK_DIR/cpu.toml"
|
||||
|
||||
cat > "$WORK_DIR/expected-rocm.toml" <<'EOF'
|
||||
[project]
|
||||
dependencies = [
|
||||
"numpy",
|
||||
]
|
||||
|
||||
[project.optional-dependencies]
|
||||
stable = [
|
||||
"torch==2.8.0",
|
||||
"torchaudio",
|
||||
]
|
||||
EOF
|
||||
|
||||
cat > "$WORK_DIR/expected-default.toml" <<'EOF'
|
||||
[project]
|
||||
dependencies = [
|
||||
"numpy",
|
||||
"torch==2.8.0",
|
||||
"torchaudio==2.8.0",
|
||||
]
|
||||
|
||||
[project.optional-dependencies]
|
||||
stable = [
|
||||
"torch==2.8.0",
|
||||
"torchaudio",
|
||||
]
|
||||
EOF
|
||||
|
||||
diff -u "$WORK_DIR/expected-rocm.toml" "$WORK_DIR/rocm.toml"
|
||||
diff -u "$WORK_DIR/expected-default.toml" "$WORK_DIR/cuda.toml"
|
||||
diff -u "$WORK_DIR/expected-default.toml" "$WORK_DIR/cpu.toml"
|
||||
|
||||
echo "PASS: source preparation preserves each platform's PyTorch dependencies"
|
||||
@@ -8,4 +8,5 @@ else
|
||||
source $backend_dir/../common/libbackend.sh
|
||||
fi
|
||||
|
||||
bash "${backend_dir}/prepare-source_test.sh"
|
||||
runUnittests
|
||||
@@ -127,7 +127,10 @@ class BackendServicer(backend_pb2_grpc.BackendServicer):
|
||||
context.set_code(grpc.StatusCode.NOT_FOUND)
|
||||
context.set_details("no face detected")
|
||||
return backend_pb2.EmbeddingResult()
|
||||
return backend_pb2.EmbeddingResult(embeddings=[float(x) for x in vec])
|
||||
return backend_pb2.EmbeddingResult(
|
||||
embeddings=[float(x) for x in vec],
|
||||
layout=backend_pb2.EMBEDDING_LAYOUT_FINAL,
|
||||
)
|
||||
|
||||
def Detect(self, request, context):
|
||||
if self.engine is None:
|
||||
|
||||
Loaded 100 of 645 files, more files were not shown because too many files have changed in this diff.
Show more
Reference in new issue
Block a user