mirror of
https://github.com/mudler/LocalAI.git
synced 2026-09-11 05:34:29 -04:00
Compare commits
No files matched your search
@@ -59,7 +59,9 @@ backend/rust/*/target
|
||||
backend-images
|
||||
local-backends
|
||||
local-ai
|
||||
.claude
|
||||
.crush
|
||||
.tools
|
||||
protoc
|
||||
tests
|
||||
|
||||
|
||||
@@ -2584,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"
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -14,6 +14,7 @@ on:
|
||||
|
||||
permissions:
|
||||
contents: write
|
||||
pull-requests: write
|
||||
|
||||
concurrency:
|
||||
group: refresh-site-counters
|
||||
@@ -30,15 +31,25 @@ jobs:
|
||||
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
|
||||
@@ -526,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'
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -676,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):
|
||||
##
|
||||
@@ -705,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" \
|
||||
@@ -724,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
|
||||
|
||||
@@ -813,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
|
||||
@@ -1505,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
|
||||
@@ -1520,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).
|
||||
|
||||
@@ -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?=7efbb58def443722ea540d931dd3debee3e4d5e8
|
||||
AUDIO_CPP_VERSION?=a61da671b6a81c79071500954eea3c91c1a383dd
|
||||
AUDIO_CPP_REPO?=https://github.com/0xShug0/audio.cpp
|
||||
|
||||
CURRENT_MAKEFILE_DIR := $(dir $(abspath $(lastword $(MAKEFILE_LIST))))
|
||||
|
||||
@@ -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,
|
||||
@@ -241,6 +243,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 +266,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;
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -1,10 +1,10 @@
|
||||
# ds4 backend Makefile.
|
||||
#
|
||||
# Upstream pin lives below as DS4_VERSION?=b0309611041655f4e45671cfd9c9886aff161406
|
||||
# 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?=b0309611041655f4e45671cfd9c9886aff161406
|
||||
DS4_VERSION?=84cc882352757baf628a1776badf7cc54d584e28
|
||||
DS4_REPO?=https://github.com/antirez/ds4
|
||||
|
||||
CURRENT_MAKEFILE_DIR := $(dir $(abspath $(lastword $(MAKEFILE_LIST))))
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
|
||||
IK_LLAMA_VERSION?=cf1aa57e1a0fabfd015831718fc99d1aec01ada5
|
||||
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?=60addddf3c567c43ec3caf70fc953fba3572d96f
|
||||
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?=21901d3f7c23554f072964828363e49ddbc2dc68
|
||||
CRISPASR_VERSION?=a153b09b37c90cd55cd9336fccbdf3ba7a289596
|
||||
SO_TARGET?=libgocrispasr.so
|
||||
|
||||
CMAKE_ARGS+=-DBUILD_SHARED_LIBS=OFF
|
||||
|
||||
@@ -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() {
|
||||
|
||||
@@ -12,7 +12,7 @@
|
||||
# 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?=2e12e2def8a98ed06666f7ee3ca94e7193e04be4
|
||||
NEMO_SPEECH_VERSION?=4f9676226f667d14608487df744f375db87127f8
|
||||
NEMO_SPEECH_REPO?=https://github.com/NVIDIA/NeMo-Speech.cpp
|
||||
|
||||
GOCMD?=go
|
||||
@@ -88,6 +88,18 @@ 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
|
||||
|
||||
@@ -152,7 +164,7 @@ else
|
||||
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 patch-itn-headers
|
||||
.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
|
||||
|
||||
@@ -266,11 +278,28 @@ patch-itn-headers:
|
||||
|
||||
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)
|
||||
NEMO_RUNTIME_PREREQS+=$(ITN_MARKER)
|
||||
endif
|
||||
|
||||
# Upstream sets CMAKE_LIBRARY_OUTPUT_DIRECTORY to ${CMAKE_BINARY_DIR}/bin, so the
|
||||
|
||||
@@ -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?=c6beeef35526c6dc94b74a7fb69f9d2e6a2a7a12
|
||||
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
|
||||
@@ -11,7 +11,7 @@ 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?=0757cac231ecd571a83c4fd2f50805c9251fc225
|
||||
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
|
||||
@@ -47,26 +47,35 @@ 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
|
||||
@@ -106,13 +115,24 @@ 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
|
||||
|
||||
ifeq ($(MLX_ENABLED),1)
|
||||
# A stamp FILE, not a phony target: a phony prerequisite is always "newer" than
|
||||
@@ -165,7 +185,7 @@ $(LIB): sources/vllm.cpp $(MLX_STAMP)
|
||||
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
|
||||
|
||||
@@ -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 v10) 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,12 @@ ABI v10) 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
|
||||
@@ -47,6 +56,68 @@ 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
|
||||
|
||||
@@ -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,14 @@ 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
|
||||
@@ -194,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 v10).
|
||||
// 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
|
||||
@@ -21,7 +21,7 @@ import (
|
||||
// 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 = 10
|
||||
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
|
||||
@@ -69,8 +69,20 @@ type cModelParams struct {
|
||||
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)
|
||||
_ [4]byte // trailing pad to the struct's 8-byte alignment
|
||||
// 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 (structured fields included).
|
||||
@@ -117,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{
|
||||
@@ -148,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 {
|
||||
@@ -175,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)
|
||||
}
|
||||
@@ -222,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
|
||||
}
|
||||
@@ -62,6 +62,66 @@ type loadOptions struct {
|
||||
// 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 {
|
||||
@@ -110,10 +170,94 @@ func applyOptionsList(lo *loadOptions, options []string) {
|
||||
if b, err := strconv.ParseBool(strings.TrimSpace(v)); err == nil {
|
||||
lo.enableJumpForward = boolTriState(b)
|
||||
}
|
||||
default:
|
||||
applyVideoOption(&lo.video, strings.TrimSpace(k), v)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 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
|
||||
@@ -160,6 +304,9 @@ func applyEngineArgs(lo *loadOptions, engineArgs string) {
|
||||
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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,7 +16,7 @@ func TestVllmCpp(t *testing.T) {
|
||||
RunSpecs(t, "vllm-cpp suite")
|
||||
}
|
||||
|
||||
// The Go POD mirrors must match the C struct layout of vllm.h (ABI v10)
|
||||
// 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() {
|
||||
@@ -24,7 +24,7 @@ var _ = Describe("C ABI struct mirrors", 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(10))
|
||||
Expect(abiVersion).To(Equal(21))
|
||||
})
|
||||
|
||||
It("cModelParams matches vllm_model_params", func() {
|
||||
@@ -42,10 +42,16 @@ var _ = Describe("C ABI struct mirrors", func() {
|
||||
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.EnableJumpForward)).To(Equal(uintptr(80)))
|
||||
// 88, not 84: the struct is 8-aligned (it holds pointers), so the
|
||||
// trailing int32 is padded out. Go pads identically.
|
||||
Expect(unsafe.Sizeof(p)).To(Equal(uintptr(88)))
|
||||
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 v8)", func() {
|
||||
|
||||
@@ -8,7 +8,7 @@ JOBS?=$(shell nproc --ignore=1)
|
||||
|
||||
# whisper.cpp version
|
||||
WHISPER_REPO?=https://github.com/ggml-org/whisper.cpp
|
||||
WHISPER_CPP_VERSION?=306c88f4d1286aec1bf96e544632897886af5501
|
||||
WHISPER_CPP_VERSION?=4834a2327d008ace3ec5a9ed00f51454bcabbc1c
|
||||
SO_TARGET?=libgowhisper.so
|
||||
|
||||
CMAKE_ARGS+=-DBUILD_SHARED_LIBS=OFF
|
||||
|
||||
@@ -11,6 +11,8 @@
|
||||
- https://github.com/ggerganov/llama.cpp
|
||||
tags:
|
||||
- text-to-text
|
||||
- text-to-speech
|
||||
- TTS
|
||||
- LLM
|
||||
- CPU
|
||||
- GPU
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -638,7 +638,10 @@ class BackendServicer(backend_pb2_grpc.BackendServicer):
|
||||
normalized = (pooled / (norm + 1e-12))
|
||||
vec = normalized.cast(dtypes.float32).tolist()
|
||||
|
||||
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,
|
||||
)
|
||||
except Exception as exc:
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
|
||||
@@ -375,7 +375,10 @@ class BackendServicer(backend_pb2_grpc.BackendServicer):
|
||||
# Pool to get sentence embeddings; i.e. generate one 1024 vector for the entire sentence
|
||||
sentence_embeddings = mean_pooling(model_output, encoded_input['attention_mask'])
|
||||
embeds = sentence_embeddings[0]
|
||||
return backend_pb2.EmbeddingResult(embeddings=embeds)
|
||||
return backend_pb2.EmbeddingResult(
|
||||
embeddings=embeds,
|
||||
layout=backend_pb2.EMBEDDING_LAYOUT_FINAL,
|
||||
)
|
||||
|
||||
async def _predict(self, request, context, streaming=False):
|
||||
set_seed(request.Seed)
|
||||
|
||||
@@ -2,9 +2,9 @@ torch==2.7.1
|
||||
llvmlite==0.43.0
|
||||
numba==0.60.0
|
||||
accelerate
|
||||
transformers>=5.14.1
|
||||
transformers>=5.15.0
|
||||
bitsandbytes
|
||||
sentence-transformers==5.6.1
|
||||
sentence-transformers==5.7.0
|
||||
diffusers
|
||||
soundfile
|
||||
protobuf==7.35.0
|
||||
@@ -2,9 +2,9 @@ torch==2.7.1
|
||||
accelerate
|
||||
llvmlite==0.43.0
|
||||
numba==0.60.0
|
||||
transformers>=5.14.1
|
||||
transformers>=5.15.0
|
||||
bitsandbytes
|
||||
sentence-transformers==5.6.1
|
||||
sentence-transformers==5.7.0
|
||||
diffusers
|
||||
soundfile
|
||||
protobuf==7.35.0
|
||||
@@ -2,9 +2,9 @@
|
||||
torch==2.9.0
|
||||
llvmlite==0.43.0
|
||||
numba==0.60.0
|
||||
transformers>=5.14.1
|
||||
transformers>=5.15.0
|
||||
bitsandbytes
|
||||
sentence-transformers==5.6.1
|
||||
sentence-transformers==5.7.0
|
||||
diffusers
|
||||
soundfile
|
||||
protobuf==7.35.0
|
||||
@@ -1,11 +1,11 @@
|
||||
--extra-index-url https://download.pytorch.org/whl/rocm7.0
|
||||
torch==2.10.0+rocm7.0
|
||||
accelerate
|
||||
transformers>=5.14.1
|
||||
transformers>=5.15.0
|
||||
llvmlite==0.43.0
|
||||
numba==0.60.0
|
||||
bitsandbytes
|
||||
sentence-transformers==5.6.1
|
||||
sentence-transformers==5.7.0
|
||||
diffusers
|
||||
soundfile
|
||||
protobuf==7.35.0
|
||||
@@ -3,9 +3,9 @@ torch
|
||||
optimum[openvino]
|
||||
llvmlite==0.43.0
|
||||
numba==0.60.0
|
||||
transformers>=5.14.1
|
||||
transformers>=5.15.0
|
||||
bitsandbytes
|
||||
sentence-transformers==5.6.1
|
||||
sentence-transformers==5.7.0
|
||||
diffusers
|
||||
soundfile
|
||||
protobuf==7.35.0
|
||||
@@ -2,9 +2,9 @@ torch==2.7.1
|
||||
llvmlite==0.43.0
|
||||
numba==0.60.0
|
||||
accelerate
|
||||
transformers>=5.14.1
|
||||
transformers>=5.15.0
|
||||
bitsandbytes
|
||||
sentence-transformers==5.6.1
|
||||
sentence-transformers==5.7.0
|
||||
diffusers
|
||||
soundfile
|
||||
protobuf==7.35.0
|
||||
@@ -336,7 +336,10 @@ class BackendServicer(backend_pb2_grpc.BackendServicer):
|
||||
context.set_code(grpc.StatusCode.INVALID_ARGUMENT)
|
||||
context.set_details("No embeddings were calculated.")
|
||||
return backend_pb2.EmbeddingResult()
|
||||
return backend_pb2.EmbeddingResult(embeddings=outputs[0].outputs.embedding)
|
||||
return backend_pb2.EmbeddingResult(
|
||||
embeddings=outputs[0].outputs.embedding,
|
||||
layout=backend_pb2.EMBEDDING_LAYOUT_FINAL,
|
||||
)
|
||||
|
||||
async def PredictStream(self, request, context):
|
||||
"""
|
||||
|
||||
@@ -119,14 +119,14 @@ if [ "$(uname -s)" = "Darwin" ]; then
|
||||
# can rewrite it. Darwin therefore follows vllm-metal and can lag the Linux
|
||||
# vllm pin (requirements-cublas13-after.txt, bumped independently against
|
||||
# vllm/vllm) until vllm-metal supports a newer vLLM.
|
||||
VLLM_METAL_VERSION="v0.3.0.dev20260726174827"
|
||||
VLLM_METAL_VERSION="v0.3.0.dev20260818075955"
|
||||
|
||||
# The coupled vLLM source version is whatever this vllm-metal release builds
|
||||
# against -- it declares it in its own installer as `vllm_v=`. Derive it from
|
||||
# against. Derive it from
|
||||
# the PINNED tag rather than hardcoding a second value that could drift. The
|
||||
# tag is immutable, so this stays reproducible across rebuilds.
|
||||
VLLM_VERSION=$(curl -fsSL "https://raw.githubusercontent.com/vllm-project/vllm-metal/${VLLM_METAL_VERSION}/install.sh" \
|
||||
| grep -oE 'vllm_v="[0-9]+\.[0-9]+\.[0-9]+"' | head -n1 | cut -d'"' -f2)
|
||||
| "$backend_dir/../../../scripts/lib/extract-vllm-metal-version.sh")
|
||||
if [ -z "${VLLM_VERSION}" ]; then
|
||||
echo "ERROR: could not derive the vLLM version from vllm-metal ${VLLM_METAL_VERSION}" >&2
|
||||
exit 1
|
||||
@@ -168,7 +168,7 @@ if [ "$(uname -s)" = "Darwin" ]; then
|
||||
|
||||
# Intel XPU has no upstream-published vllm wheels, so we always build vllm
|
||||
# from source against torch-xpu and replace the default triton with
|
||||
# triton-xpu (matching torch 2.11). Mirrors the upstream procedure:
|
||||
# triton-xpu. Mirrors the upstream procedure:
|
||||
# https://github.com/vllm-project/vllm/blob/main/docs/getting_started/installation/gpu.xpu.inc.md
|
||||
elif [ "x${BUILD_TYPE}" == "xintel" ]; then
|
||||
# Hide requirements-intel-after.txt so installRequirements doesn't
|
||||
@@ -194,18 +194,23 @@ elif [ "x${BUILD_TYPE}" == "xintel" ]; then
|
||||
|
||||
_vllm_src=$(mktemp -d)
|
||||
trap 'rm -rf "${_vllm_src}"' EXIT
|
||||
git clone --depth 1 https://github.com/vllm-project/vllm "${_vllm_src}/vllm"
|
||||
# Keep the source build aligned with the version shipped by the other
|
||||
# accelerator profiles. Building the moving main branch can silently pull
|
||||
# a newer torch/XPU runtime than the selected oneAPI base image supports.
|
||||
VLLM_VERSION="0.26.0"
|
||||
git clone --depth 1 --branch "v${VLLM_VERSION}" \
|
||||
https://github.com/vllm-project/vllm "${_vllm_src}/vllm"
|
||||
pushd "${_vllm_src}/vllm"
|
||||
# Install vllm's own runtime deps (torch-xpu, vllm_xpu_kernels,
|
||||
# pydantic, fastapi, …) from upstream's requirements/xpu.txt — the
|
||||
# canonical source of truth. Avoids re-pinning everything ourselves.
|
||||
uv pip install ${EXTRA_PIP_INSTALL_FLAGS:-} -r requirements/xpu.txt
|
||||
# Stock triton (NVIDIA-only) may have come in transitively; replace
|
||||
# with triton-xpu==3.7.0 which matches torch 2.11.
|
||||
# with the version vLLM 0.26.0 specifies for torch 2.12.
|
||||
uv pip uninstall triton triton-xpu 2>/dev/null || true
|
||||
uv pip install ${EXTRA_PIP_INSTALL_FLAGS:-} \
|
||||
--extra-index-url https://download.pytorch.org/whl/xpu \
|
||||
triton-xpu==3.7.0
|
||||
triton-xpu==3.7.1
|
||||
export CMAKE_PREFIX_PATH="$(python -c 'import site; print(site.getsitepackages()[0])'):${CMAKE_PREFIX_PATH:-}"
|
||||
VLLM_TARGET_DEVICE=xpu uv pip install ${EXTRA_PIP_INSTALL_FLAGS:-} --no-deps .
|
||||
popd
|
||||
|
||||
@@ -3,8 +3,8 @@
|
||||
# on a cu130 host. Pull the cu130-flavoured wheel from vLLM's per-tag index
|
||||
# instead — the cublas13 case in install.sh adds --index-strategy=unsafe-best-match
|
||||
# so uv consults this index alongside PyPI.
|
||||
--extra-index-url https://wheels.vllm.ai/0.26.0/cu130
|
||||
--extra-index-url https://wheels.vllm.ai/0.27.1/cu130
|
||||
# VERSION COUPLING: darwin/Apple-Silicon builds use vllm-metal (see install.sh),
|
||||
# which pins this exact vLLM version. Bumping vllm here means coordinating with a
|
||||
# vllm-metal release that supports the new version, or macOS/Metal builds break.
|
||||
vllm==0.26.0
|
||||
vllm==0.27.1
|
||||
@@ -9,4 +9,4 @@
|
||||
# memory architecture crash deterministically with an empty "Engine core init
|
||||
# failed" set (mudler/LocalAI#10722). Leaving this unpinned let the L4T image
|
||||
# drift onto whatever wheel was latest at build time.
|
||||
vllm==0.25.1
|
||||
vllm==0.26.0
|
||||
Submodule backend/rust/kokoros/sources/Kokoros updated: 7089168f0c...29e99ad5a5.
@@ -320,6 +320,13 @@ impl Backend for KokorosService {
|
||||
Err(Status::unimplemented("Not supported"))
|
||||
}
|
||||
|
||||
async fn upscale_image(
|
||||
&self,
|
||||
_: Request<backend::UpscaleImageRequest>,
|
||||
) -> Result<Response<backend::Result>, Status> {
|
||||
Err(Status::unimplemented("Not supported"))
|
||||
}
|
||||
|
||||
async fn generate_image(
|
||||
&self,
|
||||
_: Request<backend::GenerateImageRequest>,
|
||||
|
||||
@@ -1,8 +1,10 @@
|
||||
package application
|
||||
|
||||
import (
|
||||
"cmp"
|
||||
"context"
|
||||
"math/rand/v2"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
@@ -19,12 +21,14 @@ import (
|
||||
"github.com/mudler/LocalAI/core/services/nodes"
|
||||
"github.com/mudler/LocalAI/core/services/routing/admission"
|
||||
"github.com/mudler/LocalAI/core/services/routing/billing"
|
||||
"github.com/mudler/LocalAI/core/services/routing/corpus"
|
||||
"github.com/mudler/LocalAI/core/services/routing/pii"
|
||||
"github.com/mudler/LocalAI/core/services/routing/piidetector"
|
||||
"github.com/mudler/LocalAI/core/services/routing/router"
|
||||
"github.com/mudler/LocalAI/core/services/voiceprofile"
|
||||
"github.com/mudler/LocalAI/core/services/voicerecognition"
|
||||
"github.com/mudler/LocalAI/core/templates"
|
||||
"github.com/mudler/LocalAI/core/trace"
|
||||
pkggrpc "github.com/mudler/LocalAI/pkg/grpc"
|
||||
localaitools "github.com/mudler/LocalAI/pkg/mcp/localaitools"
|
||||
localaiInproc "github.com/mudler/LocalAI/pkg/mcp/localaitools/inproc"
|
||||
@@ -76,6 +80,7 @@ type Application struct {
|
||||
mitmHostConflicts atomic.Pointer[map[string][]string]
|
||||
routerDecisions router.DecisionStore
|
||||
routerRegistry *router.Registry
|
||||
routerCorpus *corpus.Manager
|
||||
admissionLimiter *admission.Limiter
|
||||
watchdogMutex sync.Mutex
|
||||
watchdogStop chan bool
|
||||
@@ -119,6 +124,8 @@ func (a *Application) Ready() bool { return a.startupComplete.Load() }
|
||||
func (a *Application) markStartupComplete() { a.startupComplete.Store(true) }
|
||||
|
||||
func newApplication(appConfig *config.ApplicationConfig) *Application {
|
||||
corebackend.ConfigureGlobalBackendAdmission(appConfig.MaxConcurrentBackendRequests)
|
||||
trace.ConfigureBackendTraceMaxInFlight(appConfig.MaxConcurrentBackendRequests)
|
||||
ml := model.NewModelLoader(appConfig.SystemState)
|
||||
|
||||
// Apply the per-model load-failure cooldown (0 disables). Set here rather
|
||||
@@ -134,7 +141,7 @@ func newApplication(appConfig *config.ApplicationConfig) *Application {
|
||||
// Record a model_load backend trace for every real backend load, so the
|
||||
// Traces UI shows which backend runtime served each model and how long
|
||||
// the load took. Load failures are traced by the modality wrappers.
|
||||
ml.SetLoadObserver(corebackend.ModelLoadTraceObserver(appConfig))
|
||||
ml.SetLoadLifecycleObserver(corebackend.ModelLoadTraceObserver(appConfig))
|
||||
|
||||
app := &Application{
|
||||
backendLoader: config.NewModelConfigLoader(
|
||||
@@ -146,6 +153,10 @@ func newApplication(appConfig *config.ApplicationConfig) *Application {
|
||||
applicationConfig: appConfig,
|
||||
templatesEvaluator: templates.NewEvaluator(appConfig.SystemState.Model.ModelsPath),
|
||||
voiceProfileStore: voiceprofile.NewStore(appConfig.DataPath),
|
||||
// KNN corpus files live under <state dir>/router-corpus (same
|
||||
// DataPath → DynamicConfigsDir precedence the agent pool uses).
|
||||
routerCorpus: corpus.NewManager(filepath.Join(
|
||||
cmp.Or(appConfig.DataPath, appConfig.DynamicConfigsDir, "."), "router-corpus")),
|
||||
}
|
||||
|
||||
// Face-recognition registry backed by LocalAI's built-in vector store.
|
||||
@@ -575,6 +586,13 @@ func (a *Application) start() error {
|
||||
assistantClient.PIIRedactor = a.piiRedactor
|
||||
assistantClient.PIIEvents = a.piiEvents
|
||||
assistantClient.RouterDecisions = a.routerDecisions
|
||||
// Router corpus tools — same factories the RouteModel middleware
|
||||
// uses, so the assistant and the request path agree on store
|
||||
// namespaces and model resolution.
|
||||
assistantClient.RouterCorpus = a.RouterCorpus()
|
||||
assistantClient.RouterEmbedder = a.Embedder
|
||||
assistantClient.RouterEmbedderFingerprint = a.EmbedderFingerprint
|
||||
assistantClient.RouterVectorStore = a.VectorStore
|
||||
if err := holder.Initialize(a.applicationConfig.Context, assistantClient, localaitools.Options{}); err != nil {
|
||||
// Why log+continue instead of fail: the assistant is an optional
|
||||
// feature; a failure here must not take down the whole server.
|
||||
|
||||
@@ -378,6 +378,9 @@ func initDistributed(cfg *config.ApplicationConfig, authDB *gorm.DB, configLoade
|
||||
cfg.Distributed.BackendInstallTimeoutOrDefault(),
|
||||
cfg.Distributed.ModelLoadTimeoutOrDefault(),
|
||||
),
|
||||
// Bounds the REQUEST, not the load: a caller out of budget gets 503 with
|
||||
// live staging progress while the job keeps running underneath.
|
||||
ModelLoadWait: cfg.Distributed.ModelLoadWait,
|
||||
})
|
||||
|
||||
// Wire staging-progress broadcasting so file-staging shows up on every
|
||||
|
||||
@@ -2,10 +2,18 @@ package application
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sort"
|
||||
"strings"
|
||||
|
||||
"github.com/mudler/LocalAI/core/backend"
|
||||
"github.com/mudler/LocalAI/core/config"
|
||||
"github.com/mudler/LocalAI/core/services/routing/corpus"
|
||||
"gopkg.in/yaml.v3"
|
||||
)
|
||||
|
||||
// adapterConfig resolves a model name to its runtime ModelConfig, or nil when
|
||||
@@ -100,6 +108,59 @@ func (a *Application) Embedder(modelName string) backend.Embedder {
|
||||
return &lazyEmbedder{app: a, modelName: modelName}
|
||||
}
|
||||
|
||||
// EmbedderFingerprint returns a stable identity for the embedding space a
|
||||
// named model currently produces. The effective config covers backend and
|
||||
// embedding-affecting options; declared download checksums are part of that
|
||||
// config, while local artifact stat data detects the common in-place file
|
||||
// replacement case. Remote services without a stable artifact identity can
|
||||
// use router.knn.embedding_revision to force invalidation explicitly.
|
||||
func (a *Application) EmbedderFingerprint(modelName string) (string, error) {
|
||||
cfg := a.adapterConfig(modelName)
|
||||
if cfg == nil {
|
||||
return "", fmt.Errorf("embedding model %q not available", modelName)
|
||||
}
|
||||
raw, err := yaml.Marshal(cfg)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("fingerprint embedding model %q config: %w", modelName, err)
|
||||
}
|
||||
h := sha256.New()
|
||||
_, _ = h.Write(raw)
|
||||
|
||||
paths := map[string]struct{}{}
|
||||
for _, name := range append([]string{cfg.Model, cfg.MMProj}, downloadFileNames(cfg)...) {
|
||||
if name == "" || strings.Contains(name, "://") {
|
||||
continue
|
||||
}
|
||||
if !filepath.IsAbs(name) {
|
||||
name = filepath.Join(a.applicationConfig.SystemState.Model.ModelsPath, name)
|
||||
}
|
||||
paths[filepath.Clean(name)] = struct{}{}
|
||||
}
|
||||
ordered := make([]string, 0, len(paths))
|
||||
for path := range paths {
|
||||
ordered = append(ordered, path)
|
||||
}
|
||||
sort.Strings(ordered)
|
||||
for _, path := range ordered {
|
||||
_, _ = h.Write([]byte("\x00" + path))
|
||||
fi, statErr := os.Stat(path)
|
||||
if statErr != nil {
|
||||
_, _ = h.Write([]byte("\x00missing"))
|
||||
continue
|
||||
}
|
||||
_, _ = fmt.Fprintf(h, "\x00%d\x00%d\x00%s", fi.Size(), fi.ModTime().UnixNano(), fi.Mode().String())
|
||||
}
|
||||
return hex.EncodeToString(h.Sum(nil)), nil
|
||||
}
|
||||
|
||||
func downloadFileNames(cfg *config.ModelConfig) []string {
|
||||
out := make([]string, 0, len(cfg.DownloadFiles))
|
||||
for _, f := range cfg.DownloadFiles {
|
||||
out = append(out, f.Filename)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
type lazyEmbedder struct {
|
||||
app *Application
|
||||
modelName string
|
||||
@@ -118,3 +179,9 @@ func (l *lazyEmbedder) Embed(ctx context.Context, text string) ([]float32, error
|
||||
func (a *Application) VectorStore(storeName string) backend.VectorStore {
|
||||
return backend.NewVectorStore(a.modelLoader, a.applicationConfig, a.backendLoader, storeName)
|
||||
}
|
||||
|
||||
// RouterCorpus returns the process-wide KNN corpus manager, built in
|
||||
// newApplication.
|
||||
func (a *Application) RouterCorpus() *corpus.Manager {
|
||||
return a.routerCorpus
|
||||
}
|
||||
@@ -94,6 +94,37 @@ var _ = Describe("router_factories lazy config resolution", func() {
|
||||
})
|
||||
})
|
||||
|
||||
Context("EmbedderFingerprint", func() {
|
||||
It("changes when the effective config or local artifact changes", func() {
|
||||
writeCfg("emb-test", "llama-cpp")
|
||||
artifact := filepath.Join(tmpDir, "emb-test.bin")
|
||||
Expect(os.WriteFile(artifact, []byte("first"), 0o644)).To(Succeed())
|
||||
|
||||
first, err := app.EmbedderFingerprint("emb-test")
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
again, err := app.EmbedderFingerprint("emb-test")
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(again).To(Equal(first))
|
||||
|
||||
Expect(os.WriteFile(artifact, []byte("replacement-with-different-size"), 0o644)).To(Succeed())
|
||||
replaced, err := app.EmbedderFingerprint("emb-test")
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(replaced).NotTo(Equal(first))
|
||||
|
||||
app.backendLoader.UpdateModelConfig("emb-test", func(c *config.ModelConfig) {
|
||||
c.Backend = "rerankers"
|
||||
})
|
||||
updated, err := app.EmbedderFingerprint("emb-test")
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(updated).NotTo(Equal(replaced))
|
||||
})
|
||||
|
||||
It("rejects an unknown model", func() {
|
||||
_, err := app.EmbedderFingerprint("missing")
|
||||
Expect(err).To(HaveOccurred())
|
||||
})
|
||||
})
|
||||
|
||||
Context("Scorer", func() {
|
||||
It("returns nil at construction for an unknown model", func() {
|
||||
Expect(app.Scorer("missing")).To(BeNil())
|
||||
|
||||
@@ -91,12 +91,20 @@ func ModelAudioTransform(
|
||||
return AudioTransformOutputs{}, nil, fmt.Errorf("persist reference: %w", err)
|
||||
}
|
||||
}
|
||||
release, err := AcquireGlobalBackendSlot()
|
||||
if err != nil {
|
||||
return AudioTransformOutputs{}, nil, err
|
||||
}
|
||||
defer release()
|
||||
|
||||
var startTime time.Time
|
||||
var traceID string
|
||||
if appConfig.EnableTracing {
|
||||
trace.InitBackendTracingIfEnabled(appConfig.TracingMaxItems, appConfig.TracingMaxBodyBytes)
|
||||
startTime = time.Now()
|
||||
traceID = trace.BeginBackendTrace(trace.BackendTrace{Timestamp: startTime, Type: trace.BackendTraceAudioTransform, ModelName: modelConfig.Name, Backend: modelConfig.Backend, Summary: trace.TruncateString(filepath.Base(audioPath), 200)})
|
||||
}
|
||||
defer trace.CancelBackendTrace(traceID)
|
||||
|
||||
res, err := transformModel.AudioTransform(ctx, &proto.AudioTransformRequest{
|
||||
ModelIdentity: modelConfig.Model,
|
||||
@@ -126,6 +134,7 @@ func ModelAudioTransform(
|
||||
}
|
||||
}
|
||||
trace.RecordBackendTrace(trace.BackendTrace{
|
||||
ID: traceID,
|
||||
Timestamp: startTime,
|
||||
Duration: time.Since(startTime),
|
||||
Type: trace.BackendTraceAudioTransform,
|
||||
@@ -198,7 +207,17 @@ func ModelAudioTransformStream(
|
||||
if transformModel == nil {
|
||||
return nil, fmt.Errorf("could not load audio-transform model %q", modelConfig.Model)
|
||||
}
|
||||
return transformModel.AudioTransformStream(ctx)
|
||||
release, err := AcquireGlobalBackendSlot()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
stream, err := transformModel.AudioTransformStream(ctx)
|
||||
if err != nil {
|
||||
release()
|
||||
return nil, err
|
||||
}
|
||||
stream.AddCleanup(release)
|
||||
return stream, nil
|
||||
}
|
||||
|
||||
// persistAudioInput copies a transient input file (typically a multipart
|
||||
|
||||
@@ -33,12 +33,20 @@ func Depth(
|
||||
if depthModel == nil {
|
||||
return nil, fmt.Errorf("could not load depth model")
|
||||
}
|
||||
release, err := AcquireGlobalBackendSlot()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer release()
|
||||
|
||||
var startTime time.Time
|
||||
var traceID string
|
||||
if appConfig.EnableTracing {
|
||||
trace.InitBackendTracingIfEnabled(appConfig.TracingMaxItems, appConfig.TracingMaxBodyBytes)
|
||||
startTime = time.Now()
|
||||
traceID = trace.BeginBackendTrace(trace.BackendTrace{Timestamp: startTime, Type: trace.BackendTraceDepth, ModelName: modelConfig.Name, Backend: modelConfig.Backend, Summary: trace.TruncateString(in.GetSrc(), 200)})
|
||||
}
|
||||
defer trace.CancelBackendTrace(traceID)
|
||||
|
||||
// Stamped here for the same reason as in rerank.go: the caller builds the
|
||||
// request without a ModelConfig, this function has the one that loaded.
|
||||
@@ -53,6 +61,7 @@ func Depth(
|
||||
}
|
||||
|
||||
trace.RecordBackendTrace(trace.BackendTrace{
|
||||
ID: traceID,
|
||||
Timestamp: startTime,
|
||||
Duration: time.Since(startTime),
|
||||
Type: trace.BackendTraceDepth,
|
||||
|
||||
@@ -33,11 +33,19 @@ func Detection(
|
||||
return nil, fmt.Errorf("could not load detection model")
|
||||
}
|
||||
|
||||
release, err := AcquireGlobalBackendSlot()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer release()
|
||||
var startTime time.Time
|
||||
var traceID string
|
||||
if appConfig.EnableTracing {
|
||||
trace.InitBackendTracingIfEnabled(appConfig.TracingMaxItems, appConfig.TracingMaxBodyBytes)
|
||||
startTime = time.Now()
|
||||
traceID = trace.BeginBackendTrace(trace.BackendTrace{Timestamp: startTime, Type: trace.BackendTraceDetection, ModelName: modelConfig.Name, Backend: modelConfig.Backend, Summary: trace.TruncateString(sourceFile, 200)})
|
||||
}
|
||||
defer trace.CancelBackendTrace(traceID)
|
||||
|
||||
res, err := detectionModel.Detect(ctx, &proto.DetectOptions{
|
||||
ModelIdentity: modelConfig.Model,
|
||||
@@ -55,6 +63,7 @@ func Detection(
|
||||
}
|
||||
|
||||
trace.RecordBackendTrace(trace.BackendTrace{
|
||||
ID: traceID,
|
||||
Timestamp: startTime,
|
||||
Duration: time.Since(startTime),
|
||||
Type: trace.BackendTraceDetection,
|
||||
|
||||
@@ -23,11 +23,19 @@ func ModelDetokenize(tokens []int32, loader *model.ModelLoader, modelConfig conf
|
||||
return schema.DetokenizeResponse{}, err
|
||||
}
|
||||
|
||||
release, err := AcquireGlobalBackendSlot()
|
||||
if err != nil {
|
||||
return schema.DetokenizeResponse{}, err
|
||||
}
|
||||
defer release()
|
||||
var startTime time.Time
|
||||
var traceID string
|
||||
if appConfig.EnableTracing {
|
||||
trace.InitBackendTracingIfEnabled(appConfig.TracingMaxItems, appConfig.TracingMaxBodyBytes)
|
||||
startTime = time.Now()
|
||||
traceID = trace.BeginBackendTrace(trace.BackendTrace{Timestamp: startTime, Type: trace.BackendTraceTokenize, ModelName: modelConfig.Name, Backend: modelConfig.Backend, Summary: "detokenize"})
|
||||
}
|
||||
defer trace.CancelBackendTrace(traceID)
|
||||
|
||||
resp, err := inferenceModel.Detokenize(appConfig.Context, &pb.DetokenizeRequest{Tokens: tokens})
|
||||
|
||||
@@ -43,6 +51,7 @@ func ModelDetokenize(tokens []int32, loader *model.ModelLoader, modelConfig conf
|
||||
}
|
||||
|
||||
trace.RecordBackendTrace(trace.BackendTrace{
|
||||
ID: traceID,
|
||||
Timestamp: startTime,
|
||||
Duration: time.Since(startTime),
|
||||
Type: trace.BackendTraceTokenize,
|
||||
|
||||
+93
-10
@@ -2,6 +2,7 @@ package backend
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
@@ -9,9 +10,80 @@ import (
|
||||
"github.com/mudler/LocalAI/core/trace"
|
||||
|
||||
"github.com/mudler/LocalAI/pkg/grpc"
|
||||
"github.com/mudler/LocalAI/pkg/grpc/proto"
|
||||
model "github.com/mudler/LocalAI/pkg/model"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
)
|
||||
|
||||
type embeddingPoolingCompatibilityError struct {
|
||||
message string
|
||||
}
|
||||
|
||||
func (e *embeddingPoolingCompatibilityError) Error() string {
|
||||
return e.message
|
||||
}
|
||||
|
||||
func poolingCompatibilityErrorf(format string, args ...any) error {
|
||||
return &embeddingPoolingCompatibilityError{message: fmt.Sprintf(format, args...)}
|
||||
}
|
||||
|
||||
// IsEmbeddingPoolingCompatibilityError reports errors caused by a requested
|
||||
// pooling scheme disagreeing with the layout declared by the loaded backend.
|
||||
// HTTP callers map these client-selectable incompatibilities to status 400.
|
||||
func IsEmbeddingPoolingCompatibilityError(err error) bool {
|
||||
var target *embeddingPoolingCompatibilityError
|
||||
return errors.As(err, &target)
|
||||
}
|
||||
|
||||
// finishEmbeddingResult applies the model's Go-side pooling scheme only when
|
||||
// the backend declares that it returned per-token vectors. Shape alone is not
|
||||
// sufficient: one raw token and one final vector are both reported as 1 x dim.
|
||||
// Legacy backends remain compatible with backend pooling, but cannot opt in to
|
||||
// Go-side pooling until they declare their result layout.
|
||||
func finishEmbeddingResult(res *proto.EmbeddingResult, modelConfig config.ModelConfig) ([]float32, error) {
|
||||
scheme := modelConfig.Pooling
|
||||
if scheme == "" || scheme == PoolingBackend {
|
||||
switch res.GetLayout() {
|
||||
case proto.EmbeddingLayout_EMBEDDING_LAYOUT_UNSPECIFIED, proto.EmbeddingLayout_EMBEDDING_LAYOUT_FINAL:
|
||||
return res.Embeddings, nil
|
||||
case proto.EmbeddingLayout_EMBEDDING_LAYOUT_PER_TOKEN:
|
||||
return nil, poolingCompatibilityErrorf(
|
||||
"pooling %q cannot pass through per-token embeddings: choose %q, %q or %q, or load the backend with pooling enabled",
|
||||
PoolingBackend, PoolingMean, PoolingLast, PoolingDecayedMean)
|
||||
default:
|
||||
return nil, poolingCompatibilityErrorf("pooling %q cannot use unknown embedding layout %d", PoolingBackend, res.GetLayout())
|
||||
}
|
||||
}
|
||||
switch res.GetLayout() {
|
||||
case proto.EmbeddingLayout_EMBEDDING_LAYOUT_PER_TOKEN:
|
||||
// Pool below after validating the reported matrix shape.
|
||||
case proto.EmbeddingLayout_EMBEDDING_LAYOUT_FINAL:
|
||||
return nil, poolingCompatibilityErrorf(
|
||||
"pooling %q needs per-token embeddings but this backend returned a final vector; configure raw per-token output if the backend supports it (llama.cpp: options [\"pooling:none\"])",
|
||||
scheme)
|
||||
case proto.EmbeddingLayout_EMBEDDING_LAYOUT_UNSPECIFIED:
|
||||
return nil, poolingCompatibilityErrorf(
|
||||
"pooling %q needs per-token embeddings but the backend did not declare its embedding layout: rebuild/update the backend to report EmbeddingResult.layout",
|
||||
scheme)
|
||||
default:
|
||||
return nil, poolingCompatibilityErrorf("pooling %q cannot use unknown embedding layout %d", scheme, res.GetLayout())
|
||||
}
|
||||
return PoolEmbeddingResult(res, scheme,
|
||||
float64(modelConfig.PoolingHalfLifeTokens),
|
||||
embdNormalizeFromOptions(modelConfig.Options))
|
||||
}
|
||||
|
||||
// mapEmbeddingGRPCError turns a gRPC ResourceExhausted — the per-token
|
||||
// payload of a very long conversation exceeding the 50MB message cap —
|
||||
// into an actionable message; everything else passes through unchanged.
|
||||
func mapEmbeddingGRPCError(err error) error {
|
||||
if status.Code(err) == codes.ResourceExhausted {
|
||||
return fmt.Errorf("conversation too long for per-token embeddings (gRPC message limit exceeded): %w", err)
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
// Embedder produces a fixed-dimension vector from a prompt. The
|
||||
// router's L2 embedding cache uses it to look up semantically-similar
|
||||
// past decisions.
|
||||
@@ -66,19 +138,19 @@ func ModelEmbedding(ctx context.Context, s string, tokens []int, loader *model.M
|
||||
|
||||
res, err := model.Embeddings(appConfig.Context, predictOptions)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return nil, mapEmbeddingGRPCError(err)
|
||||
}
|
||||
|
||||
return res.Embeddings, nil
|
||||
return finishEmbeddingResult(res, modelConfig)
|
||||
}
|
||||
predictOptions.Embeddings = s
|
||||
|
||||
res, err := model.Embeddings(appConfig.Context, predictOptions)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return nil, mapEmbeddingGRPCError(err)
|
||||
}
|
||||
|
||||
return res.Embeddings, nil
|
||||
return finishEmbeddingResult(res, modelConfig)
|
||||
}
|
||||
default:
|
||||
fn = func() ([]float32, error) {
|
||||
@@ -109,9 +181,15 @@ func ModelEmbedding(ctx context.Context, s string, tokens []int, loader *model.M
|
||||
traceData["input_tokens_count"] = len(tokens)
|
||||
}
|
||||
|
||||
startTime := time.Now()
|
||||
summary := trace.TruncateString(s, 200)
|
||||
if summary == "" {
|
||||
summary = fmt.Sprintf("tokens[%d]", len(tokens))
|
||||
}
|
||||
originalFn := wrappedFn
|
||||
wrappedFn = func() ([]float32, error) {
|
||||
startTime := time.Now()
|
||||
traceID := trace.BeginBackendTrace(trace.BackendTrace{Timestamp: startTime, Type: trace.BackendTraceEmbedding, ModelName: modelConfig.Name, Backend: modelConfig.Backend, Summary: summary})
|
||||
defer trace.CancelBackendTrace(traceID)
|
||||
result, err := originalFn()
|
||||
duration := time.Since(startTime)
|
||||
|
||||
@@ -122,12 +200,8 @@ func ModelEmbedding(ctx context.Context, s string, tokens []int, loader *model.M
|
||||
errStr = err.Error()
|
||||
}
|
||||
|
||||
summary := trace.TruncateString(s, 200)
|
||||
if summary == "" {
|
||||
summary = fmt.Sprintf("tokens[%d]", len(tokens))
|
||||
}
|
||||
|
||||
trace.RecordBackendTrace(trace.BackendTrace{
|
||||
ID: traceID,
|
||||
Timestamp: startTime,
|
||||
Duration: duration,
|
||||
Type: trace.BackendTraceEmbedding,
|
||||
@@ -141,6 +215,15 @@ func ModelEmbedding(ctx context.Context, s string, tokens []int, loader *model.M
|
||||
return result, err
|
||||
}
|
||||
}
|
||||
originalFn := wrappedFn
|
||||
wrappedFn = func() ([]float32, error) {
|
||||
release, err := AcquireGlobalBackendSlot()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer release()
|
||||
return originalFn()
|
||||
}
|
||||
|
||||
return wrappedFn, nil
|
||||
}
|
||||
@@ -30,11 +30,19 @@ func FaceAnalyze(
|
||||
return nil, fmt.Errorf("could not load face recognition model")
|
||||
}
|
||||
|
||||
release, err := AcquireGlobalBackendSlot()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer release()
|
||||
var startTime time.Time
|
||||
var traceID string
|
||||
if appConfig.EnableTracing {
|
||||
trace.InitBackendTracingIfEnabled(appConfig.TracingMaxItems, appConfig.TracingMaxBodyBytes)
|
||||
startTime = time.Now()
|
||||
traceID = trace.BeginBackendTrace(trace.BackendTrace{Timestamp: startTime, Type: trace.BackendTraceFaceAnalyze, ModelName: modelConfig.Name, Backend: modelConfig.Backend, Summary: "face analysis"})
|
||||
}
|
||||
defer trace.CancelBackendTrace(traceID)
|
||||
|
||||
res, err := faceModel.FaceAnalyze(ctx, &proto.FaceAnalyzeRequest{
|
||||
ModelIdentity: modelConfig.Model,
|
||||
@@ -49,6 +57,7 @@ func FaceAnalyze(
|
||||
errStr = err.Error()
|
||||
}
|
||||
trace.RecordBackendTrace(trace.BackendTrace{
|
||||
ID: traceID,
|
||||
Timestamp: startTime,
|
||||
Duration: time.Since(startTime),
|
||||
Type: trace.BackendTraceFaceAnalyze,
|
||||
|
||||
@@ -32,6 +32,11 @@ func FaceEmbed(
|
||||
|
||||
predictOpts := gRPCPredictOpts(modelConfig, loader.ModelPath)
|
||||
predictOpts.Images = []string{imgBase64}
|
||||
release, err := AcquireGlobalBackendSlot()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer release()
|
||||
|
||||
res, err := faceModel.Embeddings(ctx, predictOpts)
|
||||
if err != nil {
|
||||
|
||||
@@ -30,11 +30,19 @@ func FaceVerify(
|
||||
return nil, fmt.Errorf("could not load face recognition model")
|
||||
}
|
||||
|
||||
release, err := AcquireGlobalBackendSlot()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer release()
|
||||
var startTime time.Time
|
||||
var traceID string
|
||||
if appConfig.EnableTracing {
|
||||
trace.InitBackendTracingIfEnabled(appConfig.TracingMaxItems, appConfig.TracingMaxBodyBytes)
|
||||
startTime = time.Now()
|
||||
traceID = trace.BeginBackendTrace(trace.BackendTrace{Timestamp: startTime, Type: trace.BackendTraceFaceVerify, ModelName: modelConfig.Name, Backend: modelConfig.Backend, Summary: "face verification"})
|
||||
}
|
||||
defer trace.CancelBackendTrace(traceID)
|
||||
|
||||
res, err := faceModel.FaceVerify(ctx, &proto.FaceVerifyRequest{
|
||||
ModelIdentity: modelConfig.Model,
|
||||
@@ -50,6 +58,7 @@ func FaceVerify(
|
||||
errStr = err.Error()
|
||||
}
|
||||
trace.RecordBackendTrace(trace.BackendTrace{
|
||||
ID: traceID,
|
||||
Timestamp: startTime,
|
||||
Duration: time.Since(startTime),
|
||||
Type: trace.BackendTraceFaceVerify,
|
||||
|
||||
@@ -0,0 +1,72 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package backend
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/mudler/LocalAI/core/config"
|
||||
)
|
||||
|
||||
// BackendAdmissionError reports that the process-wide backend execution
|
||||
// ceiling is full. HTTP callers map it to 503; internal callers receive the
|
||||
// same typed error instead of silently queueing and growing in-flight state.
|
||||
type BackendAdmissionError struct {
|
||||
Limit int
|
||||
RetryAfter time.Duration
|
||||
}
|
||||
|
||||
func (e *BackendAdmissionError) Error() string {
|
||||
return fmt.Sprintf("backend inference capacity reached (max_concurrent=%d); retry after %s", e.Limit, e.RetryAfter)
|
||||
}
|
||||
|
||||
var backendAdmission = struct {
|
||||
sync.RWMutex
|
||||
limit int
|
||||
slots chan struct{}
|
||||
}{}
|
||||
|
||||
// ConfigureGlobalBackendAdmission sets the process-wide ceiling. It is called
|
||||
// during application construction, before backend work can begin.
|
||||
func ConfigureGlobalBackendAdmission(limit int) {
|
||||
if limit <= 0 {
|
||||
limit = config.DefaultMaxConcurrentBackendRequests
|
||||
}
|
||||
backendAdmission.Lock()
|
||||
backendAdmission.limit = limit
|
||||
backendAdmission.slots = make(chan struct{}, limit)
|
||||
backendAdmission.Unlock()
|
||||
}
|
||||
|
||||
// AcquireGlobalBackendSlot admits one backend operation without queueing.
|
||||
// Callers must invoke release on every completion path.
|
||||
func AcquireGlobalBackendSlot() (release func(), err error) {
|
||||
backendAdmission.RLock()
|
||||
limit, slots := backendAdmission.limit, backendAdmission.slots
|
||||
backendAdmission.RUnlock()
|
||||
if slots == nil {
|
||||
backendAdmission.Lock()
|
||||
if backendAdmission.slots == nil {
|
||||
backendAdmission.limit = config.DefaultMaxConcurrentBackendRequests
|
||||
backendAdmission.slots = make(chan struct{}, backendAdmission.limit)
|
||||
}
|
||||
limit, slots = backendAdmission.limit, backendAdmission.slots
|
||||
backendAdmission.Unlock()
|
||||
}
|
||||
select {
|
||||
case slots <- struct{}{}:
|
||||
var once sync.Once
|
||||
return func() { once.Do(func() { <-slots }) }, nil
|
||||
default:
|
||||
return nil, &BackendAdmissionError{Limit: limit, RetryAfter: time.Second}
|
||||
}
|
||||
}
|
||||
|
||||
// GlobalBackendInFlight is the current number of admitted backend operations.
|
||||
func GlobalBackendInFlight() int {
|
||||
backendAdmission.RLock()
|
||||
defer backendAdmission.RUnlock()
|
||||
return len(backendAdmission.slots)
|
||||
}
|
||||
@@ -0,0 +1,39 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package backend_test
|
||||
|
||||
import (
|
||||
"errors"
|
||||
|
||||
"github.com/mudler/LocalAI/core/backend"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
var _ = Describe("global backend admission", func() {
|
||||
BeforeEach(func() {
|
||||
backend.ConfigureGlobalBackendAdmission(1)
|
||||
})
|
||||
|
||||
It("rejects excess backend work without queueing", func() {
|
||||
release, err := backend.AcquireGlobalBackendSlot()
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(backend.GlobalBackendInFlight()).To(Equal(1))
|
||||
|
||||
_, err = backend.AcquireGlobalBackendSlot()
|
||||
var capacityErr *backend.BackendAdmissionError
|
||||
Expect(errors.As(err, &capacityErr)).To(BeTrue())
|
||||
Expect(capacityErr.Limit).To(Equal(1))
|
||||
|
||||
release()
|
||||
Expect(backend.GlobalBackendInFlight()).To(BeZero())
|
||||
})
|
||||
|
||||
It("makes release idempotent", func() {
|
||||
release, err := backend.AcquireGlobalBackendSlot()
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
release()
|
||||
release()
|
||||
Expect(backend.GlobalBackendInFlight()).To(BeZero())
|
||||
})
|
||||
})
|
||||
+13
-1
@@ -63,9 +63,11 @@ func ImageGeneration(ctx context.Context, height, width, step, seed int, positiv
|
||||
"destination": dst,
|
||||
}
|
||||
|
||||
startTime := time.Now()
|
||||
originalFn := fn
|
||||
fn = func() error {
|
||||
startTime := time.Now()
|
||||
traceID := trace.BeginBackendTrace(trace.BackendTrace{Timestamp: startTime, Type: trace.BackendTraceImageGeneration, ModelName: modelConfig.Name, Backend: modelConfig.Backend, Summary: trace.TruncateString(positive_prompt, 200)})
|
||||
defer trace.CancelBackendTrace(traceID)
|
||||
err := originalFn()
|
||||
duration := time.Since(startTime)
|
||||
|
||||
@@ -75,6 +77,7 @@ func ImageGeneration(ctx context.Context, height, width, step, seed int, positiv
|
||||
}
|
||||
|
||||
trace.RecordBackendTrace(trace.BackendTrace{
|
||||
ID: traceID,
|
||||
Timestamp: startTime,
|
||||
Duration: duration,
|
||||
Type: trace.BackendTraceImageGeneration,
|
||||
@@ -88,6 +91,15 @@ func ImageGeneration(ctx context.Context, height, width, step, seed int, positiv
|
||||
return err
|
||||
}
|
||||
}
|
||||
originalFn := fn
|
||||
fn = func() error {
|
||||
release, err := AcquireGlobalBackendSlot()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer release()
|
||||
return originalFn()
|
||||
}
|
||||
|
||||
return fn, nil
|
||||
}
|
||||
|
||||
+19
-1
@@ -378,9 +378,17 @@ func ModelInference(ctx context.Context, s string, messages schema.Messages, ima
|
||||
"xml_format_preset": c.FunctionsConfig.XMLFormatPreset,
|
||||
}
|
||||
|
||||
startTime := time.Now()
|
||||
originalFn := fn
|
||||
fn = func() (LLMResponse, error) {
|
||||
startTime := time.Now()
|
||||
traceID := trace.BeginBackendTrace(trace.BackendTrace{
|
||||
Timestamp: startTime,
|
||||
Type: trace.BackendTraceLLM,
|
||||
ModelName: c.Name,
|
||||
Backend: c.Backend,
|
||||
Summary: trace.GenerateLLMSummary(messages, s),
|
||||
})
|
||||
defer trace.CancelBackendTrace(traceID)
|
||||
resp, err := originalFn()
|
||||
duration := time.Since(startTime)
|
||||
|
||||
@@ -432,6 +440,7 @@ func ModelInference(ctx context.Context, s string, messages schema.Messages, ima
|
||||
}
|
||||
|
||||
trace.RecordBackendTrace(trace.BackendTrace{
|
||||
ID: traceID,
|
||||
Timestamp: startTime,
|
||||
Duration: duration,
|
||||
Type: trace.BackendTraceLLM,
|
||||
@@ -445,6 +454,15 @@ func ModelInference(ctx context.Context, s string, messages schema.Messages, ima
|
||||
return resp, err
|
||||
}
|
||||
}
|
||||
originalFn := fn
|
||||
fn = func() (LLMResponse, error) {
|
||||
release, err := AcquireGlobalBackendSlot()
|
||||
if err != nil {
|
||||
return LLMResponse{}, err
|
||||
}
|
||||
defer release()
|
||||
return originalFn()
|
||||
}
|
||||
|
||||
return fn, nil
|
||||
}
|
||||
|
||||
+13
-1
@@ -74,9 +74,11 @@ func Model3DGeneration(options Model3DGenerationOptions, loader *model.ModelLoad
|
||||
}
|
||||
}
|
||||
|
||||
startTime := time.Now()
|
||||
originalFn := fn
|
||||
fn = func() error {
|
||||
startTime := time.Now()
|
||||
traceID := trace.BeginBackendTrace(trace.BackendTrace{Timestamp: startTime, Type: traceType, ModelName: modelConfig.Name, Backend: modelConfig.Backend, Summary: trace.TruncateString(traceSummary, 200)})
|
||||
defer trace.CancelBackendTrace(traceID)
|
||||
err := originalFn()
|
||||
duration := time.Since(startTime)
|
||||
|
||||
@@ -86,6 +88,7 @@ func Model3DGeneration(options Model3DGenerationOptions, loader *model.ModelLoad
|
||||
}
|
||||
|
||||
trace.RecordBackendTrace(trace.BackendTrace{
|
||||
ID: traceID,
|
||||
Timestamp: startTime,
|
||||
Duration: duration,
|
||||
Type: traceType,
|
||||
@@ -99,6 +102,15 @@ func Model3DGeneration(options Model3DGenerationOptions, loader *model.ModelLoad
|
||||
return err
|
||||
}
|
||||
}
|
||||
originalFn := fn
|
||||
fn = func() error {
|
||||
release, err := AcquireGlobalBackendSlot()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer release()
|
||||
return originalFn()
|
||||
}
|
||||
|
||||
return fn, nil
|
||||
}
|
||||
@@ -30,6 +30,7 @@ var _ = Describe("ModelLoadTraceObserver", func() {
|
||||
}
|
||||
|
||||
BeforeEach(func() {
|
||||
backend.ConfigureGlobalBackendAdmission(64)
|
||||
appConfig = &config.ApplicationConfig{
|
||||
EnableTracing: true,
|
||||
TracingMaxItems: 64,
|
||||
@@ -39,7 +40,11 @@ var _ = Describe("ModelLoadTraceObserver", func() {
|
||||
})
|
||||
|
||||
It("records a model_load trace with the backend runtime on success", func() {
|
||||
backend.ModelLoadTraceObserver(appConfig)(successEvent)
|
||||
finish, err := backend.ModelLoadTraceObserver(appConfig)(successEvent)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(finish).NotTo(BeNil())
|
||||
Expect(trace.GetBackendTraces()).To(ContainElement(HaveField("Status", trace.BackendTraceRunning)))
|
||||
finish(successEvent)
|
||||
|
||||
Eventually(trace.GetBackendTraces).Should(HaveLen(1))
|
||||
got := trace.GetBackendTraces()[0]
|
||||
@@ -57,7 +62,10 @@ var _ = Describe("ModelLoadTraceObserver", func() {
|
||||
failed := successEvent
|
||||
failed.Err = errors.New("grpc service not ready")
|
||||
|
||||
backend.ModelLoadTraceObserver(appConfig)(failed)
|
||||
finish, err := backend.ModelLoadTraceObserver(appConfig)(failed)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(finish).NotTo(BeNil())
|
||||
finish(failed)
|
||||
|
||||
Consistently(trace.GetBackendTraces, "100ms", "20ms").Should(BeEmpty())
|
||||
})
|
||||
@@ -65,7 +73,10 @@ var _ = Describe("ModelLoadTraceObserver", func() {
|
||||
It("records nothing when tracing is disabled", func() {
|
||||
appConfig.EnableTracing = false
|
||||
|
||||
backend.ModelLoadTraceObserver(appConfig)(successEvent)
|
||||
finish, err := backend.ModelLoadTraceObserver(appConfig)(successEvent)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(finish).NotTo(BeNil())
|
||||
finish(successEvent)
|
||||
|
||||
Consistently(trace.GetBackendTraces, "100ms", "20ms").Should(BeEmpty())
|
||||
})
|
||||
|
||||
+30
-16
@@ -33,24 +33,38 @@ import (
|
||||
// backend's launcher path, which names the variant directory) — that is what
|
||||
// identifies WHICH build served the load. A stale installed backend is
|
||||
// invisible in the model config but obvious here.
|
||||
func ModelLoadTraceObserver(appConfig *config.ApplicationConfig) func(model.BackendLoadEvent) {
|
||||
return func(ev model.BackendLoadEvent) {
|
||||
if ev.Err != nil || !appConfig.EnableTracing {
|
||||
return
|
||||
func ModelLoadTraceObserver(appConfig *config.ApplicationConfig) func(model.BackendLoadEvent) (func(model.BackendLoadEvent), error) {
|
||||
return func(ev model.BackendLoadEvent) (func(model.BackendLoadEvent), error) {
|
||||
release, err := AcquireGlobalBackendSlot()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if !appConfig.EnableTracing {
|
||||
return func(model.BackendLoadEvent) { release() }, nil
|
||||
}
|
||||
trace.InitBackendTracingIfEnabled(appConfig.TracingMaxItems, appConfig.TracingMaxBodyBytes)
|
||||
trace.RecordBackendTrace(trace.BackendTrace{
|
||||
Timestamp: time.Now(),
|
||||
Duration: ev.Duration,
|
||||
Type: trace.BackendTraceModelLoad,
|
||||
ModelName: ev.ModelID,
|
||||
Backend: ev.Backend,
|
||||
Summary: "Model loaded",
|
||||
Data: map[string]any{
|
||||
"model_file": ev.ModelName,
|
||||
"backend_runtime": ev.BackendURI,
|
||||
},
|
||||
})
|
||||
started := time.Now()
|
||||
id := trace.BeginBackendTrace(trace.BackendTrace{Timestamp: started, Type: trace.BackendTraceModelLoad, ModelName: ev.ModelID, Backend: ev.Backend, Summary: "Loading model"})
|
||||
return func(done model.BackendLoadEvent) {
|
||||
defer release()
|
||||
if done.Err != nil {
|
||||
trace.CancelBackendTrace(id)
|
||||
return
|
||||
}
|
||||
trace.RecordBackendTrace(trace.BackendTrace{
|
||||
ID: id,
|
||||
Timestamp: started,
|
||||
Duration: done.Duration,
|
||||
Type: trace.BackendTraceModelLoad,
|
||||
ModelName: done.ModelID,
|
||||
Backend: done.Backend,
|
||||
Summary: "Model loaded",
|
||||
Data: map[string]any{
|
||||
"model_file": done.ModelName,
|
||||
"backend_runtime": done.BackendURI,
|
||||
},
|
||||
})
|
||||
}, nil
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,187 @@
|
||||
package backend
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"math"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/mudler/LocalAI/core/config"
|
||||
"github.com/mudler/LocalAI/pkg/grpc/proto"
|
||||
)
|
||||
|
||||
// Go-side embedding pooling schemes. The canonical strings live on the
|
||||
// config package (config.Pooling* — mirroring the ScoreNormalization*
|
||||
// pattern) so ModelConfig.Validate can reject unknown values without
|
||||
// importing this package; these aliases are the API the backend layer
|
||||
// (and the HTTP endpoints) program against.
|
||||
const (
|
||||
// PoolingBackend leaves pooling to the inference backend — the exact
|
||||
// pre-existing behavior of the embeddings path (also selected by "").
|
||||
PoolingBackend = config.PoolingBackend
|
||||
// PoolingMean averages the per-token vectors.
|
||||
PoolingMean = config.PoolingMean
|
||||
// PoolingLast selects the last token's vector.
|
||||
PoolingLast = config.PoolingLast
|
||||
// PoolingDecayedMean is a mean weighted toward the most recent tokens:
|
||||
// w_i = 2^(-(T-1-i)/H) for token i of T, with half-life H tokens.
|
||||
PoolingDecayedMean = config.PoolingDecayedMean
|
||||
|
||||
// DefaultPoolingHalfLifeTokens is the half-life used by
|
||||
// PoolingDecayedMean when the model config / request doesn't set one.
|
||||
DefaultPoolingHalfLifeTokens = 256
|
||||
)
|
||||
|
||||
// reshapeEmbeddings views the flat float payload of an EmbeddingResult as
|
||||
// tokens rows of dim columns. The gRPC contract packs vectors row-major
|
||||
// (vector 0 first), so row i aliases flat[i*dim : (i+1)*dim].
|
||||
func reshapeEmbeddings(flat []float32, tokens, dim int) ([][]float32, error) {
|
||||
if tokens <= 0 || dim <= 0 {
|
||||
return nil, fmt.Errorf("invalid embedding shape: %d vectors x %d dims", tokens, dim)
|
||||
}
|
||||
if len(flat) != tokens*dim {
|
||||
return nil, fmt.Errorf("embedding payload of %d floats does not match reported shape %d vectors x %d dims", len(flat), tokens, dim)
|
||||
}
|
||||
vecs := make([][]float32, tokens)
|
||||
for i := range vecs {
|
||||
vecs[i] = flat[i*dim : (i+1)*dim]
|
||||
}
|
||||
return vecs, nil
|
||||
}
|
||||
|
||||
// poolMean averages the per-token vectors with float64 accumulators.
|
||||
func poolMean(vecs [][]float32) []float32 {
|
||||
dim := len(vecs[0])
|
||||
acc := make([]float64, dim)
|
||||
for _, v := range vecs {
|
||||
for j, x := range v {
|
||||
acc[j] += float64(x)
|
||||
}
|
||||
}
|
||||
out := make([]float32, dim)
|
||||
for j := range out {
|
||||
out[j] = float32(acc[j] / float64(len(vecs)))
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// poolLast returns (a copy of) the last token's vector.
|
||||
func poolLast(vecs [][]float32) []float32 {
|
||||
last := vecs[len(vecs)-1]
|
||||
out := make([]float32, len(last))
|
||||
copy(out, last)
|
||||
return out
|
||||
}
|
||||
|
||||
// poolDecayedMean computes a weighted mean over the per-token vectors with
|
||||
// exponentially decaying weights anchored at the last token: token i of T
|
||||
// gets w_i = 2^(-(T-1-i)/halfLife), so the last token always weighs 1 and a
|
||||
// token halfLife positions earlier weighs 0.5. Accumulation is in float64.
|
||||
func poolDecayedMean(vecs [][]float32, halfLife float64) []float32 {
|
||||
dim := len(vecs[0])
|
||||
T := len(vecs)
|
||||
acc := make([]float64, dim)
|
||||
wsum := 0.0
|
||||
for i, v := range vecs {
|
||||
w := math.Exp2(-float64(T-1-i) / halfLife)
|
||||
wsum += w
|
||||
for j, x := range v {
|
||||
acc[j] += w * float64(x)
|
||||
}
|
||||
}
|
||||
out := make([]float32, dim)
|
||||
for j := range out {
|
||||
out[j] = float32(acc[j] / wsum)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// normalizeEmbedding is an exact port of llama.cpp's common_embd_normalize
|
||||
// (backend/cpp/llama-cpp/llama.cpp/common/common.cpp), applied after Go-side
|
||||
// pooling because per-token vectors arrive RAW from llama.cpp with
|
||||
// pooling:none (the server only normalizes vectors it pooled itself).
|
||||
// embdNorm: <0 none, 0 max-abs scaled to int16 range (/32760.0), 2 L2
|
||||
// (llama.cpp default), anything else p-norm with p=embdNorm (1 = taxicab).
|
||||
func normalizeEmbedding(v []float32, embdNorm int) []float32 {
|
||||
sum := 0.0
|
||||
switch {
|
||||
case embdNorm < 0: // no normalisation
|
||||
sum = 1.0
|
||||
case embdNorm == 0: // max absolute
|
||||
for _, x := range v {
|
||||
if a := math.Abs(float64(x)); sum < a {
|
||||
sum = a
|
||||
}
|
||||
}
|
||||
sum /= 32760.0 // make an int16 range
|
||||
case embdNorm == 2: // euclidean
|
||||
for _, x := range v {
|
||||
sum += float64(x) * float64(x)
|
||||
}
|
||||
sum = math.Sqrt(sum)
|
||||
default: // p-norm (euclidean is p-norm p=2)
|
||||
for _, x := range v {
|
||||
sum += math.Pow(math.Abs(float64(x)), float64(embdNorm))
|
||||
}
|
||||
sum = math.Pow(sum, 1.0/float64(embdNorm))
|
||||
}
|
||||
|
||||
// llama.cpp computes the reciprocal as a float32 and multiplies in
|
||||
// float32; mirror that so both paths yield bit-identical vectors.
|
||||
var norm float32
|
||||
if sum > 0.0 {
|
||||
norm = float32(1.0 / sum)
|
||||
}
|
||||
out := make([]float32, len(v))
|
||||
for i, x := range v {
|
||||
out[i] = x * norm
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// embdNormalizeFromOptions extracts the load-time embd_normalize backend
|
||||
// option ("embd_normalize:<n>", alias "embedding_normalize:<n>") the same
|
||||
// way the llama-cpp gRPC server parses it, so Go-side pooling normalizes
|
||||
// with the exact norm the backend would have applied had it pooled
|
||||
// server-side. Defaults to 2 (L2) like llama.cpp; unparsable values are
|
||||
// ignored (llama.cpp swallows std::stoi failures).
|
||||
func embdNormalizeFromOptions(options []string) int {
|
||||
embdNorm := 2
|
||||
for _, opt := range options {
|
||||
name, val, found := strings.Cut(opt, ":")
|
||||
if !found || (name != "embd_normalize" && name != "embedding_normalize") {
|
||||
continue
|
||||
}
|
||||
if n, err := strconv.Atoi(strings.TrimSpace(val)); err == nil {
|
||||
embdNorm = n
|
||||
}
|
||||
}
|
||||
return embdNorm
|
||||
}
|
||||
|
||||
// PoolEmbeddingResult reduces a per-token EmbeddingResult (the backend ran
|
||||
// with pooling:none) to a single vector using scheme, then normalizes it
|
||||
// with llama.cpp's common_embd_normalize semantics. halfLife only applies
|
||||
// to PoolingDecayedMean; non-positive values fall back to
|
||||
// DefaultPoolingHalfLifeTokens.
|
||||
func PoolEmbeddingResult(res *proto.EmbeddingResult, scheme string, halfLife float64, embdNorm int) ([]float32, error) {
|
||||
vecs, err := reshapeEmbeddings(res.GetEmbeddings(), int(res.GetTokens()), int(res.GetDim()))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var pooled []float32
|
||||
switch scheme {
|
||||
case PoolingMean:
|
||||
pooled = poolMean(vecs)
|
||||
case PoolingLast:
|
||||
pooled = poolLast(vecs)
|
||||
case PoolingDecayedMean:
|
||||
if halfLife <= 0 {
|
||||
halfLife = DefaultPoolingHalfLifeTokens
|
||||
}
|
||||
pooled = poolDecayedMean(vecs, halfLife)
|
||||
default:
|
||||
return nil, fmt.Errorf("unknown Go-side pooling scheme %q (expected %q, %q or %q)", scheme, PoolingMean, PoolingLast, PoolingDecayedMean)
|
||||
}
|
||||
return normalizeEmbedding(pooled, embdNorm), nil
|
||||
}
|
||||
@@ -0,0 +1,229 @@
|
||||
package backend
|
||||
|
||||
import (
|
||||
"math"
|
||||
|
||||
"github.com/mudler/LocalAI/core/config"
|
||||
"github.com/mudler/LocalAI/pkg/grpc/proto"
|
||||
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
var _ = Describe("Go-side embedding pooling", func() {
|
||||
Describe("reshapeEmbeddings", func() {
|
||||
It("views a flat payload as tokens x dim rows", func() {
|
||||
vecs, err := reshapeEmbeddings([]float32{1, 2, 3, 4, 5, 6}, 2, 3)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(vecs).To(Equal([][]float32{{1, 2, 3}, {4, 5, 6}}))
|
||||
})
|
||||
|
||||
It("rejects a payload that does not match the reported shape", func() {
|
||||
_, err := reshapeEmbeddings([]float32{1, 2, 3, 4, 5}, 2, 3)
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(err.Error()).To(ContainSubstring("does not match reported shape"))
|
||||
})
|
||||
|
||||
It("rejects a non-positive shape", func() {
|
||||
_, err := reshapeEmbeddings(nil, 0, 3)
|
||||
Expect(err).To(HaveOccurred())
|
||||
_, err = reshapeEmbeddings(nil, 3, 0)
|
||||
Expect(err).To(HaveOccurred())
|
||||
})
|
||||
})
|
||||
|
||||
Describe("poolMean", func() {
|
||||
It("averages the per-token vectors", func() {
|
||||
Expect(poolMean([][]float32{{1, 2}, {3, 4}})).To(Equal([]float32{2, 3}))
|
||||
})
|
||||
})
|
||||
|
||||
Describe("poolLast", func() {
|
||||
It("returns the last token's vector", func() {
|
||||
Expect(poolLast([][]float32{{1, 2}, {3, 4}})).To(Equal([]float32{3, 4}))
|
||||
})
|
||||
})
|
||||
|
||||
Describe("poolDecayedMean", func() {
|
||||
It("weights tokens by 2^(-(T-1-i)/H): H=1, T=3 gives [0.25 0.5 1]/1.75", func() {
|
||||
vecs := [][]float32{{1, 0}, {0, 1}, {1, 1}}
|
||||
got := poolDecayedMean(vecs, 1)
|
||||
Expect(got[0]).To(BeNumerically("~", 1.25/1.75, 1e-6))
|
||||
Expect(got[1]).To(BeNumerically("~", 1.5/1.75, 1e-6))
|
||||
})
|
||||
|
||||
It("approaches the plain mean as the half-life grows", func() {
|
||||
vecs := [][]float32{{1, 0}, {0, 1}, {1, 1}}
|
||||
got := poolDecayedMean(vecs, 1e12)
|
||||
want := poolMean(vecs)
|
||||
Expect(got[0]).To(BeNumerically("~", want[0], 1e-6))
|
||||
Expect(got[1]).To(BeNumerically("~", want[1], 1e-6))
|
||||
})
|
||||
})
|
||||
|
||||
Describe("single-token conversations", func() {
|
||||
It("agree across mean, last and decayed_mean", func() {
|
||||
vecs := [][]float32{{3, 4}}
|
||||
Expect(poolMean(vecs)).To(Equal([]float32{3, 4}))
|
||||
Expect(poolLast(vecs)).To(Equal([]float32{3, 4}))
|
||||
got := poolDecayedMean(vecs, 256)
|
||||
Expect(got[0]).To(BeNumerically("~", 3, 1e-6))
|
||||
Expect(got[1]).To(BeNumerically("~", 4, 1e-6))
|
||||
})
|
||||
})
|
||||
|
||||
Describe("normalizeEmbedding (common_embd_normalize port)", func() {
|
||||
v := []float32{3, -4}
|
||||
|
||||
It("passes through untouched for negative embd_norm", func() {
|
||||
Expect(normalizeEmbedding(v, -1)).To(Equal([]float32{3, -4}))
|
||||
})
|
||||
|
||||
It("scales to the int16 range for embd_norm 0 (max-abs)", func() {
|
||||
// max-abs = 4, sum = 4/32760, norm = 32760/4 = 8190
|
||||
got := normalizeEmbedding(v, 0)
|
||||
Expect(got[0]).To(BeNumerically("~", 3*8190.0, 1e-2))
|
||||
Expect(got[1]).To(BeNumerically("~", -4*8190.0, 1e-2))
|
||||
})
|
||||
|
||||
It("applies the taxicab norm for embd_norm 1", func() {
|
||||
got := normalizeEmbedding(v, 1)
|
||||
Expect(got[0]).To(BeNumerically("~", 3.0/7.0, 1e-6))
|
||||
Expect(got[1]).To(BeNumerically("~", -4.0/7.0, 1e-6))
|
||||
})
|
||||
|
||||
It("applies the L2 norm for embd_norm 2", func() {
|
||||
got := normalizeEmbedding(v, 2)
|
||||
Expect(got[0]).To(BeNumerically("~", 0.6, 1e-6))
|
||||
Expect(got[1]).To(BeNumerically("~", -0.8, 1e-6))
|
||||
})
|
||||
|
||||
It("applies a p-norm for embd_norm > 2", func() {
|
||||
p3 := math.Cbrt(27 + 64) // (|3|^3 + |-4|^3)^(1/3)
|
||||
got := normalizeEmbedding(v, 3)
|
||||
Expect(got[0]).To(BeNumerically("~", 3.0/p3, 1e-5))
|
||||
Expect(got[1]).To(BeNumerically("~", -4.0/p3, 1e-5))
|
||||
})
|
||||
|
||||
It("maps the all-zero vector to all zeros instead of dividing by zero", func() {
|
||||
Expect(normalizeEmbedding([]float32{0, 0, 0}, 2)).To(Equal([]float32{0, 0, 0}))
|
||||
})
|
||||
})
|
||||
|
||||
Describe("embdNormalizeFromOptions", func() {
|
||||
It("defaults to 2 (L2) like llama.cpp", func() {
|
||||
Expect(embdNormalizeFromOptions(nil)).To(Equal(2))
|
||||
Expect(embdNormalizeFromOptions([]string{"pooling:none", "gpu"})).To(Equal(2))
|
||||
})
|
||||
|
||||
It("parses embd_normalize and its embedding_normalize alias", func() {
|
||||
Expect(embdNormalizeFromOptions([]string{"embd_normalize:0"})).To(Equal(0))
|
||||
Expect(embdNormalizeFromOptions([]string{"embedding_normalize:-1"})).To(Equal(-1))
|
||||
Expect(embdNormalizeFromOptions([]string{"embd_normalize: 3"})).To(Equal(3))
|
||||
})
|
||||
|
||||
It("keeps the default when the value does not parse", func() {
|
||||
Expect(embdNormalizeFromOptions([]string{"embd_normalize:junk"})).To(Equal(2))
|
||||
Expect(embdNormalizeFromOptions([]string{"embd_normalize"})).To(Equal(2))
|
||||
})
|
||||
})
|
||||
|
||||
Describe("PoolEmbeddingResult", func() {
|
||||
res := &proto.EmbeddingResult{
|
||||
Embeddings: []float32{1, 2, 3, 4},
|
||||
Tokens: 2,
|
||||
Dim: 2,
|
||||
}
|
||||
|
||||
It("reshapes, pools and normalizes", func() {
|
||||
got, err := PoolEmbeddingResult(res, PoolingMean, 0, -1)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(got).To(Equal([]float32{2, 3}))
|
||||
})
|
||||
|
||||
It("L2-normalizes by default norm 2", func() {
|
||||
got, err := PoolEmbeddingResult(res, PoolingLast, 0, 2)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(got[0]).To(BeNumerically("~", 0.6, 1e-6))
|
||||
Expect(got[1]).To(BeNumerically("~", 0.8, 1e-6))
|
||||
})
|
||||
|
||||
It("rejects unknown pooling schemes", func() {
|
||||
_, err := PoolEmbeddingResult(res, "sideways", 0, 2)
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(err.Error()).To(ContainSubstring("unknown Go-side pooling scheme"))
|
||||
})
|
||||
|
||||
It("propagates shape mismatches", func() {
|
||||
bad := &proto.EmbeddingResult{Embeddings: []float32{1, 2, 3}, Tokens: 2, Dim: 2}
|
||||
_, err := PoolEmbeddingResult(bad, PoolingMean, 0, 2)
|
||||
Expect(err).To(HaveOccurred())
|
||||
})
|
||||
})
|
||||
|
||||
Describe("finishEmbeddingResult", func() {
|
||||
finish := func(res *proto.EmbeddingResult, scheme string) ([]float32, error) {
|
||||
cfg := config.ModelConfig{}
|
||||
cfg.Pooling = scheme
|
||||
cfg.Options = []string{"embd_normalize:-1"}
|
||||
return finishEmbeddingResult(res, cfg)
|
||||
}
|
||||
final := func() *proto.EmbeddingResult {
|
||||
return &proto.EmbeddingResult{
|
||||
Embeddings: []float32{1, 2},
|
||||
Tokens: 1,
|
||||
Dim: 2,
|
||||
Layout: proto.EmbeddingLayout_EMBEDDING_LAYOUT_FINAL,
|
||||
}
|
||||
}
|
||||
perToken := func() *proto.EmbeddingResult {
|
||||
return &proto.EmbeddingResult{
|
||||
Embeddings: []float32{1, 2, 3, 4},
|
||||
Tokens: 2,
|
||||
Dim: 2,
|
||||
Layout: proto.EmbeddingLayout_EMBEDDING_LAYOUT_PER_TOKEN,
|
||||
}
|
||||
}
|
||||
|
||||
It("passes a final vector through to backend pooling", func() {
|
||||
got, err := finish(final(), PoolingBackend)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(got).To(Equal([]float32{1, 2}))
|
||||
})
|
||||
|
||||
It("rejects Go-side pooling after the backend returned a final vector", func() {
|
||||
_, err := finish(final(), PoolingMean)
|
||||
Expect(err).To(MatchError(ContainSubstring("final vector")))
|
||||
})
|
||||
|
||||
It("pools a backend-declared per-token matrix", func() {
|
||||
got, err := finish(perToken(), PoolingMean)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(got).To(Equal([]float32{2, 3}))
|
||||
})
|
||||
|
||||
It("rejects backend pass-through of a per-token matrix", func() {
|
||||
_, err := finish(perToken(), PoolingBackend)
|
||||
Expect(err).To(MatchError(ContainSubstring("cannot pass through per-token")))
|
||||
})
|
||||
|
||||
It("allows legacy layout only for backend pooling", func() {
|
||||
legacy := &proto.EmbeddingResult{Embeddings: []float32{1, 2}, Tokens: 1, Dim: 2}
|
||||
got, err := finish(legacy, PoolingBackend)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(got).To(Equal([]float32{1, 2}))
|
||||
|
||||
_, err = finish(legacy, PoolingMean)
|
||||
Expect(err).To(MatchError(ContainSubstring("did not declare")))
|
||||
})
|
||||
|
||||
It("fails closed for an unknown layout value", func() {
|
||||
res := perToken()
|
||||
res.Layout = proto.EmbeddingLayout(99)
|
||||
_, err := finish(res, PoolingMean)
|
||||
Expect(err).To(MatchError(ContainSubstring("unknown embedding layout")))
|
||||
_, err = finish(res, PoolingBackend)
|
||||
Expect(err).To(MatchError(ContainSubstring("unknown embedding layout")))
|
||||
})
|
||||
})
|
||||
})
|
||||
@@ -71,11 +71,19 @@ func Rerank(ctx context.Context, request *proto.RerankRequest, loader *model.Mod
|
||||
return nil, fmt.Errorf("could not load rerank model")
|
||||
}
|
||||
|
||||
release, err := AcquireGlobalBackendSlot()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer release()
|
||||
var startTime time.Time
|
||||
var traceID string
|
||||
if appConfig.EnableTracing {
|
||||
trace.InitBackendTracingIfEnabled(appConfig.TracingMaxItems, appConfig.TracingMaxBodyBytes)
|
||||
startTime = time.Now()
|
||||
traceID = trace.BeginBackendTrace(trace.BackendTrace{Timestamp: startTime, Type: trace.BackendTraceRerank, ModelName: modelConfig.Name, Backend: modelConfig.Backend, Summary: trace.TruncateString(request.Query, 200)})
|
||||
}
|
||||
defer trace.CancelBackendTrace(traceID)
|
||||
|
||||
// Stamped here, not at the HTTP handler: this is the function that also
|
||||
// builds ModelOptions from the same config, so the two values are equal by
|
||||
@@ -91,6 +99,7 @@ func Rerank(ctx context.Context, request *proto.RerankRequest, loader *model.Mod
|
||||
}
|
||||
|
||||
trace.RecordBackendTrace(trace.BackendTrace{
|
||||
ID: traceID,
|
||||
Timestamp: startTime,
|
||||
Duration: time.Since(startTime),
|
||||
Type: trace.BackendTraceRerank,
|
||||
|
||||
@@ -96,15 +96,23 @@ func ModelScore(prompt string, candidates []string, opts ScoreOptions, loader *m
|
||||
return nil, fmt.Errorf("Score: candidates must be non-empty")
|
||||
}
|
||||
return func(ctx context.Context) ([]CandidateScore, error) {
|
||||
release, err := AcquireGlobalBackendSlot()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer release()
|
||||
// Surface score calls in the Traces UI alongside the LLM calls
|
||||
// they typically gate (router classifier, eval scoring). Without
|
||||
// this, a router-classified request shows only the downstream LLM
|
||||
// trace with no record of the classification that picked it.
|
||||
var startTime time.Time
|
||||
var traceID string
|
||||
if appConfig.EnableTracing {
|
||||
trace.InitBackendTracingIfEnabled(appConfig.TracingMaxItems, appConfig.TracingMaxBodyBytes)
|
||||
startTime = time.Now()
|
||||
traceID = trace.BeginBackendTrace(trace.BackendTrace{Timestamp: startTime, Type: trace.BackendTraceScore, ModelName: modelConfig.Name, Backend: modelConfig.Backend, Summary: trace.TruncateString(prompt, 200)})
|
||||
}
|
||||
defer trace.CancelBackendTrace(traceID)
|
||||
resp, err := b.Score(ctx, &pb.ScoreRequest{
|
||||
ModelIdentity: modelConfig.Model,
|
||||
Prompt: prompt,
|
||||
@@ -120,6 +128,7 @@ func ModelScore(prompt string, candidates []string, opts ScoreOptions, loader *m
|
||||
errStr = err.Error()
|
||||
}
|
||||
trace.RecordBackendTrace(trace.BackendTrace{
|
||||
ID: traceID,
|
||||
Timestamp: startTime,
|
||||
Duration: time.Since(startTime),
|
||||
Type: trace.BackendTraceScore,
|
||||
|
||||
@@ -101,11 +101,23 @@ func SoundGeneration(
|
||||
req.Instrumental = instrumental
|
||||
}
|
||||
|
||||
release, err := AcquireGlobalBackendSlot()
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
defer release()
|
||||
var startTime time.Time
|
||||
var traceID string
|
||||
if appConfig.EnableTracing {
|
||||
trace.InitBackendTracingIfEnabled(appConfig.TracingMaxItems, appConfig.TracingMaxBodyBytes)
|
||||
startTime = time.Now()
|
||||
summary := trace.TruncateString(text, 200)
|
||||
if summary == "" && caption != "" {
|
||||
summary = trace.TruncateString(caption, 200)
|
||||
}
|
||||
traceID = trace.BeginBackendTrace(trace.BackendTrace{Timestamp: startTime, Type: trace.BackendTraceSoundGeneration, ModelName: modelConfig.Name, Backend: modelConfig.Backend, Summary: summary})
|
||||
}
|
||||
defer trace.CancelBackendTrace(traceID)
|
||||
|
||||
res, err := soundGenModel.SoundGeneration(ctx, req)
|
||||
|
||||
@@ -135,6 +147,7 @@ func SoundGeneration(
|
||||
}
|
||||
|
||||
trace.RecordBackendTrace(trace.BackendTrace{
|
||||
ID: traceID,
|
||||
Timestamp: startTime,
|
||||
Duration: time.Since(startTime),
|
||||
Type: trace.BackendTraceSoundGeneration,
|
||||
|
||||
+137
-19
@@ -15,13 +15,23 @@ import (
|
||||
)
|
||||
|
||||
// VectorStore is the narrowed KNN store used by the router's embedding
|
||||
// cache. Search returns the top-1 match (cosine similarity in [-1, 1])
|
||||
// and the serialised payload, or ok=false on a clean miss.
|
||||
// cache and the KNN classifier. Search returns the top-1 match (cosine
|
||||
// similarity in [-1, 1]) and the serialised payload, or ok=false on a
|
||||
// clean miss. SearchK returns up to k nearest neighbours ordered by
|
||||
// descending similarity; an empty slice is a clean miss.
|
||||
type VectorStore interface {
|
||||
Search(ctx context.Context, vec []float32) (similarity float64, payload []byte, ok bool, err error)
|
||||
SearchK(ctx context.Context, vec []float32, k int) ([]Neighbor, error)
|
||||
Insert(ctx context.Context, vec []float32, payload []byte) error
|
||||
}
|
||||
|
||||
// Neighbor is one SearchK result — the stored payload and its cosine
|
||||
// similarity to the query vector.
|
||||
type Neighbor struct {
|
||||
Similarity float64
|
||||
Payload []byte
|
||||
}
|
||||
|
||||
// NewVectorStore returns a VectorStore backed by the local-store
|
||||
// gRPC backend, namespaced by storeName so two routers don't collide.
|
||||
// cl resolves the per-store model config (backend + options); it may be nil,
|
||||
@@ -45,40 +55,73 @@ func (s *localVectorStore) backend(_ context.Context) (grpc.Backend, error) {
|
||||
return StoreBackend(s.loader, s.appConfig, s.cl, s.storeName, "")
|
||||
}
|
||||
|
||||
func (s *localVectorStore) Search(ctx context.Context, vec []float32) (sim float64, payload []byte, ok bool, err error) {
|
||||
start := time.Now()
|
||||
// Search is the top-1 special case of SearchK; delegating keeps the
|
||||
// backend-load/Find/trace plumbing in one place (SearchK records the
|
||||
// identically-shaped trace, so /api/backend-traces sees no difference).
|
||||
func (s *localVectorStore) Search(ctx context.Context, vec []float32) (float64, []byte, bool, error) {
|
||||
neighbors, err := s.SearchK(ctx, vec, 1)
|
||||
if err != nil || len(neighbors) == 0 {
|
||||
return 0, nil, false, err
|
||||
}
|
||||
return neighbors[0].Similarity, neighbors[0].Payload, true, nil
|
||||
}
|
||||
|
||||
func (s *localVectorStore) SearchK(ctx context.Context, vec []float32, k int) (neighbors []Neighbor, err error) {
|
||||
outcome := "hit"
|
||||
defer func() {
|
||||
s.recordTrace(start, "search", len(vec), sim, outcome, err)
|
||||
}()
|
||||
sim := 0.0
|
||||
be, berr := s.backend(ctx)
|
||||
if berr != nil {
|
||||
outcome = "backend_load_error"
|
||||
return 0, nil, false, fmt.Errorf("vector store load: %w", berr)
|
||||
err = fmt.Errorf("vector store load: %w", berr)
|
||||
s.recordTrace("", time.Now(), "search", len(vec), 0, outcome, err)
|
||||
return nil, err
|
||||
}
|
||||
_, values, similarities, ferr := store.Find(ctx, be, vec, 1)
|
||||
release, err := AcquireGlobalBackendSlot()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer release()
|
||||
start := time.Now()
|
||||
traceID := s.beginTrace(start, "search")
|
||||
defer func() {
|
||||
s.recordTrace(traceID, start, "search", len(vec), sim, outcome, err)
|
||||
}()
|
||||
_, values, similarities, ferr := store.Find(ctx, be, vec, k)
|
||||
if ferr != nil {
|
||||
outcome = "find_error"
|
||||
return 0, nil, false, fmt.Errorf("vector store find: %w", ferr)
|
||||
return nil, fmt.Errorf("vector store find: %w", ferr)
|
||||
}
|
||||
if len(values) == 0 || len(similarities) == 0 {
|
||||
if len(values) == 0 {
|
||||
outcome = "miss"
|
||||
return 0, nil, false, nil
|
||||
return nil, nil
|
||||
}
|
||||
return float64(similarities[0]), values[0], true, nil
|
||||
neighbors = make([]Neighbor, 0, len(values))
|
||||
for i, v := range values {
|
||||
neighbors = append(neighbors, Neighbor{Similarity: float64(similarities[i]), Payload: v})
|
||||
}
|
||||
sim = neighbors[0].Similarity
|
||||
return neighbors, nil
|
||||
}
|
||||
|
||||
func (s *localVectorStore) Insert(ctx context.Context, vec []float32, payload []byte) (err error) {
|
||||
start := time.Now()
|
||||
outcome := "ok"
|
||||
defer func() {
|
||||
s.recordTrace(start, "insert", len(vec), 0, outcome, err)
|
||||
}()
|
||||
be, berr := s.backend(ctx)
|
||||
if berr != nil {
|
||||
outcome = "backend_load_error"
|
||||
return fmt.Errorf("vector store load: %w", berr)
|
||||
err = fmt.Errorf("vector store load: %w", berr)
|
||||
s.recordTrace("", time.Now(), "insert", len(vec), 0, outcome, err)
|
||||
return err
|
||||
}
|
||||
release, err := AcquireGlobalBackendSlot()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer release()
|
||||
start := time.Now()
|
||||
traceID := s.beginTrace(start, "insert")
|
||||
defer func() {
|
||||
s.recordTrace(traceID, start, "insert", len(vec), 0, outcome, err)
|
||||
}()
|
||||
if serr := store.SetSingle(ctx, be, vec, payload); serr != nil {
|
||||
outcome = "insert_error"
|
||||
return serr
|
||||
@@ -86,12 +129,86 @@ func (s *localVectorStore) Insert(ctx context.Context, vec []float32, payload []
|
||||
return nil
|
||||
}
|
||||
|
||||
// InsertBatch upserts many vectors in one gRPC round-trip. Not part of
|
||||
// the VectorStore interface — the corpus manager type-asserts for it
|
||||
// and falls back to per-entry Insert on stores that lack it.
|
||||
func (s *localVectorStore) InsertBatch(ctx context.Context, vecs [][]float32, payloads [][]byte) (err error) {
|
||||
outcome := "ok"
|
||||
dim := 0
|
||||
if len(vecs) > 0 {
|
||||
dim = len(vecs[0])
|
||||
}
|
||||
be, berr := s.backend(ctx)
|
||||
if berr != nil {
|
||||
outcome = "backend_load_error"
|
||||
err = fmt.Errorf("vector store load: %w", berr)
|
||||
s.recordTrace("", time.Now(), "insert_batch", dim, 0, outcome, err)
|
||||
return err
|
||||
}
|
||||
release, err := AcquireGlobalBackendSlot()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer release()
|
||||
start := time.Now()
|
||||
traceID := s.beginTrace(start, "insert_batch")
|
||||
defer func() {
|
||||
s.recordTrace(traceID, start, "insert_batch", dim, 0, outcome, err)
|
||||
}()
|
||||
if serr := store.SetCols(ctx, be, vecs, payloads); serr != nil {
|
||||
outcome = "insert_error"
|
||||
return serr
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Delete removes vectors by key. Optional capability like InsertBatch;
|
||||
// used by the corpus manager's Clear so a wiped corpus also leaves the
|
||||
// live index.
|
||||
func (s *localVectorStore) Delete(ctx context.Context, vecs [][]float32) (err error) {
|
||||
outcome := "ok"
|
||||
dim := 0
|
||||
if len(vecs) > 0 {
|
||||
dim = len(vecs[0])
|
||||
}
|
||||
be, berr := s.backend(ctx)
|
||||
if berr != nil {
|
||||
outcome = "backend_load_error"
|
||||
err = fmt.Errorf("vector store load: %w", berr)
|
||||
s.recordTrace("", time.Now(), "delete", dim, 0, outcome, err)
|
||||
return err
|
||||
}
|
||||
release, err := AcquireGlobalBackendSlot()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer release()
|
||||
start := time.Now()
|
||||
traceID := s.beginTrace(start, "delete")
|
||||
defer func() {
|
||||
s.recordTrace(traceID, start, "delete", dim, 0, outcome, err)
|
||||
}()
|
||||
if serr := store.DeleteCols(ctx, be, vecs); serr != nil {
|
||||
outcome = "delete_error"
|
||||
return serr
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// recordTrace surfaces vector-store calls in /api/backend-traces, including
|
||||
// the backend-load-failure path that otherwise vanishes into an xlog.Warn.
|
||||
// modelName uses the store namespace (e.g. "router-cache-smart-router") so
|
||||
// admins can tell which router's cache misbehaved; the backend is always
|
||||
// "local-store" and can't disambiguate.
|
||||
func (s *localVectorStore) recordTrace(start time.Time, op string, vecDim int, sim float64, outcome string, err error) {
|
||||
func (s *localVectorStore) beginTrace(start time.Time, op string) string {
|
||||
if s.appConfig == nil || !s.appConfig.EnableTracing {
|
||||
return ""
|
||||
}
|
||||
trace.InitBackendTracingIfEnabled(s.appConfig.TracingMaxItems, s.appConfig.TracingMaxBodyBytes)
|
||||
return trace.BeginBackendTrace(trace.BackendTrace{Timestamp: start, Type: trace.BackendTraceVectorStore, ModelName: s.storeName, Backend: model.LocalStoreBackend, Summary: op})
|
||||
}
|
||||
|
||||
func (s *localVectorStore) recordTrace(traceID string, start time.Time, op string, vecDim int, sim float64, outcome string, err error) {
|
||||
if s.appConfig == nil || !s.appConfig.EnableTracing {
|
||||
return
|
||||
}
|
||||
@@ -115,6 +232,7 @@ func (s *localVectorStore) recordTrace(start time.Time, op string, vecDim int, s
|
||||
data["similarity"] = sim
|
||||
}
|
||||
trace.RecordBackendTrace(trace.BackendTrace{
|
||||
ID: traceID,
|
||||
Timestamp: start,
|
||||
Duration: time.Since(start),
|
||||
Type: trace.BackendTraceVectorStore,
|
||||
|
||||
@@ -81,11 +81,19 @@ func ModelTokenClassify(text string, opts TokenClassifyOptions, loader *model.Mo
|
||||
return nil, err
|
||||
}
|
||||
return func(ctx context.Context) ([]TokenEntity, error) {
|
||||
release, err := AcquireGlobalBackendSlot()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer release()
|
||||
var startTime time.Time
|
||||
var traceID string
|
||||
if appConfig.EnableTracing {
|
||||
trace.InitBackendTracingIfEnabled(appConfig.TracingMaxItems, appConfig.TracingMaxBodyBytes)
|
||||
startTime = time.Now()
|
||||
traceID = trace.BeginBackendTrace(trace.BackendTrace{Timestamp: startTime, Type: trace.BackendTraceTokenClassify, ModelName: modelConfig.Name, Backend: modelConfig.Backend, Summary: trace.TruncateString(text, 200)})
|
||||
}
|
||||
defer trace.CancelBackendTrace(traceID)
|
||||
resp, err := inferenceModel.TokenClassify(ctx, &pb.TokenClassifyRequest{
|
||||
ModelIdentity: modelConfig.Model,
|
||||
Text: text,
|
||||
@@ -93,7 +101,9 @@ func ModelTokenClassify(text string, opts TokenClassifyOptions, loader *model.Mo
|
||||
})
|
||||
entities := tokenClassifyResponseToEntities(resp)
|
||||
if appConfig.EnableTracing {
|
||||
trace.RecordBackendTrace(tokenClassifyTrace(modelConfig, text, opts.Threshold, entities, startTime, err))
|
||||
bt := tokenClassifyTrace(modelConfig, text, opts.Threshold, entities, startTime, err)
|
||||
bt.ID = traceID
|
||||
trace.RecordBackendTrace(bt)
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
|
||||
@@ -26,6 +26,11 @@ func TokenMetrics(
|
||||
if model == nil {
|
||||
return nil, fmt.Errorf("could not loadmodel model")
|
||||
}
|
||||
release, err := AcquireGlobalBackendSlot()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer release()
|
||||
|
||||
res, err := model.GetTokenMetrics(ctx, &proto.MetricsRequest{})
|
||||
|
||||
|
||||
@@ -39,11 +39,19 @@ func ModelTokenize(s string, loader *model.ModelLoader, modelConfig config.Model
|
||||
predictOptions := gRPCPredictOpts(modelConfig, loader.ModelPath)
|
||||
predictOptions.Prompt = s
|
||||
|
||||
release, err := AcquireGlobalBackendSlot()
|
||||
if err != nil {
|
||||
return schema.TokenizeResponse{}, err
|
||||
}
|
||||
defer release()
|
||||
var startTime time.Time
|
||||
var traceID string
|
||||
if appConfig.EnableTracing {
|
||||
trace.InitBackendTracingIfEnabled(appConfig.TracingMaxItems, appConfig.TracingMaxBodyBytes)
|
||||
startTime = time.Now()
|
||||
traceID = trace.BeginBackendTrace(trace.BackendTrace{Timestamp: startTime, Type: trace.BackendTraceTokenize, ModelName: modelConfig.Name, Backend: modelConfig.Backend, Summary: trace.TruncateString(s, 200)})
|
||||
}
|
||||
defer trace.CancelBackendTrace(traceID)
|
||||
|
||||
// tokenize the string
|
||||
resp, err := inferenceModel.TokenizeString(appConfig.Context, predictOptions)
|
||||
@@ -57,6 +65,7 @@ func ModelTokenize(s string, loader *model.ModelLoader, modelConfig config.Model
|
||||
tokenCount := tokenizeTokenCount(resp)
|
||||
|
||||
trace.RecordBackendTrace(trace.BackendTrace{
|
||||
ID: traceID,
|
||||
Timestamp: startTime,
|
||||
Duration: time.Since(startTime),
|
||||
Type: trace.BackendTraceTokenize,
|
||||
|
||||
Loaded 100 of 420 files, more files were not shown because too many files have changed in this diff.
Show more
Reference in new issue
Block a user